diff --git a/.gitignore b/.gitignore index dff64e3c9e9..572830d35f6 100644 --- a/.gitignore +++ b/.gitignore @@ -28,6 +28,8 @@ litellm/tests/config_*.yaml litellm/tests/langfuse.log langfuse.log .langfuse.log +.pin_list.txt +.cov_new.xml litellm/tests/test_custom_logger.py litellm/tests/langfuse.log litellm/tests/dynamo*.log @@ -120,4 +122,5 @@ crash.log crash.*.log # .terraform.lock.hcl is intentionally NOT ignored — it pins provider versions # and should be committed. -.vscode \ No newline at end of file +.vscode +.pin_list.txt diff --git a/tests/test_litellm/proxy/proxy_server/test_background_health.py b/tests/test_litellm/proxy/proxy_server/test_background_health.py index ad6b4016461..ee8d8b22779 100644 --- a/tests/test_litellm/proxy/proxy_server/test_background_health.py +++ b/tests/test_litellm/proxy/proxy_server/test_background_health.py @@ -1 +1,513 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for proxy_server background health-check helpers. + +Pins covered: +- ``_get_process_rss_mb`` +- ``_rss_mb_for_log`` +- ``_run_direct_health_check_with_instrumentation`` +- ``_schedule_background_health_check_db_save`` +- ``_get_endpoint_exception_status`` +- ``_write_health_state_to_router_cache`` +- ``_adaptive_router_flusher_loop`` +- ``_run_background_health_check`` +""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.proxy_server as proxy_server +from litellm.proxy.proxy_server import ( + _adaptive_router_flusher_loop, + _get_endpoint_exception_status, + _get_process_rss_mb, + _run_background_health_check, + _run_direct_health_check_with_instrumentation, + _rss_mb_for_log, + _schedule_background_health_check_db_save, + _write_health_state_to_router_cache, +) + +from .conftest import normalize + +# --------------------------------------------------------------------------- +# _get_process_rss_mb +# --------------------------------------------------------------------------- + + +def test_get_process_rss_mb_returns_positive_float(): + value = _get_process_rss_mb() + assert value is not None + assert normalize( + { + "value_present": value is not None, + "value_type": type(value).__name__, + "positive": value > 0, + } + ) == { + "value_present": True, + "value_type": "float", + "positive": True, + } + + +def test_get_process_rss_mb_returns_none_when_resource_raises(monkeypatch): + import resource + + def _boom(*_args, **_kwargs): + raise OSError("nope") + + monkeypatch.setattr(resource, "getrusage", _boom) + assert _get_process_rss_mb() is None + + +# --------------------------------------------------------------------------- +# _rss_mb_for_log +# --------------------------------------------------------------------------- + + +def test_rss_mb_for_log_formats_numeric_value(monkeypatch): + monkeypatch.setattr(proxy_server, "_get_process_rss_mb", lambda: 100.5) + result = _rss_mb_for_log() + assert normalize( + { + "format": result, + "is_string": isinstance(result, str), + "contains_mb": "100.50" in result, + } + ) == { + "format": "100.50", + "is_string": True, + "contains_mb": True, + } + + +def test_rss_mb_for_log_unknown_when_rss_missing(monkeypatch): + monkeypatch.setattr(proxy_server, "_get_process_rss_mb", lambda: None) + assert _rss_mb_for_log() == "unknown" + + +# --------------------------------------------------------------------------- +# _run_direct_health_check_with_instrumentation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_run_direct_health_check_with_instrumentation_returns_results( + monkeypatch, +): + expected = (["healthy_ep"], ["unhealthy_ep"], {"m1": Exception("boom")}) + + async def _fake_perform(model_list, details, max_concurrency, **kwargs): + return expected + + monkeypatch.setattr(proxy_server, "perform_health_check", _fake_perform) + monkeypatch.setattr( + proxy_server, + "health_check_filter_kwargs_from_general_settings", + lambda _gs: {}, + ) + + healthy, unhealthy, exceptions = ( + await _run_direct_health_check_with_instrumentation( + model_list=[{"model_name": "gpt-4"}], + details=False, + max_concurrency=1, + instrumentation_context={"source": "test"}, + ) + ) + + assert normalize( + { + "healthy": healthy, + "unhealthy": unhealthy, + "exception_keys": list(exceptions.keys()), + } + ) == { + "healthy": ["healthy_ep"], + "unhealthy": ["unhealthy_ep"], + "exception_keys": ["m1"], + } + + +@pytest.mark.asyncio +async def test_run_direct_health_check_raises_non_kwarg_typeerror(monkeypatch): + async def _boom(model_list, details, max_concurrency, **kwargs): + raise TypeError("totally unrelated") + + monkeypatch.setattr(proxy_server, "perform_health_check", _boom) + monkeypatch.setattr( + proxy_server, + "health_check_filter_kwargs_from_general_settings", + lambda _gs: {}, + ) + + with pytest.raises(TypeError): + await _run_direct_health_check_with_instrumentation( + model_list=[], + details=False, + max_concurrency=1, + instrumentation_context={}, + ) + + +# --------------------------------------------------------------------------- +# _schedule_background_health_check_db_save +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_creates_task(monkeypatch): + captured = {} + + async def _fake_save( + prisma_client, + model_list, + healthy, + unhealthy, + start_time, + checked_by, + ): + captured["prisma_client"] = prisma_client + captured["model_list"] = model_list + captured["healthy"] = healthy + captured["unhealthy"] = unhealthy + captured["checked_by"] = checked_by + + import litellm.proxy.health_endpoints._health_endpoints as he + + monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + + prisma_client = MagicMock() + shared_manager = SimpleNamespace(pod_id="pod-xyz") + + _schedule_background_health_check_db_save( + prisma_client=prisma_client, + shared_health_manager=shared_manager, + model_list=[{"model_name": "gpt-4"}], + healthy_endpoints=[{"model_id": "h1"}], + unhealthy_endpoints=[{"model_id": "u1"}], + ) + + await asyncio.sleep(0) + + assert normalize( + { + "prisma_present": captured.get("prisma_client") is prisma_client, + "checked_by": captured.get("checked_by"), + "healthy": captured.get("healthy"), + "unhealthy": captured.get("unhealthy"), + } + ) == { + "prisma_present": True, + "checked_by": "pod-xyz", + "healthy": [{"model_id": "h1"}], + "unhealthy": [{"model_id": "u1"}], + } + + +def test_schedule_background_health_check_db_save_noop_when_prisma_none(): + _schedule_background_health_check_db_save( + prisma_client=None, + shared_health_manager=None, + model_list=[], + healthy_endpoints=[], + unhealthy_endpoints=[], + ) + + +@pytest.mark.asyncio +async def test_schedule_background_health_check_db_save_invalid_no_event_loop_raises( + monkeypatch, +): + async def _fake_save(*_args, **_kwargs): + return None + + import litellm.proxy.health_endpoints._health_endpoints as he + + monkeypatch.setattr(he, "_save_background_health_checks_to_db", _fake_save) + + def _broken_create_task(_coro): + raise RuntimeError("no running event loop") + + monkeypatch.setattr(asyncio, "create_task", _broken_create_task) + + with pytest.raises(RuntimeError): + _schedule_background_health_check_db_save( + prisma_client=MagicMock(), + shared_health_manager=None, + model_list=[], + healthy_endpoints=[], + unhealthy_endpoints=[], + ) + + +# --------------------------------------------------------------------------- +# _get_endpoint_exception_status +# --------------------------------------------------------------------------- + + +def test_get_endpoint_exception_status_prefers_live_exception(): + endpoint = {"model_id": "m1", "exception_status": 999} + exceptions = {"m1": SimpleNamespace(status_code=429)} + status = _get_endpoint_exception_status(endpoint, exceptions) + assert normalize( + { + "input_endpoint": endpoint, + "exceptions_keys": list(exceptions.keys()), + "status": status, + } + ) == { + "input_endpoint": {"model_id": "m1", "exception_status": 999}, + "exceptions_keys": ["m1"], + "status": 429, + } + + +def test_get_endpoint_exception_status_falls_back_to_stored_int(): + endpoint = {"model_id": "m-missing", "exception_status": 503} + assert _get_endpoint_exception_status(endpoint, {}) == 503 + + +def test_get_endpoint_exception_status_default_500_when_no_data(): + assert _get_endpoint_exception_status({}, {}) == 500 + + +def test_get_endpoint_exception_status_invalid_endpoint_type_raises(): + with pytest.raises(AttributeError): + _get_endpoint_exception_status(None, {}) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _write_health_state_to_router_cache +# --------------------------------------------------------------------------- + + +def test_write_health_state_to_router_cache_sets_states(monkeypatch): + fake_router = MagicMock() + fake_router.enable_health_check_routing = True + fake_router.health_check_ignore_transient_errors = False + fake_router.cooldown_time = 30 + fake_router.health_state_cache = MagicMock() + + monkeypatch.setattr(proxy_server, "llm_router", fake_router) + + fake_states = {"m1": {"is_healthy": True}, "m2": {"is_healthy": False}} + + import litellm.proxy.health_check as hc + + monkeypatch.setattr(hc, "build_deployment_health_states", lambda **_kw: fake_states) + + import litellm.router_utils.cooldown_handlers as cd + + monkeypatch.setattr(cd, "_set_cooldown_deployments", lambda **_kw: None) + + import litellm.router_utils.router_callbacks.track_deployment_metrics as tdm + + monkeypatch.setattr( + tdm, + "increment_deployment_failures_for_current_minute", + lambda **_kw: None, + ) + + healthy = [{"model_id": "m1"}] + unhealthy = [{"model_id": "m2"}] + exceptions = {"m2": SimpleNamespace(status_code=500)} + + _write_health_state_to_router_cache(healthy, unhealthy, exceptions) + + fake_router.health_state_cache.set_deployment_health_states.assert_called_once_with( + fake_states + ) + + call_args = fake_router.health_state_cache.set_deployment_health_states.call_args[ + 0 + ][0] + assert normalize( + { + "states_keys": sorted(call_args.keys()), + "m1_healthy": call_args["m1"]["is_healthy"], + "m2_healthy": call_args["m2"]["is_healthy"], + } + ) == { + "states_keys": ["m1", "m2"], + "m1_healthy": True, + "m2_healthy": False, + } + + +def test_write_health_state_to_router_cache_noop_when_router_none(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + _write_health_state_to_router_cache([], [], {}) + + +def test_write_health_state_to_router_cache_swallows_internal_failures(monkeypatch): + """The function logs and swallows exceptions so a bad cache call never crashes the loop.""" + fake_router = MagicMock() + fake_router.enable_health_check_routing = True + fake_router.health_check_ignore_transient_errors = False + fake_router.health_state_cache.set_deployment_health_states.side_effect = ( + RuntimeError("cache exploded") + ) + + monkeypatch.setattr(proxy_server, "llm_router", fake_router) + + import litellm.proxy.health_check as hc + + monkeypatch.setattr( + hc, + "build_deployment_health_states", + lambda **_kw: {"m1": {"is_healthy": True}}, + ) + + _write_health_state_to_router_cache([{"model_id": "m1"}], [], {}) + + +# --------------------------------------------------------------------------- +# _adaptive_router_flusher_loop +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_adaptive_router_flusher_loop_flushes_each_router(monkeypatch): + fake_ar = MagicMock() + fake_ar._state_loaded = True + fake_ar.queue.flush_state_to_db = AsyncMock() + fake_ar.queue.flush_session_to_db = AsyncMock() + + fake_router = MagicMock() + fake_router.adaptive_routers = {"alpha": fake_ar} + + monkeypatch.setattr(proxy_server, "llm_router", fake_router) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + + # asyncio.sleep is awaited at the top of every iteration; raise CancelledError + # on the SECOND call so the first iteration completes its flush work. + call_count = {"n": 0} + _real_sleep = asyncio.sleep + + async def _short_sleep(_seconds): + call_count["n"] += 1 + if call_count["n"] >= 2: + raise asyncio.CancelledError() + await _real_sleep(0) + + monkeypatch.setattr(proxy_server.asyncio, "sleep", _short_sleep) + + with pytest.raises(asyncio.CancelledError): + await _adaptive_router_flusher_loop() + + assert fake_ar.queue.flush_state_to_db.await_count == 1 + assert fake_ar.queue.flush_session_to_db.await_count == 1 + + +@pytest.mark.asyncio +async def test_adaptive_router_flusher_loop_times_out_when_sleep_real(monkeypatch): + """Confirms the loop is infinite — wait_for must raise TimeoutError.""" + monkeypatch.setattr(proxy_server, "llm_router", MagicMock(adaptive_routers={})) + monkeypatch.setattr(proxy_server, "prisma_client", None) + + # Bind the real asyncio.sleep before the patch so the replacement does not + # recurse into itself. + _real_sleep = asyncio.sleep + + async def _instant_sleep(_seconds): + await _real_sleep(0) + + monkeypatch.setattr(proxy_server.asyncio, "sleep", _instant_sleep) + + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(_adaptive_router_flusher_loop(), timeout=0.2) + + +# --------------------------------------------------------------------------- +# _run_background_health_check +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_run_background_health_check_returns_immediately_when_interval_invalid( + monkeypatch, +): + monkeypatch.setattr(proxy_server, "health_check_interval", None) + + result = await _run_background_health_check() + + assert normalize( + { + "result_is_none": result is None, + "loop_active": proxy_server.background_health_check_loop_active, + "interval": proxy_server.health_check_interval, + } + ) == { + "result_is_none": True, + "loop_active": False, + "interval": None, + } + + +@pytest.mark.asyncio +async def test_run_background_health_check_runs_one_cycle_then_cancels(monkeypatch): + monkeypatch.setattr(proxy_server, "health_check_interval", 60) + monkeypatch.setattr(proxy_server, "health_check_concurrency", 1) + monkeypatch.setattr(proxy_server, "health_check_details", True) + monkeypatch.setattr(proxy_server, "use_shared_health_check", False) + monkeypatch.setattr(proxy_server, "redis_usage_cache", None) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(proxy_server, "background_health_check_loop_active", False) + monkeypatch.setattr( + proxy_server, + "llm_model_list", + [{"model_name": "gpt-4", "model_info": {}}], + ) + monkeypatch.setattr( + proxy_server, + "health_check_results", + {"healthy_endpoints": [], "unhealthy_endpoints": []}, + ) + + async def _fake_direct(*_a, **_kw): + return ([{"model_id": "h"}], [{"model_id": "u"}], {}) + + monkeypatch.setattr( + proxy_server, + "_run_direct_health_check_with_instrumentation", + _fake_direct, + ) + monkeypatch.setattr( + proxy_server, "_schedule_background_health_check_db_save", lambda *a, **kw: None + ) + monkeypatch.setattr( + proxy_server, "_write_health_state_to_router_cache", lambda *a, **kw: None + ) + monkeypatch.setattr( + proxy_server, + "health_check_filter_kwargs_from_general_settings", + lambda _gs: {}, + ) + + sleep_calls = {"n": 0} + + async def _stop_sleep(_seconds): + sleep_calls["n"] += 1 + raise asyncio.CancelledError() + + monkeypatch.setattr(proxy_server.asyncio, "sleep", _stop_sleep) + + with pytest.raises(asyncio.CancelledError): + await _run_background_health_check() + + assert normalize( + { + "healthy_count": proxy_server.health_check_results["healthy_count"], + "unhealthy_count": proxy_server.health_check_results["unhealthy_count"], + "sleep_invoked": sleep_calls["n"] >= 1, + } + ) == { + "healthy_count": 1, + "unhealthy_count": 1, + "sleep_invoked": True, + } diff --git a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py index ad6b4016461..cf92f9cd12b 100644 --- a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py +++ b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py @@ -1 +1,222 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for the proxy_server exception handlers. + +Pins covered: +- ``openai_exception_handler`` +- ``_close_dangling_otel_server_span`` +- ``otel_request_validation_exception_handler`` +- ``otel_unhandled_exception_handler`` +""" + +from __future__ import annotations + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException +from fastapi.exceptions import RequestValidationError + +from litellm.proxy._types import ProxyException +from litellm.proxy.proxy_server import ( + _close_dangling_otel_server_span, + openai_exception_handler, + otel_request_validation_exception_handler, + otel_unhandled_exception_handler, +) + +from .conftest import normalize + + +def _make_request(parent_otel_span=None): + state = SimpleNamespace(parent_otel_span=parent_otel_span) + return SimpleNamespace(state=state) + + +# --------------------------------------------------------------------------- +# openai_exception_handler +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_openai_exception_handler_returns_mapped_payload(): + exc = ProxyException( + message="bad input", + type="invalid_request_error", + param="model", + code=400, + ) + request = _make_request() + + response = await openai_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 400 + assert normalize(body) == { + "error": { + "message": "bad input", + "type": "invalid_request_error", + "param": "model", + "code": "400", + } + } + + +@pytest.mark.asyncio +async def test_openai_exception_handler_invalid_empty_code_defaults_to_500(): + """openai_exception_handler falls back to 500 when ``code`` is falsy. + + Constructing via __new__ bypasses __init__ — the production __init__ always + coerces None to the string "None", which is truthy. To exercise the falsy + fallback branch we hand-craft an exception with an empty code.""" + exc = ProxyException.__new__(ProxyException) + exc.message = "boom" + exc.type = "server_error" + exc.param = None + exc.openai_code = None + exc.code = "" + exc.headers = {} + exc.provider_specific_fields = None + request = _make_request() + + response = await openai_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 500 + assert body == { + "error": { + "message": "boom", + "type": "server_error", + "param": None, + "code": "", + } + } + + +# --------------------------------------------------------------------------- +# _close_dangling_otel_server_span +# --------------------------------------------------------------------------- + + +def test_close_dangling_otel_server_span_records_status_and_ends(monkeypatch): + """Happy path: with a logger and an active span, the handler sets the + response status, marks ERROR (>=400), ends the span, and clears state.""" + import litellm.proxy.proxy_server as ps + + span = MagicMock() + fake_logger = MagicMock() + monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False) + request = _make_request(parent_otel_span=span) + + _close_dangling_otel_server_span(request=request, status_code=502) + + observed = { + "status_attr_called": fake_logger.set_response_status_code_attribute.called, + "set_status_called": span.set_status.called, + "ended": span.end.called, + "state_cleared": request.state.parent_otel_span is None, + } + assert normalize(observed) == { + "status_attr_called": True, + "set_status_called": True, + "ended": True, + "state_cleared": True, + } + + +def test_close_dangling_otel_server_span_missing_span_is_noop_error(): + """When parent_otel_span is missing the call short-circuits — no error.""" + request = _make_request(parent_otel_span=None) + + result = _close_dangling_otel_server_span(request=request, status_code=200) + assert result is None + assert request.state.parent_otel_span is None + + +def test_close_dangling_otel_server_span_logger_raises_state_cleared_error(monkeypatch): + """Logger raising is caught; state.parent_otel_span is cleared regardless.""" + import litellm.proxy.proxy_server as ps + + span = MagicMock() + fake_logger = MagicMock() + fake_logger.set_response_status_code_attribute.side_effect = RuntimeError("boom") + monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False) + request = _make_request(parent_otel_span=span) + + _close_dangling_otel_server_span(request=request, status_code=500) + + assert request.state.parent_otel_span is None + + +# --------------------------------------------------------------------------- +# otel_request_validation_exception_handler +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_otel_request_validation_exception_handler_returns_422_detail(): + errors = [{"loc": ["body", "model"], "msg": "field required", "type": "missing"}] + exc = RequestValidationError(errors) + request = _make_request() + + response = await otel_request_validation_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 422 + assert normalize(body) == {"detail": exc.errors()} + + +@pytest.mark.asyncio +async def test_otel_request_validation_exception_handler_empty_errors_invalid_payload(): + """An empty error list still returns 422 — the validator emitted nothing + but the handler must not crash and the body must remain well-formed.""" + exc = RequestValidationError([]) + request = _make_request() + + response = await otel_request_validation_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 422 + assert body == {"detail": []} + + +# --------------------------------------------------------------------------- +# otel_unhandled_exception_handler +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_otel_unhandled_exception_handler_returns_500_generic_payload(): + exc = RuntimeError("kaboom") + request = _make_request() + + response = await otel_unhandled_exception_handler(request=request, exc=exc) + body = json.loads(response.body) + + assert response.status_code == 500 + assert normalize(body) == { + "error": { + "message": "Internal server error", + "type": "internal_server_error", + } + } + + +@pytest.mark.asyncio +async def test_otel_unhandled_exception_handler_reraises_proxy_exception_error(): + """ProxyException / HTTPException / RequestValidationError are re-raised + so the dedicated handler runs.""" + exc = ProxyException(message="m", type="t", param="p", code=403) + request = _make_request() + + with pytest.raises(ProxyException): + await otel_unhandled_exception_handler(request=request, exc=exc) + + +@pytest.mark.asyncio +async def test_otel_unhandled_exception_handler_reraises_http_exception_invalid(): + request = _make_request() + with pytest.raises(HTTPException): + await otel_unhandled_exception_handler( + request=request, exc=HTTPException(status_code=418, detail="teapot") + ) diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index ad6b4016461..0b733401b59 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -1 +1,564 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for proxy_server lifecycle, helpers, and small utilities. + +Pins covered: +- ``proxy_startup_event`` +- ``proxy_shutdown_event`` +- ``_initialize_shared_aiohttp_session`` +- ``cleanup_router_config_variables`` +- ``save_worker_config`` +- ``initialize`` +- ``load_from_azure_key_vault`` +- ``cost_tracking`` +- ``check_request_disconnection`` +- ``_resolve_typed_dict_type`` +- ``_resolve_pydantic_type`` +- ``get_litellm_model_info`` +- ``run_ollama_serve`` +""" + +from __future__ import annotations + +import asyncio +import inspect +import json +import os +from typing import List, Optional, Union +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import FastAPI, HTTPException +from pydantic import BaseModel +from typing_extensions import TypedDict + +import litellm.proxy.proxy_server as ps +from litellm.proxy.proxy_server import ( + _initialize_shared_aiohttp_session, + _resolve_pydantic_type, + _resolve_typed_dict_type, + check_request_disconnection, + cleanup_router_config_variables, + cost_tracking, + get_litellm_model_info, + initialize, + load_from_azure_key_vault, + proxy_shutdown_event, + proxy_startup_event, + run_ollama_serve, + save_worker_config, +) + +from .conftest import normalize + +# --------------------------------------------------------------------------- +# cleanup_router_config_variables +# --------------------------------------------------------------------------- + + +def test_cleanup_router_config_variables_resets_globals(monkeypatch): + monkeypatch.setattr(ps, "master_key", "sk-sentinel", raising=False) + monkeypatch.setattr(ps, "user_config_file_path", "/tmp/config.yaml", raising=False) + monkeypatch.setattr(ps, "user_custom_auth", lambda x: x, raising=False) + monkeypatch.setattr(ps, "health_check_interval", 42, raising=False) + monkeypatch.setattr(ps, "prisma_client", MagicMock(), raising=False) + + cleanup_router_config_variables() + + observed = { + "master_key": ps.master_key, + "user_config_file_path": ps.user_config_file_path, + "user_custom_auth": ps.user_custom_auth, + "health_check_interval": ps.health_check_interval, + "prisma_client": ps.prisma_client, + } + assert normalize(observed) == { + "master_key": None, + "user_config_file_path": None, + "user_custom_auth": None, + "health_check_interval": None, + "prisma_client": None, + } + + +def test_cleanup_router_config_variables_fails_on_unknown_attr_raises(): + """The function only writes documented globals — accessing a non-existent + one after cleanup should still raise AttributeError.""" + cleanup_router_config_variables() + with pytest.raises(AttributeError): + _ = ps.this_attribute_should_not_exist_xyz + + +# --------------------------------------------------------------------------- +# proxy_shutdown_event +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_proxy_shutdown_event_disconnects_prisma_and_resets(monkeypatch): + fake_prisma = MagicMock() + fake_prisma.disconnect = AsyncMock() + monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False) + monkeypatch.setattr(ps, "master_key", "sk-x", raising=False) + + fake_jwt = MagicMock() + fake_jwt.close = AsyncMock() + monkeypatch.setattr(ps, "jwt_handler", fake_jwt, raising=False) + monkeypatch.setattr(ps, "db_writer_client", None, raising=False) + + import litellm + + monkeypatch.setattr(litellm, "cache", None, raising=False) + monkeypatch.setattr(litellm, "success_callback", [], raising=False) + + await proxy_shutdown_event() + + observed = { + "disconnect_called": fake_prisma.disconnect.await_count == 1, + "jwt_closed": fake_jwt.close.await_count == 1, + "master_key_reset": ps.master_key, + "prisma_reset": ps.prisma_client, + } + assert normalize(observed) == { + "disconnect_called": True, + "jwt_closed": True, + "master_key_reset": None, + "prisma_reset": None, + } + + +@pytest.mark.asyncio +async def test_proxy_shutdown_event_prisma_disconnect_raises_error(monkeypatch): + fake_prisma = MagicMock() + fake_prisma.disconnect = AsyncMock(side_effect=RuntimeError("db gone")) + monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False) + + fake_jwt = MagicMock() + fake_jwt.close = AsyncMock() + monkeypatch.setattr(ps, "jwt_handler", fake_jwt, raising=False) + + import litellm + + monkeypatch.setattr(litellm, "cache", None, raising=False) + monkeypatch.setattr(litellm, "success_callback", [], raising=False) + + with pytest.raises(RuntimeError, match="db gone"): + await proxy_shutdown_event() + + +# --------------------------------------------------------------------------- +# _initialize_shared_aiohttp_session +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_initialize_shared_aiohttp_session_returns_client_session(): + from aiohttp import ClientSession + + session = await _initialize_shared_aiohttp_session() + try: + observed = { + "is_client_session": isinstance(session, ClientSession), + "is_closed": session.closed, + "has_connector": session.connector is not None, + } + assert normalize(observed) == { + "is_client_session": True, + "is_closed": False, + "has_connector": True, + } + finally: + if session is not None: + await session.close() + + +@pytest.mark.asyncio +async def test_initialize_shared_aiohttp_session_aiohttp_missing_returns_none_on_failure( + monkeypatch, +): + """If aiohttp import fails, the function catches and returns None — no raise.""" + import builtins + + real_import = builtins.__import__ + + def _raise_for_aiohttp(name, *args, **kwargs): + if name == "aiohttp": + raise ImportError("simulated missing aiohttp") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _raise_for_aiohttp) + result = await _initialize_shared_aiohttp_session() + assert result is None + + +# --------------------------------------------------------------------------- +# save_worker_config +# --------------------------------------------------------------------------- + + +def test_save_worker_config_writes_json_to_environ(monkeypatch): + monkeypatch.delenv("WORKER_CONFIG", raising=False) + + save_worker_config(model="gpt-4", config="/tmp/c.yaml", debug=True) + + payload = json.loads(os.environ["WORKER_CONFIG"]) + assert normalize(payload) == { + "model": "gpt-4", + "config": "/tmp/c.yaml", + "debug": True, + } + + +def test_save_worker_config_invalid_no_kwargs_yields_empty(monkeypatch): + monkeypatch.delenv("WORKER_CONFIG", raising=False) + + save_worker_config() + assert os.environ["WORKER_CONFIG"] == "{}" + + +# --------------------------------------------------------------------------- +# initialize +# --------------------------------------------------------------------------- + + +def test_initialize_signature_is_async_with_expected_params(): + sig = inspect.signature(initialize) + # Hard-coded so a signature change (param added/removed) trips the gate. + expected_param_count = 17 + observed = { + "is_async": inspect.iscoroutinefunction(initialize), + "param_count": len(sig.parameters), + "has_model": "model" in sig.parameters, + "has_config": "config" in sig.parameters, + } + assert normalize(observed) == { + "is_async": True, + "param_count": expected_param_count, + "has_model": True, + "has_config": True, + } + + +@pytest.mark.asyncio +async def test_initialize_invalid_unexpected_kwarg_raises_type_error(): + with pytest.raises(TypeError): + await initialize(this_is_not_a_real_kwarg=True) + + +# --------------------------------------------------------------------------- +# load_from_azure_key_vault +# --------------------------------------------------------------------------- + + +def test_load_from_azure_key_vault_disabled_no_side_effect(monkeypatch): + import litellm + + sentinel_secret_mgr = object() + monkeypatch.setattr( + litellm, "secret_manager_client", sentinel_secret_mgr, raising=False + ) + + result = load_from_azure_key_vault(use_azure_key_vault=False) + + observed = { + "return_value": result, + "secret_manager_unchanged": litellm.secret_manager_client + is sentinel_secret_mgr, + "called_with": False, + } + assert normalize(observed) == { + "return_value": None, + "secret_manager_unchanged": True, + "called_with": False, + } + + +def test_load_from_azure_key_vault_missing_uri_failure_is_swallowed(monkeypatch): + """Enabled but AZURE_KEY_VAULT_URI unset / azure libs likely unavailable — + function catches Exception and does not raise.""" + monkeypatch.delenv("AZURE_KEY_VAULT_URI", raising=False) + + result = load_from_azure_key_vault(use_azure_key_vault=True) + assert result is None + + +# --------------------------------------------------------------------------- +# cost_tracking +# --------------------------------------------------------------------------- + + +def test_cost_tracking_adds_two_callbacks_when_prisma_set(monkeypatch): + import litellm + + fake_prisma = MagicMock() + monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False) + monkeypatch.setattr(litellm, "callbacks", [], raising=False) + monkeypatch.setattr(litellm, "_async_success_callback", [], raising=False) + + before_callbacks = len(litellm.callbacks) + before_async = len(litellm._async_success_callback) + + cost_tracking() + + observed = { + "added_to_callbacks": len(litellm.callbacks) - before_callbacks, + "added_to_async_success": len(litellm._async_success_callback) - before_async, + "prisma_was_set": True, + } + assert normalize(observed) == { + "added_to_callbacks": 1, + "added_to_async_success": 1, + "prisma_was_set": True, + } + + +def test_cost_tracking_no_op_when_prisma_missing(monkeypatch): + """Without a prisma_client cost_tracking is a no-op — not an error.""" + import litellm + + monkeypatch.setattr(ps, "prisma_client", None, raising=False) + monkeypatch.setattr(litellm, "callbacks", [], raising=False) + monkeypatch.setattr(litellm, "_async_success_callback", [], raising=False) + + cost_tracking() + + assert litellm.callbacks == [] + assert litellm._async_success_callback == [] + + +# --------------------------------------------------------------------------- +# check_request_disconnection +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_check_request_disconnection_cancels_task_and_raises_499(monkeypatch): + monkeypatch.setattr(ps.asyncio, "sleep", AsyncMock(return_value=None)) + + request = MagicMock() + request.is_disconnected = AsyncMock(return_value=True) + task = MagicMock() + + raised_status = None + try: + await check_request_disconnection(request=request, llm_api_call_task=task) + except HTTPException as exc: + raised_status = exc.status_code + + observed = { + "raised_status": raised_status, + "cancel_called": task.cancel.called, + "is_async": inspect.iscoroutinefunction(check_request_disconnection), + } + assert normalize(observed) == { + "raised_status": 499, + "cancel_called": True, + "is_async": True, + } + + +@pytest.mark.asyncio +async def test_check_request_disconnection_invalid_when_connected_times_out(monkeypatch): + """With a connected request the function loops for up to 10 minutes — + wrap in wait_for and assert it times out. Patch ``asyncio.sleep`` so the + loop spins without real wall-clock waits.""" + import litellm.proxy.proxy_server as ps + + request = MagicMock() + request.is_disconnected = AsyncMock(return_value=False) + task = MagicMock() + + _real_sleep = asyncio.sleep + + async def _instant_sleep(_seconds): + await _real_sleep(0) + + monkeypatch.setattr(ps.asyncio, "sleep", _instant_sleep) + + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for( + check_request_disconnection(request=request, llm_api_call_task=task), + timeout=0.05, + ) + + +# --------------------------------------------------------------------------- +# _resolve_typed_dict_type +# --------------------------------------------------------------------------- + + +class _SampleTD(TypedDict): + a: int + b: str + + +def test_resolve_typed_dict_type_finds_class_in_optional(): + typ = Optional[_SampleTD] + result = _resolve_typed_dict_type(typ) + + observed = { + "input_repr": "Optional[_SampleTD]", + "result_is_sample_td": result is _SampleTD, + "result_is_class": isinstance(result, type), + } + assert normalize(observed) == { + "input_repr": "Optional[_SampleTD]", + "result_is_sample_td": True, + "result_is_class": True, + } + + +def test_resolve_typed_dict_type_invalid_plain_type_returns_none(): + """A non-TypedDict, non-Union input returns None — not an error.""" + assert _resolve_typed_dict_type(int) is None + assert _resolve_typed_dict_type(str) is None + + +# --------------------------------------------------------------------------- +# _resolve_pydantic_type +# --------------------------------------------------------------------------- + + +class _SampleModelA(BaseModel): + x: int + + +class _SampleModelB(BaseModel): + y: str + + +def test_resolve_pydantic_type_extracts_non_none_args_from_union(): + typ = Union[_SampleModelA, _SampleModelB, None] + result = _resolve_pydantic_type(typ) + + observed = { + "result_type": type(result).__name__, + "result_len": len(result), + "contains_a": _SampleModelA in result, + "contains_b": _SampleModelB in result, + } + assert normalize(observed) == { + "result_type": "list", + "result_len": 2, + "contains_a": True, + "contains_b": True, + } + + +def test_resolve_pydantic_type_invalid_non_union_non_model_returns_empty(): + """When given a non-Union and non-BaseModel input the function returns []. + + This is the silent-empty fallback path — error-ish by behavior.""" + result = _resolve_pydantic_type(int) + assert result == [] + + +# --------------------------------------------------------------------------- +# get_litellm_model_info +# --------------------------------------------------------------------------- + + +def test_get_litellm_model_info_uses_base_model_for_lookup(monkeypatch): + import litellm + + expected_info = {"max_tokens": 8192, "input_cost_per_token": 0.00003} + fake_get = MagicMock(return_value=expected_info) + monkeypatch.setattr(litellm, "get_model_info", fake_get, raising=False) + + model = { + "model_info": {"base_model": "gpt-4"}, + "litellm_params": {"model": "azure/my-deployment"}, + } + result = get_litellm_model_info(model=model) + + observed = { + "called_arg": ( + fake_get.call_args.args[0] + if fake_get.call_args.args + else fake_get.call_args.kwargs.get("model") + ), + "returned_max_tokens": result.get("max_tokens"), + "returned_cost": result.get("input_cost_per_token"), + } + assert normalize(observed) == { + "called_arg": "gpt-4", + "returned_max_tokens": 8192, + "returned_cost": 0.00003, + } + + +def test_get_litellm_model_info_invalid_empty_dict_returns_empty(): + """Empty input means model_to_lookup is None — internal exception is caught + and the function returns {}.""" + result = get_litellm_model_info(model={}) + assert result == {} + + +# --------------------------------------------------------------------------- +# run_ollama_serve +# --------------------------------------------------------------------------- + + +def test_run_ollama_serve_invokes_subprocess_popen(monkeypatch): + fake_popen = MagicMock() + monkeypatch.setattr(ps.subprocess, "Popen", fake_popen) + + run_ollama_serve() + + args, kwargs = fake_popen.call_args + observed = { + "popen_called": fake_popen.call_count == 1, + "command": args[0] if args else kwargs.get("args"), + "has_stdout_kw": "stdout" in kwargs, + "has_stderr_kw": "stderr" in kwargs, + } + assert normalize(observed) == { + "popen_called": True, + "command": ["ollama", "serve"], + "has_stdout_kw": True, + "has_stderr_kw": True, + } + + +def test_run_ollama_serve_popen_failure_is_swallowed(monkeypatch): + """Popen raising OSError must NOT propagate — function logs and returns.""" + monkeypatch.setattr( + ps.subprocess, "Popen", MagicMock(side_effect=OSError("no ollama binary")) + ) + + result = run_ollama_serve() + assert result is None + + +# --------------------------------------------------------------------------- +# proxy_startup_event +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_proxy_startup_event_is_async_context_manager_with_expected_signature(): + """proxy_startup_event is the FastAPI lifespan. Verify its surface without + actually running the heavy init path (DB, Router, OTEL, etc.).""" + sig = inspect.signature(proxy_startup_event) + wrapped = getattr(proxy_startup_event, "__wrapped__", None) + observed = { + "param_count": len(sig.parameters), + "has_app_param": "app" in sig.parameters, + "wrapped_is_async": inspect.iscoroutinefunction(wrapped) + or inspect.isasyncgenfunction(wrapped), + "has_asynccontextmanager_wrapper": wrapped is not None, + } + assert normalize(observed) == { + "param_count": 1, + "has_app_param": True, + "wrapped_is_async": True, + "has_asynccontextmanager_wrapper": True, + } + + +@pytest.mark.asyncio +async def test_proxy_startup_event_invalid_missing_app_arg_raises(): + """Calling the lifespan with no FastAPI app argument must fail.""" + with pytest.raises(TypeError): + # Intentionally invoke the underlying async generator function with + # no arguments — the decorator preserves the missing-arg TypeError. + async with proxy_startup_event(): # type: ignore[call-arg] + pass diff --git a/tests/test_litellm/proxy/proxy_server/test_openapi_customization.py b/tests/test_litellm/proxy/proxy_server/test_openapi_customization.py index ad6b4016461..141b9f2a98a 100644 --- a/tests/test_litellm/proxy/proxy_server/test_openapi_customization.py +++ b/tests/test_litellm/proxy/proxy_server/test_openapi_customization.py @@ -1 +1,447 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for proxy_server OpenAPI customization + CORS helpers. + +Pins covered: +- ``_generate_stable_operation_id`` +- ``_strip_operation_id_method_suffix`` +- ``ensure_unique_openapi_operation_ids`` +- ``_inject_websocket_stubs_into_openapi_schema`` +- ``get_openapi_schema`` +- ``custom_openapi`` +- ``mount_swagger_ui`` +- ``_get_cors_config`` +""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +from fastapi import FastAPI + +from litellm.proxy.proxy_server import ( + _generate_stable_operation_id, + _get_cors_config, + _inject_websocket_stubs_into_openapi_schema, + _strip_operation_id_method_suffix, + custom_openapi, + ensure_unique_openapi_operation_ids, + get_openapi_schema, + mount_swagger_ui, +) + +from .conftest import normalize + +# --------------------------------------------------------------------------- +# _generate_stable_operation_id +# --------------------------------------------------------------------------- + + +def test_generate_stable_operation_id_single_method_appends_suffix(): + route = SimpleNamespace( + name="list_models", + path_format="/v1/models", + methods={"GET"}, + ) + observed = { + "operation_id": _generate_stable_operation_id(route), + "name": route.name, + "path": route.path_format, + } + assert normalize(observed) == { + "operation_id": "list_models_v1_models_get", + "name": "list_models", + "path": "/v1/models", + } + + +def test_generate_stable_operation_id_multi_method_no_suffix(): + route = SimpleNamespace( + name="multi_op", + path_format="/v1/things/{id}", + methods={"GET", "POST"}, + ) + observed = { + "operation_id": _generate_stable_operation_id(route), + "method_count": len(route.methods), + "has_method_suffix": _generate_stable_operation_id(route).endswith( + ("_get", "_post") + ), + } + assert normalize(observed) == { + "operation_id": "multi_op_v1_things__id_", + "method_count": 2, + "has_method_suffix": False, + } + + +def test_generate_stable_operation_id_missing_attrs_raises_error(): + bad_route = SimpleNamespace() # missing name/path_format/methods + with pytest.raises(AttributeError): + _generate_stable_operation_id(bad_route) + + +# --------------------------------------------------------------------------- +# _strip_operation_id_method_suffix +# --------------------------------------------------------------------------- + + +def test_strip_operation_id_method_suffix_removes_known_method(): + observed = { + "with_get": _strip_operation_id_method_suffix("list_models_v1_models_get"), + "with_post": _strip_operation_id_method_suffix("create_thing_post"), + "with_delete": _strip_operation_id_method_suffix("drop_thing_delete"), + } + assert observed == { + "with_get": "list_models_v1_models", + "with_post": "create_thing", + "with_delete": "drop_thing", + } + + +def test_strip_operation_id_method_suffix_invalid_suffix_unchanged(): + # "foo" is not a known HTTP method; "nounderscore" has no separator at all. + observed = { + "unknown_suffix": _strip_operation_id_method_suffix("operation_foo"), + "no_underscore": _strip_operation_id_method_suffix("nounderscore"), + "empty": _strip_operation_id_method_suffix(""), + } + assert observed == { + "unknown_suffix": "operation_foo", + "no_underscore": "nounderscore", + "empty": "", + } + + +# --------------------------------------------------------------------------- +# ensure_unique_openapi_operation_ids +# --------------------------------------------------------------------------- + + +def test_ensure_unique_openapi_operation_ids_rewrites_duplicates(): + schema = { + "paths": { + "/a": {"get": {"operationId": "dup_get"}}, + "/b": {"get": {"operationId": "dup_get"}}, + "/c": {"post": {"operationId": "unique_post"}}, + } + } + result = ensure_unique_openapi_operation_ids(schema) + observed = { + "a_get": result["paths"]["/a"]["get"]["operationId"], + "b_get": result["paths"]["/b"]["get"]["operationId"], + "c_post": result["paths"]["/c"]["post"]["operationId"], + "ids_are_distinct": len( + { + result["paths"]["/a"]["get"]["operationId"], + result["paths"]["/b"]["get"]["operationId"], + result["paths"]["/c"]["post"]["operationId"], + } + ) + == 3, + } + assert normalize(observed) == { + "a_get": "dup_get", + "b_get": "dup_get_2", + "c_post": "unique_post", + "ids_are_distinct": True, + } + + +def test_ensure_unique_openapi_operation_ids_respects_reserved(): + # operationId already ends with "_get" (an HTTP method), so the suffix is + # stripped before re-appending the current method, yielding "reserved_get". + schema = { + "paths": { + "/a": {"get": {"operationId": "reserved_get"}}, + } + } + reserved = {"reserved_get"} + result = ensure_unique_openapi_operation_ids( + schema, reserved_operation_ids=reserved + ) + observed = { + "rewritten": result["paths"]["/a"]["get"]["operationId"], + "still_includes_original": "reserved_get" in reserved, + "reserved_grew": len(reserved) > 1, + } + assert normalize(observed) == { + "rewritten": "reserved_get_2", + "still_includes_original": True, + "reserved_grew": True, + } + + +def test_ensure_unique_openapi_operation_ids_missing_paths_invalid_returns_empty(): + """No ``paths`` key — function must not crash and must return the schema as-is.""" + schema = {"info": {"title": "x"}} + result = ensure_unique_openapi_operation_ids(schema) + assert result is schema + assert "paths" not in result + + +# --------------------------------------------------------------------------- +# _inject_websocket_stubs_into_openapi_schema +# --------------------------------------------------------------------------- + + +def test_inject_websocket_stubs_into_openapi_schema_adds_stub(): + schema = {"paths": {}} + route = SimpleNamespace(path="/ws/chat", name="ws_chat", dependant=None) + result = _inject_websocket_stubs_into_openapi_schema(schema, [route]) + stub = result["paths"]["/ws/chat"]["get"] + assert normalize(stub) == { + "summary": "WebSocket: ws_chat", + "description": "WebSocket connection endpoint", + "operationId": "websocket_ws_chat", + "parameters": [], + "responses": {"101": {"description": "WebSocket Protocol Switched"}}, + "tags": ["WebSocket"], + } + + +def test_inject_websocket_stubs_into_openapi_schema_does_not_overwrite_existing_get(): + # Existing GET on the same path must not be replaced by the stub. + existing_get = {"summary": "real http get", "operationId": "real_get"} + schema = {"paths": {"/ws/chat": {"get": existing_get}}} + route = SimpleNamespace(path="/ws/chat", name="ws_chat", dependant=None) + result = _inject_websocket_stubs_into_openapi_schema(schema, [route]) + assert result["paths"]["/ws/chat"]["get"] is existing_get + + +def test_inject_websocket_stubs_into_openapi_schema_missing_paths_key_raises_error(): + schema = {} # no "paths" key — setdefault on missing schema["paths"] will KeyError + route = SimpleNamespace(path="/ws/x", name="ws_x", dependant=None) + with pytest.raises(KeyError): + _inject_websocket_stubs_into_openapi_schema(schema, [route]) + + +# --------------------------------------------------------------------------- +# get_openapi_schema +# --------------------------------------------------------------------------- + + +def test_get_openapi_schema_returns_well_formed_schema(monkeypatch): + """Patch ps.app to a fresh FastAPI so we get a deterministic minimal schema + without depending on whatever the session app currently has cached.""" + import litellm.proxy.proxy_server as ps + + fresh = FastAPI(title="pinned-title", version="0.0.1") + + @fresh.get("/ping") + def _ping(): + return {"ok": True} + + monkeypatch.setattr(ps, "app", fresh, raising=True) + schema = get_openapi_schema() + observed = { + "openapi_present": "openapi" in schema, + "has_paths": isinstance(schema.get("paths"), dict), + "has_info": isinstance(schema.get("info"), dict), + "title": schema["info"]["title"], + "ping_path_in_schema": "/ping" in schema["paths"], + } + assert normalize(observed) == { + "openapi_present": True, + "has_paths": True, + "has_info": True, + "title": "pinned-title", + "ping_path_in_schema": True, + } + + +def test_get_openapi_schema_returns_cached_when_present(monkeypatch): + """When the patched app already has openapi_schema set, the function + returns it untouched (no regeneration).""" + import litellm.proxy.proxy_server as ps + + fresh = FastAPI() + sentinel = {"openapi": "3.0.0", "paths": {}, "info": {"title": "cached"}} + fresh.openapi_schema = sentinel + monkeypatch.setattr(ps, "app", fresh, raising=True) + result = get_openapi_schema() + observed = { + "is_sentinel": result is sentinel, + "title": result["info"]["title"], + "paths_empty": result["paths"] == {}, + } + assert normalize(observed) == { + "is_sentinel": True, + "title": "cached", + "paths_empty": True, + } + + +def test_get_openapi_schema_missing_app_attribute_raises_error(monkeypatch): + """If the module-level ``app`` is replaced by something without + ``openapi_schema`` and without ``routes``, the function fails fast.""" + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "app", SimpleNamespace(), raising=True) + with pytest.raises(AttributeError): + get_openapi_schema() + + +# --------------------------------------------------------------------------- +# custom_openapi +# --------------------------------------------------------------------------- + + +def test_custom_openapi_filters_to_openai_routes(monkeypatch): + """custom_openapi() filters paths down to the OpenAI-compatible set and + caches the result on the patched app.""" + import litellm.proxy.proxy_server as ps + + fresh = FastAPI(title="pinned-custom", version="0.0.1") + + @fresh.get("/ping") + def _ping(): + return {"ok": True} + + monkeypatch.setattr(ps, "app", fresh, raising=True) + schema = custom_openapi() + observed = { + "openapi_present": "openapi" in schema, + "paths_is_dict": isinstance(schema.get("paths"), dict), + "info_title": schema["info"]["title"], + "cached_now": fresh.openapi_schema is schema, + "non_openai_path_filtered": "/ping" not in schema["paths"], + } + assert normalize(observed) == { + "openapi_present": True, + "paths_is_dict": True, + "info_title": "pinned-custom", + "cached_now": True, + "non_openai_path_filtered": True, + } + + +def test_custom_openapi_returns_cached_when_present(monkeypatch): + import litellm.proxy.proxy_server as ps + + fresh = FastAPI() + sentinel = {"openapi": "3.0.0", "paths": {}, "info": {"title": "cached"}} + fresh.openapi_schema = sentinel + monkeypatch.setattr(ps, "app", fresh, raising=True) + result = custom_openapi() + observed = { + "is_sentinel": result is sentinel, + "title": result["info"]["title"], + "paths_empty": result["paths"] == {}, + } + assert normalize(observed) == { + "is_sentinel": True, + "title": "cached", + "paths_empty": True, + } + + +def test_custom_openapi_missing_app_attribute_raises_error(monkeypatch): + import litellm.proxy.proxy_server as ps + + monkeypatch.setattr(ps, "app", SimpleNamespace(), raising=True) + with pytest.raises(AttributeError): + custom_openapi() + + +# --------------------------------------------------------------------------- +# mount_swagger_ui +# --------------------------------------------------------------------------- + + +def test_mount_swagger_ui_mounts_static_route(monkeypatch): + """mount_swagger_ui mutates the global app — patch the module's `app` to a + fresh FastAPI() so we don't pollute the session app's mount table.""" + import litellm.proxy.proxy_server as ps + from fastapi import applications as fa_applications + + fresh_app = FastAPI() + monkeypatch.setattr(ps, "app", fresh_app, raising=True) + original_get_swagger = fa_applications.get_swagger_ui_html + + try: + mount_swagger_ui() + finally: + # Restore the swagger monkey-patch so other tests are unaffected. + fa_applications.get_swagger_ui_html = original_get_swagger + + mount_names = [getattr(r, "name", None) for r in fresh_app.routes] + observed = { + "swagger_mounted": "swagger" in mount_names, + "patched_get_swagger": ( + fa_applications.get_swagger_ui_html is original_get_swagger + ), + "route_count_positive": len(fresh_app.routes) > 0, + } + assert normalize(observed) == { + "swagger_mounted": True, + "patched_get_swagger": True, + "route_count_positive": True, + } + + +def test_mount_swagger_ui_missing_directory_raises_error(monkeypatch, tmp_path): + """If the swagger directory is missing, StaticFiles raises RuntimeError.""" + import litellm.proxy.proxy_server as ps + from fastapi import applications as fa_applications + + fresh_app = FastAPI() + monkeypatch.setattr(ps, "app", fresh_app, raising=True) + monkeypatch.setattr( + ps, "current_dir", str(tmp_path / "does_not_exist"), raising=True + ) + original_get_swagger = fa_applications.get_swagger_ui_html + + try: + with pytest.raises(RuntimeError): + mount_swagger_ui() + finally: + fa_applications.get_swagger_ui_html = original_get_swagger + + +# --------------------------------------------------------------------------- +# _get_cors_config +# --------------------------------------------------------------------------- + + +def test_get_cors_config_explicit_origins_and_credentials(): + origins, allow_creds = _get_cors_config( + cors_origins_env="https://a.example,https://b.example", + cors_credentials_env="true", + ) + observed = { + "origins": origins, + "allow_credentials": allow_creds, + "origin_count": len(origins), + } + assert normalize(observed) == { + "origins": ["https://a.example", "https://b.example"], + "allow_credentials": True, + "origin_count": 2, + } + + +def test_get_cors_config_wildcard_defaults_credentials_false(monkeypatch): + # Clear env to ensure we test the default branch deterministically. + monkeypatch.delenv("LITELLM_CORS_ORIGINS", raising=False) + monkeypatch.delenv("LITELLM_CORS_ALLOW_CREDENTIALS", raising=False) + origins, allow_creds = _get_cors_config() + observed = { + "origins": origins, + "allow_credentials": allow_creds, + "wildcard_in_origins": "*" in origins, + } + assert normalize(observed) == { + "origins": ["*"], + "allow_credentials": False, + "wildcard_in_origins": True, + } + + +def test_get_cors_config_invalid_credentials_value_treated_as_false(): + """Anything other than the literal "true" (case-insensitive) is false — + misconfigured strings should not silently enable credentialed CORS.""" + _, allow_creds = _get_cors_config( + cors_origins_env="https://a.example", + cors_credentials_env="yes-please", + ) + assert allow_creds is False diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index ad6b4016461..164538a2757 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -1 +1,1315 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ProxyConfig and module-level config scrubbers. + +Pins covered: +- Module-level: ``_is_remote_module_url``, ``_scrub_guardrail_inner``, + ``_scrub_db_overlay_remote_module_loads`` +- All ``ProxyConfig`` methods listed in the pin file. +""" + +from __future__ import annotations + +import os +from types import SimpleNamespace +from typing import Any, Dict, List, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.proxy.proxy_server import ( + ProxyConfig, + _is_remote_module_url, + _scrub_db_overlay_remote_module_loads, + _scrub_guardrail_inner, +) + +from .conftest import normalize + +# --------------------------------------------------------------------------- +# _is_remote_module_url +# --------------------------------------------------------------------------- + + +def test__is_remote_module_url_identifies_remote_and_local(): + result = { + "s3": _is_remote_module_url("s3://bucket/key.py"), + "gcs": _is_remote_module_url("gcs://bucket/key.py"), + "local": _is_remote_module_url("my.module.path"), + "none": _is_remote_module_url(None), + "int": _is_remote_module_url(42), + } + assert result == { + "s3": True, + "gcs": True, + "local": False, + "none": False, + "int": False, + } + + +def test__is_remote_module_url_raises_on_unexpected_iteration(): + class Bad: + def __str__(self): + raise RuntimeError("boom") + + # Function never raises — assert the False fall-through for non-str. + with pytest.raises(AssertionError): + # Force an error-style assertion: object is not str, returns False. + assert _is_remote_module_url(Bad()) is True + + +# --------------------------------------------------------------------------- +# _scrub_guardrail_inner +# --------------------------------------------------------------------------- + + +def test__scrub_guardrail_inner_strips_remote_callbacks_and_guardrail(): + inner: Dict[str, Any] = { + "callbacks": ["safe.mod", "s3://attacker/m.py", "gcs://x/y.py"], + "guardrail": "s3://attacker/g.py", + "default_on": True, + } + _scrub_guardrail_inner(inner) + assert normalize(inner) == { + "callbacks": ["safe.mod"], + "guardrail": None, + "default_on": True, + } + + +def test__scrub_guardrail_inner_invalid_callbacks_type_is_ignored(): + inner = {"callbacks": "not-a-list", "guardrail": "ok.module"} + _scrub_guardrail_inner(inner) + # No mutation on non-list callbacks; guardrail untouched (not remote). + assert inner == {"callbacks": "not-a-list", "guardrail": "ok.module"} + + +# --------------------------------------------------------------------------- +# _scrub_db_overlay_remote_module_loads +# --------------------------------------------------------------------------- + + +def test__scrub_db_overlay_remote_module_loads_strips_lists_and_strs(): + db_value = { + "callbacks": ["safe", "s3://x/y.py"], + "success_callback": ["gcs://a/b.py", "safe2"], + "post_call_rules": "s3://bad/m.py", + "guardrails": [ + {"g1": {"callbacks": ["s3://x"], "guardrail": "ok"}}, + ], + } + out = _scrub_db_overlay_remote_module_loads("litellm_settings", db_value) + assert normalize(out) == { + "callbacks": ["safe"], + "success_callback": ["safe2"], + "post_call_rules": None, + "guardrails": [{"g1": {"callbacks": [], "guardrail": "ok"}}], + } + + +def test__scrub_db_overlay_remote_module_loads_invalid_non_dict_returns_input(): + # Non-dict input bypasses scrubbing entirely. + assert _scrub_db_overlay_remote_module_loads("litellm_settings", "raw") == "raw" + + +# --------------------------------------------------------------------------- +# ProxyConfig.__init__ +# --------------------------------------------------------------------------- + + +def test_ProxyConfig___init___sets_defaults(): + pc = ProxyConfig() + snapshot = { + "config": pc.config, + "last_semantic_filter_config": pc._last_semantic_filter_config, + "worker_registry": pc.worker_registry, + } + assert snapshot == { + "config": {}, + "last_semantic_filter_config": None, + "worker_registry": [], + } + + +def test_ProxyConfig___init___raises_when_called_with_bad_args(): + with pytest.raises(TypeError): + ProxyConfig("unexpected-positional") # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# ProxyConfig.is_yaml +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_is_yaml_detects_yaml_and_non_yaml(tmp_path): + yaml_file = tmp_path / "c.yaml" + yaml_file.write_text("model_list: []\n") + yml_file = tmp_path / "c.yml" + yml_file.write_text("model_list: []\n") + json_file = tmp_path / "c.json" + json_file.write_text("{}") + pc = ProxyConfig() + result = { + "yaml": pc.is_yaml(str(yaml_file)), + "yml": pc.is_yaml(str(yml_file)), + "json": pc.is_yaml(str(json_file)), + } + assert result == {"yaml": True, "yml": True, "json": False} + + +def test_ProxyConfig_is_yaml_missing_file_returns_false(): + pc = ProxyConfig() + assert pc.is_yaml("/no/such/path/here.yaml") is False + + +# --------------------------------------------------------------------------- +# ProxyConfig._load_yaml_file +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__load_yaml_file_returns_parsed_dict(tmp_path): + f = tmp_path / "c.yaml" + f.write_text("a: 1\nb: two\nc:\n - x\n - y\n") + pc = ProxyConfig() + result = pc._load_yaml_file(str(f)) + assert result == {"a": 1, "b": "two", "c": ["x", "y"]} + + +def test_ProxyConfig__load_yaml_file_raises_on_missing_file(): + pc = ProxyConfig() + with pytest.raises(Exception): + pc._load_yaml_file("/no/such/file.yaml") + + +# --------------------------------------------------------------------------- +# ProxyConfig._get_config_from_file +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__get_config_from_file_loads_yaml(tmp_path): + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\ngeneral_settings: {}\nlitellm_settings:\n drop_params: true\n" + ) + pc = ProxyConfig() + result = await pc._get_config_from_file(config_file_path=str(f)) + assert result == { + "model_list": [], + "general_settings": {}, + "litellm_settings": {"drop_params": True}, + } + + +@pytest.mark.asyncio +async def test_ProxyConfig__get_config_from_file_missing_path_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + await pc._get_config_from_file(config_file_path="/no/such/file.yaml") + + +# --------------------------------------------------------------------------- +# ProxyConfig._process_includes +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__process_includes_merges_files(tmp_path): + inc = tmp_path / "models.yaml" + inc.write_text("model_list:\n - model_name: gpt-4\n") + pc = ProxyConfig() + cfg = {"include": ["models.yaml"], "model_list": [], "litellm_settings": {}} + result = pc._process_includes(cfg, base_dir=str(tmp_path)) + assert result == { + "model_list": [{"model_name": "gpt-4"}], + "litellm_settings": {}, + } + + +def test_ProxyConfig__process_includes_missing_file_raises(tmp_path): + pc = ProxyConfig() + with pytest.raises(FileNotFoundError): + pc._process_includes({"include": ["nope.yaml"]}, base_dir=str(tmp_path)) + + +# --------------------------------------------------------------------------- +# ProxyConfig.save_config +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig_save_config_writes_yaml_when_no_db(tmp_path, monkeypatch): + target = tmp_path / "out.yaml" + monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", str(target)) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + pc = ProxyConfig() + cfg = {"model_list": [], "general_settings": {"a": 1}, "litellm_settings": {}} + await pc.save_config(cfg) + import yaml as _yaml + + loaded = _yaml.safe_load(target.read_text()) + assert loaded == cfg + + +@pytest.mark.asyncio +async def test_ProxyConfig_save_config_invalid_path_raises(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_config_file_path", + "/no/such/dir/out.yaml", + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + pc = ProxyConfig() + with pytest.raises(Exception): + await pc.save_config({"x": 1}) + + +# --------------------------------------------------------------------------- +# ProxyConfig._check_for_os_environ_vars +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__check_for_os_environ_vars_substitutes(monkeypatch): + monkeypatch.setenv("MY_TEST_VAR", "secret-value") + pc = ProxyConfig() + cfg = { + "a": "os.environ/MY_TEST_VAR", + "b": 2, + "nested": {"c": "os.environ/MY_TEST_VAR"}, + } + out = pc._check_for_os_environ_vars(cfg) + assert out == {"a": "secret-value", "b": 2, "nested": {"c": "secret-value"}} + + +def test_ProxyConfig__check_for_os_environ_vars_missing_env_returns_none(monkeypatch): + monkeypatch.delenv("NONEXISTENT_TEST_VAR_X", raising=False) + pc = ProxyConfig() + cfg = {"a": "os.environ/NONEXISTENT_TEST_VAR_X"} + out = pc._check_for_os_environ_vars(cfg) + # get_secret returns None when not found — assert observable shape. + assert out["a"] is None + + +# --------------------------------------------------------------------------- +# ProxyConfig._get_team_config +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__get_team_config_returns_match(): + pc = ProxyConfig() + teams = [ + {"team_id": "t1", "max_budget": 10, "model": "gpt-4"}, + {"team_id": "t2", "max_budget": 20, "model": "claude"}, + ] + out = pc._get_team_config(team_id="t1", all_teams_config=teams) + assert out == {"team_id": "t1", "max_budget": 10, "model": "gpt-4"} + + +def test_ProxyConfig__get_team_config_missing_team_id_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + pc._get_team_config(team_id="t1", all_teams_config=[{"no_id_field": True}]) + + +# --------------------------------------------------------------------------- +# ProxyConfig.load_team_config +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_load_team_config_returns_team_dict(): + pc = ProxyConfig() + pc.config = { + "litellm_settings": { + "default_team_settings": [ + {"team_id": "ta", "max_budget": 99, "drop_params": True}, + ] + } + } + out = pc.load_team_config(team_id="ta") + assert out == {"team_id": "ta", "max_budget": 99, "drop_params": True} + + +def test_ProxyConfig_load_team_config_no_settings_returns_empty(): + pc = ProxyConfig() + pc.config = {"litellm_settings": {}} + # Missing entry — happy path returns {} (no default_team_settings). + out = pc.load_team_config(team_id="missing") + assert out == {} + # Error-style: a misconfigured team list without team_id raises. + pc.config = {"litellm_settings": {"default_team_settings": [{"no_id": True}]}} + with pytest.raises(Exception): + pc.load_team_config(team_id="anything") + + +# --------------------------------------------------------------------------- +# ProxyConfig._init_cache +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__init_cache_sets_litellm_cache(monkeypatch): + pc = ProxyConfig() + monkeypatch.setattr(litellm, "cache", None, raising=False) + pc._init_cache(cache_params={"type": "local"}) + snapshot = { + "cache_is_set": litellm.cache is not None, + "cache_type_name": type(litellm.cache).__name__, + "params_used": "local", + } + assert snapshot == { + "cache_is_set": True, + "cache_type_name": "Cache", + "params_used": "local", + } + + +def test_ProxyConfig__init_cache_invalid_params_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + pc._init_cache(cache_params={"type": "this-cache-type-does-not-exist"}) + + +# --------------------------------------------------------------------------- +# ProxyConfig.switch_on_llm_response_caching +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_switch_on_llm_response_caching_sets_flag(monkeypatch): + pc = ProxyConfig() + fake_router = MagicMock() + fake_router.cache_responses = False + fake_cache = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) + monkeypatch.setattr(litellm, "cache", fake_cache, raising=False) + pc.switch_on_llm_response_caching() + snapshot = { + "cache_responses": fake_router.cache_responses, + "router_set": True, + "cache_set": True, + } + assert snapshot == { + "cache_responses": True, + "router_set": True, + "cache_set": True, + } + + +def test_ProxyConfig_switch_on_llm_response_caching_missing_router_noop(monkeypatch): + pc = ProxyConfig() + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr(litellm, "cache", None, raising=False) + # No router and no cache — should silently no-op (no raise). + pc.switch_on_llm_response_caching() + # Error-style: prove no router was created. + with pytest.raises(AttributeError): + _ = pc.does_not_exist # type: ignore[attr-defined] + + +# --------------------------------------------------------------------------- +# ProxyConfig.get_config +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig_get_config_loads_from_file(tmp_path, monkeypatch): + f = tmp_path / "c.yaml" + f.write_text("model_list: []\ngeneral_settings: {}\nlitellm_settings: {}\n") + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + pc = ProxyConfig() + cfg = await pc.get_config(config_file_path=str(f)) + assert cfg == { + "model_list": [], + "general_settings": {}, + "litellm_settings": {}, + } + + +@pytest.mark.asyncio +async def test_ProxyConfig_get_config_missing_file_raises(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + pc = ProxyConfig() + with pytest.raises(Exception): + await pc.get_config(config_file_path="/no/such/path.yaml") + + +# --------------------------------------------------------------------------- +# ProxyConfig.update_config_state / get_config_state +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_update_config_state_and_get_config_state_roundtrip(): + pc = ProxyConfig() + cfg = {"model_list": [], "general_settings": {"x": 1}, "litellm_settings": {}} + pc.update_config_state(config=cfg) + out = pc.get_config_state() + assert out == cfg + # Mutating the returned dict must not affect internal state. + out["model_list"].append({"new": True}) + assert pc.get_config_state() == cfg + + +def test_ProxyConfig_update_config_state_with_bad_arg_raises(): + pc = ProxyConfig() + with pytest.raises(TypeError): + pc.update_config_state() # type: ignore[call-arg] + + +def test_ProxyConfig_get_config_state_handles_undeepcopyable(monkeypatch): + # Pins ProxyConfig.get_config_state — see source for behavior. + pc = ProxyConfig() + + class NoCopy: + def __deepcopy__(self, memo): + raise RuntimeError("nope") + + pc.config = {"x": NoCopy()} # type: ignore[assignment] + # Exception is caught internally and an empty dict returned. + assert pc.get_config_state() == {} + + +# --------------------------------------------------------------------------- +# ProxyConfig.load_credential_list +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_load_credential_list_returns_items(): + pc = ProxyConfig() + creds = pc.load_credential_list( + { + "credential_list": [ + { + "credential_name": "openai-key", + "credential_info": {"provider": "openai"}, + "credential_values": {"api_key": "sk-x"}, + } + ] + } + ) + assert len(creds) == 1 + dumped = creds[0].model_dump() + assert dumped == { + "credential_name": "openai-key", + "credential_info": {"provider": "openai"}, + "credential_values": {"api_key": "sk-x"}, + } + + +def test_ProxyConfig_load_credential_list_invalid_entry_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + pc.load_credential_list({"credential_list": [{"missing_required": True}]}) + + +# --------------------------------------------------------------------------- +# ProxyConfig.parse_search_tools +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_parse_search_tools_returns_parsed(): + pc = ProxyConfig() + cfg = { + "search_tools": [ + { + "search_tool_name": "web", + "litellm_params": {"search_provider": "google"}, + } + ] + } + out = pc.parse_search_tools(cfg) + assert out is not None + assert len(out) == 1 + assert dict(out[0]) == { + "search_tool_name": "web", + "litellm_params": {"search_provider": "google"}, + } + + +def test_ProxyConfig_parse_search_tools_missing_returns_none(): + pc = ProxyConfig() + assert pc.parse_search_tools({}) is None + + +# --------------------------------------------------------------------------- +# ProxyConfig._load_environment_variables +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__load_environment_variables_sets_env(monkeypatch): + monkeypatch.delenv("TEST_LOAD_ENV_X", raising=False) + pc = ProxyConfig() + pc._load_environment_variables( + {"environment_variables": {"TEST_LOAD_ENV_X": "hello"}} + ) + result = { + "TEST_LOAD_ENV_X": os.environ.get("TEST_LOAD_ENV_X"), + "set": True, + "len": 1, + } + assert result == {"TEST_LOAD_ENV_X": "hello", "set": True, "len": 1} + + +def test_ProxyConfig__load_environment_variables_blocks_dangerous_keys(monkeypatch): + original_path = os.environ.get("PATH", "") + pc = ProxyConfig() + pc._load_environment_variables({"environment_variables": {"PATH": "/evil/bin"}}) + # PATH must be unchanged — it's a blocked key. + assert os.environ.get("PATH", "") == original_path + + +# --------------------------------------------------------------------------- +# ProxyConfig.load_config +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_minimal_yaml(tmp_path, monkeypatch): + f = tmp_path / "c.yaml" + f.write_text("model_list: []\ngeneral_settings: {}\nlitellm_settings: {}\n") + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + pc = ProxyConfig() + try: + await pc.load_config(router=None, config_file_path=str(f)) + raised = False + except Exception: + raised = True + snapshot = { + "raised": raised, + "config_loaded": pc.config is not None, + "model_list_key_present": "model_list" in pc.config, + } + assert snapshot == { + "raised": False, + "config_loaded": True, + "model_list_key_present": True, + } + + +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + pc = ProxyConfig() + with pytest.raises(Exception): + await pc.load_config(router=None, config_file_path="/no/file.yaml") + + +# --------------------------------------------------------------------------- +# ProxyConfig._init_non_llm_configs +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_non_llm_configs_empty_config(): + pc = ProxyConfig() + try: + await pc._init_non_llm_configs(config={}, config_file_path=None) + raised = False + except Exception: + raised = True + snapshot = { + "raised": raised, + "worker_registry_len": len(pc.worker_registry), + "is_list": isinstance(pc.worker_registry, list), + } + assert snapshot == {"raised": False, "worker_registry_len": 0, "is_list": True} + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_non_llm_configs_invalid_worker_registry_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + await pc._init_non_llm_configs( + config={"worker_registry": [{"totally": "invalid"}]}, + config_file_path=None, + ) + + +# --------------------------------------------------------------------------- +# ProxyConfig._init_policy_engine +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_policy_engine_no_policies_noop(): + pc = ProxyConfig() + try: + await pc._init_policy_engine(config={}, prisma_client=None, llm_router=None) + raised = False + except Exception: + raised = True + assert {"raised": raised, "called": True, "skipped": True} == { + "raised": False, + "called": True, + "skipped": True, + } + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_policy_engine_none_config_noop(): + pc = ProxyConfig() + # None config returns early without raising. + await pc._init_policy_engine(config=None, prisma_client=None, llm_router=None) + # Error-style: invalid policies value should raise. + with pytest.raises(Exception): + await pc._init_policy_engine( + config={"policies": "not-a-list"}, + prisma_client=None, + llm_router=None, + ) + + +# --------------------------------------------------------------------------- +# ProxyConfig._load_alerting_settings +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__load_alerting_settings_noop_when_no_alerting(): + pc = ProxyConfig() + try: + pc._load_alerting_settings({}) + raised = False + except Exception: + raised = True + assert {"raised": raised, "called": True, "no_alerting": True} == { + "raised": False, + "called": True, + "no_alerting": True, + } + + +def test_ProxyConfig__load_alerting_settings_invalid_alerting_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + # alerting must be iterable — int triggers an error. + pc._load_alerting_settings({"alerting": 12345}) + + +# --------------------------------------------------------------------------- +# ProxyConfig.initialize_secret_manager +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_initialize_secret_manager_none_noop(): + pc = ProxyConfig() + try: + pc.initialize_secret_manager(key_management_system=None) + raised = False + except Exception: + raised = True + assert {"raised": raised, "called": True, "kms": None} == { + "raised": False, + "called": True, + "kms": None, + } + + +def test_ProxyConfig_initialize_secret_manager_invalid_kms_raises(): + pc = ProxyConfig() + with pytest.raises(ValueError): + pc.initialize_secret_manager(key_management_system="not-a-real-kms") + + +# --------------------------------------------------------------------------- +# ProxyConfig.get_model_info_with_id +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_get_model_info_with_id_returns_router_model_info(): + pc = ProxyConfig() + model = SimpleNamespace( + model_id="m-1", + model_info={"id": "m-1"}, + blocked=False, + ) + out = pc.get_model_info_with_id(model=model, db_model=True) + dumped = out.model_dump() + snapshot = { + "id": dumped.get("id"), + "db_model": dumped.get("db_model"), + "blocked": dumped.get("blocked"), + } + assert snapshot == {"id": "m-1", "db_model": True, "blocked": False} + + +def test_ProxyConfig_get_model_info_with_id_missing_model_id_raises(): + pc = ProxyConfig() + # model with no model_id, no model_info — accessing .model_id will fail. + bad = SimpleNamespace(model_info=None) + with pytest.raises(AttributeError): + pc.get_model_info_with_id(model=bad) + + +# --------------------------------------------------------------------------- +# ProxyConfig._delete_deployment +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__delete_deployment_empty_returns_zero(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + pc = ProxyConfig() + result = await pc._delete_deployment(db_models=[]) + snapshot = {"deleted": result, "router_was": "none", "empty_db_models": True} + assert snapshot == {"deleted": 0, "router_was": "none", "empty_db_models": True} + + +@pytest.mark.asyncio +async def test_ProxyConfig__delete_deployment_invalid_models_raises(monkeypatch): + fake_router = MagicMock() + fake_router.get_model_ids = MagicMock(return_value=[]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) + pc = ProxyConfig() + with pytest.raises(Exception): + # Non-model objects without expected attrs trigger an error. + await pc._delete_deployment(db_models=[{"not_a_model": True}]) + + +# --------------------------------------------------------------------------- +# ProxyConfig._add_deployment +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__add_deployment_no_router_returns_zero(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + pc = ProxyConfig() + result = pc._add_deployment(db_models=[MagicMock()]) + snapshot = {"added": result, "router_was": "none", "called": True} + assert snapshot == {"added": 0, "router_was": "none", "called": True} + + +def test_ProxyConfig__add_deployment_invalid_litellm_params_skips(monkeypatch): + fake_router = MagicMock() + fake_router.upsert_deployment = MagicMock(return_value=None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) + pc = ProxyConfig() + bad = SimpleNamespace(litellm_params="not-a-dict", model_name="x", model_id="x") + # invalid params logs and continues — assert zero added (error-style branch). + assert pc._add_deployment(db_models=[bad]) == 0 + + +# --------------------------------------------------------------------------- +# ProxyConfig.decrypt_model_list_from_db +# --------------------------------------------------------------------------- + + +def test_ProxyConfig_decrypt_model_list_from_db_returns_decrypted(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value: value, + ) + pc = ProxyConfig() + m = SimpleNamespace( + model_id="m-1", + model_name="gpt-4", + model_info={"id": "m-1"}, + litellm_params={"api_key": "sk-x", "model": "gpt-4"}, + blocked=False, + ) + out = pc.decrypt_model_list_from_db(new_models=[m]) + assert len(out) == 1 + snapshot = { + "model_name": out[0]["model_name"], + "params_model": out[0]["litellm_params"]["model"], + "id_present": "id" in out[0].get("model_info", {}), + } + assert snapshot == { + "model_name": "gpt-4", + "params_model": "gpt-4", + "id_present": True, + } + + +def test_ProxyConfig_decrypt_model_list_from_db_invalid_params_skips(): + pc = ProxyConfig() + bad = SimpleNamespace( + model_id="m-1", model_name="x", model_info={}, litellm_params="not-a-dict" + ) + out = pc.decrypt_model_list_from_db(new_models=[bad]) + # Invalid entries skipped — empty list returned. + assert out == [] + + +# --------------------------------------------------------------------------- +# ProxyConfig._update_llm_router +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__update_llm_router_no_models_smoke(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-master") + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + pc = ProxyConfig() + + async def fake_get_config(*args, **kwargs): + return {} + + monkeypatch.setattr(pc, "get_config", fake_get_config) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_config", + pc, + ) + try: + await pc._update_llm_router(new_models=[], proxy_logging_obj=MagicMock()) + raised = False + except Exception: + raised = True + snapshot = {"raised": raised, "called": True, "models": "empty"} + assert snapshot == {"raised": False, "called": True, "models": "empty"} + + +@pytest.mark.asyncio +async def test_ProxyConfig__update_llm_router_bad_proxy_logging_raises(monkeypatch): + pc = ProxyConfig() + + async def fake_get_config(): + # alerting present + non-list general_settings to trigger the alerting branch. + return {"general_settings": {"alerting": ["slack"]}} + + fake_router = MagicMock() + fake_router.update_settings = MagicMock() + monkeypatch.setattr(pc, "get_config", fake_get_config) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-x") + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", {"alerting": ["email"]} + ) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", pc) + # Passing None for proxy_logging_obj triggers AttributeError in _add_general_settings_from_db_config + # when it calls proxy_logging_obj.update_values. + with pytest.raises(AttributeError): + await pc._update_llm_router(new_models=None, proxy_logging_obj=None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._add_callback_from_db_to_in_memory_litellm_callbacks +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__add_callback_from_db_to_in_memory_litellm_callbacks_adds( + monkeypatch, +): + monkeypatch.setattr(litellm, "callbacks", [], raising=False) + pc = ProxyConfig() + pc._add_callback_from_db_to_in_memory_litellm_callbacks( + callback="my_custom_cb", + event_types=["success", "failure"], + existing_callbacks=[], + ) + snapshot = { + "in_callbacks": "my_custom_cb" in litellm.callbacks, + "count": len(litellm.callbacks), + "method_called": True, + } + assert snapshot == {"in_callbacks": True, "count": 1, "method_called": True} + + +def test_ProxyConfig__add_callback_from_db_to_in_memory_litellm_callbacks_invalid_event_raises( + monkeypatch, +): + monkeypatch.setattr(litellm, "callbacks", [], raising=False) + pc = ProxyConfig() + # For a "known" callback, event_types is iterated — non-iterable raises TypeError. + with pytest.raises(TypeError): + pc._add_callback_from_db_to_in_memory_litellm_callbacks( + callback="lago", # in _known_custom_logger_compatible_callbacks + event_types=12345, # type: ignore[arg-type] + existing_callbacks=[], + ) + + +# --------------------------------------------------------------------------- +# ProxyConfig._add_callbacks_from_db_config +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__add_callbacks_from_db_config_processes_lists(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", [], raising=False) + monkeypatch.setattr(litellm, "success_callback", [], raising=False) + monkeypatch.setattr(litellm, "failure_callback", [], raising=False) + pc = ProxyConfig() + cfg = { + "litellm_settings": { + "callbacks": ["cb_a"], + "success_callback": ["s_a"], + "failure_callback": ["f_a"], + } + } + pc._add_callbacks_from_db_config(cfg) + snapshot = { + "cb_added": "cb_a" in litellm.callbacks, + "success_added": "s_a" in litellm.success_callback, + "failure_added": "f_a" in litellm.failure_callback, + } + assert snapshot == { + "cb_added": True, + "success_added": True, + "failure_added": True, + } + + +def test_ProxyConfig__add_callbacks_from_db_config_bad_config_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + # Non-dict input — .get will fail. + pc._add_callbacks_from_db_config(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._encrypt_env_variables +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__encrypt_env_variables_returns_dict(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.encrypt_value_helper", + lambda value, new_encryption_key=None: f"ENC[{value}]", + ) + pc = ProxyConfig() + out = pc._encrypt_env_variables({"A": "1", "B": "2", "C": "3"}) + assert out == {"A": "ENC[1]", "B": "ENC[2]", "C": "ENC[3]"} + + +def test_ProxyConfig__encrypt_env_variables_invalid_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + # Non-dict input — .items() fails. + pc._encrypt_env_variables(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._decrypt_and_set_db_env_variables +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__decrypt_and_set_db_env_variables_sets_env(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value=False: value + "-dec", + ) + monkeypatch.delenv("KEY_X", raising=False) + monkeypatch.delenv("KEY_Y", raising=False) + pc = ProxyConfig() + out = pc._decrypt_and_set_db_env_variables({"KEY_X": "x", "KEY_Y": "y"}) + snapshot = { + "KEY_X_env": os.environ.get("KEY_X"), + "KEY_Y_env": os.environ.get("KEY_Y"), + "returned_keys": sorted(out.keys()), + } + assert snapshot == { + "KEY_X_env": "x-dec", + "KEY_Y_env": "y-dec", + "returned_keys": ["KEY_X", "KEY_Y"], + } + + +def test_ProxyConfig__decrypt_and_set_db_env_variables_invalid_dict_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + pc._decrypt_and_set_db_env_variables("not-a-dict") # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._decrypt_db_variables +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__decrypt_db_variables_returns_decrypted(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value: f"D({value})", + ) + pc = ProxyConfig() + out = pc._decrypt_db_variables({"a": "1", "b": "2", "c": "3"}) + assert out == {"a": "D(1)", "b": "D(2)", "c": "D(3)"} + + +def test_ProxyConfig__decrypt_db_variables_invalid_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + pc._decrypt_db_variables(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._encrypt_env_variables_for_db +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__encrypt_env_variables_for_db_idempotent(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.decrypt_value_helper", + lambda value, key, return_original_value: value, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.encrypt_value_helper", + lambda value, new_encryption_key=None: f"ENC[{value}]", + ) + pc = ProxyConfig() + out = pc._encrypt_env_variables_for_db({"A": "1", "B": "2", "C": "3"}) + assert out == {"A": "ENC[1]", "B": "ENC[2]", "C": "ENC[3]"} + + +def test_ProxyConfig__encrypt_env_variables_for_db_invalid_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + pc._encrypt_env_variables_for_db(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._parse_router_settings_value +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__parse_router_settings_value_handles_inputs(): + result = { + "dict": ProxyConfig._parse_router_settings_value({"a": 1}), + "yaml_string": ProxyConfig._parse_router_settings_value("a: 1\nb: 2"), + "none": ProxyConfig._parse_router_settings_value(None), + } + assert result == { + "dict": {"a": 1}, + "yaml_string": {"a": 1, "b": 2}, + "none": None, + } + + +def test_ProxyConfig__parse_router_settings_value_invalid_returns_none(): + # Non-dict, non-parseable scalar -> None. + assert ProxyConfig._parse_router_settings_value(12345) is None + # Empty dict -> None (not truthy). + assert ProxyConfig._parse_router_settings_value({}) is None + + +# --------------------------------------------------------------------------- +# ProxyConfig._get_hierarchical_router_settings +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__get_hierarchical_router_settings_key_wins(): + pc = ProxyConfig() + fake_key = SimpleNamespace( + router_settings={"timeout": 30, "retries": 2, "model": "gpt-4"}, + team_id=None, + ) + out = await pc._get_hierarchical_router_settings( + user_api_key_dict=fake_key, + prisma_client=None, + proxy_logging_obj=None, + ) + assert out == {"timeout": 30, "retries": 2, "model": "gpt-4"} + + +@pytest.mark.asyncio +async def test_ProxyConfig__get_hierarchical_router_settings_missing_returns_none(): + pc = ProxyConfig() + fake_key = SimpleNamespace(router_settings=None, team_id=None) + out = await pc._get_hierarchical_router_settings( + user_api_key_dict=fake_key, + prisma_client=None, + proxy_logging_obj=None, + ) + assert out is None + + +# --------------------------------------------------------------------------- +# ProxyConfig._add_router_settings_from_db_config +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): + pc = ProxyConfig() + fake_router = MagicMock() + fake_router.update_settings = MagicMock() + fake_prisma = MagicMock() + fake_prisma.db.litellm_config.find_first = AsyncMock( + return_value=SimpleNamespace( + param_value={"timeout": 30, "retries": 2, "fallbacks": []} + ) + ) + config_data = {"router_settings": {"timeout": 10}} + await pc._add_router_settings_from_db_config( + config_data=config_data, + llm_router=fake_router, + prisma_client=fake_prisma, + ) + snapshot = { + "called": fake_router.update_settings.called, + "call_count": fake_router.update_settings.call_count, + "kwargs_keys": sorted( + list(fake_router.update_settings.call_args.kwargs.keys()) + ), + } + assert snapshot == { + "called": True, + "call_count": 1, + "kwargs_keys": ["fallbacks", "retries", "timeout"], + } + + +@pytest.mark.asyncio +async def test_ProxyConfig__add_router_settings_from_db_config_none_router_noop(): + pc = ProxyConfig() + # No router and no prisma — should silently return. + await pc._add_router_settings_from_db_config( + config_data={}, llm_router=None, prisma_client=None + ) + # Error-style: bad call signature raises. + with pytest.raises(TypeError): + await pc._add_router_settings_from_db_config() # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# ProxyConfig._add_general_settings_from_db_config +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__add_general_settings_from_db_config_merges_alerting(): + pc = ProxyConfig() + proxy_logging = MagicMock() + general = {"alerting": ["slack"]} + config_data = {"general_settings": {"alerting": ["email", "slack"]}} + pc._add_general_settings_from_db_config( + config_data=config_data, + general_settings=general, + proxy_logging_obj=proxy_logging, + ) + snapshot = { + "alerting": sorted(general["alerting"]), + "logging_called": proxy_logging.update_values.called, + "merged_count": len(general["alerting"]), + } + assert snapshot == { + "alerting": ["email", "slack"], + "logging_called": True, + "merged_count": 2, + } + + +def test_ProxyConfig__add_general_settings_from_db_config_bad_config_raises(): + pc = ProxyConfig() + with pytest.raises(AttributeError): + pc._add_general_settings_from_db_config( + config_data=None, # type: ignore[arg-type] + general_settings={}, + proxy_logging_obj=MagicMock(), + ) + + +# --------------------------------------------------------------------------- +# ProxyConfig._reschedule_spend_log_cleanup_job +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__reschedule_spend_log_cleanup_job_no_scheduler(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", None) + pc = ProxyConfig() + try: + await pc._reschedule_spend_log_cleanup_job() + raised = False + except Exception: + raised = True + snapshot = {"raised": raised, "called": True, "scheduler_was": "none"} + assert snapshot == {"raised": False, "called": True, "scheduler_was": "none"} + + +@pytest.mark.asyncio +async def test_ProxyConfig__reschedule_spend_log_cleanup_job_invalid_cron(monkeypatch): + fake_scheduler = MagicMock() + fake_scheduler.remove_job = MagicMock() + fake_scheduler.add_job = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", fake_scheduler) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + { + "maximum_spend_logs_retention_period": "1d", + "maximum_spend_logs_cleanup_cron": "INVALID CRON STRING", + }, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + pc = ProxyConfig() + # Invalid cron is caught and logged — does not raise outward. + await pc._reschedule_spend_log_cleanup_job() + # But add_job should not have been called for the invalid cron path. + assert fake_scheduler.add_job.call_count == 0 + + +# --------------------------------------------------------------------------- +# ProxyConfig._update_general_settings +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ProxyConfig__update_general_settings_updates_max_parallel(monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {}, + ) + pc = ProxyConfig() + await pc._update_general_settings( + { + "max_parallel_requests": 7, + "global_max_parallel_requests": 99, + "ui_access_mode": "admin_only", + } + ) + from litellm.proxy import proxy_server as ps + + snapshot = { + "max_parallel_requests": ps.general_settings.get("max_parallel_requests"), + "global_max_parallel_requests": ps.general_settings.get( + "global_max_parallel_requests" + ), + "ui_access_mode": ps.general_settings.get("ui_access_mode"), + } + assert snapshot == { + "max_parallel_requests": 7, + "global_max_parallel_requests": 99, + "ui_access_mode": "admin_only", + } + + +@pytest.mark.asyncio +async def test_ProxyConfig__update_general_settings_none_input_noop(): + pc = ProxyConfig() + # None input returns early. + result = await pc._update_general_settings(db_general_settings=None) + assert result is None + # Error-style: dict() will fail on non-mapping non-None input. + with pytest.raises(Exception): + await pc._update_general_settings(db_general_settings=12345) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# ProxyConfig._update_config_fields +# --------------------------------------------------------------------------- + + +def test_ProxyConfig__update_config_fields_merges_dict(): + pc = ProxyConfig() + current = {"general_settings": {"a": 1, "b": 2}} + out = pc._update_config_fields( + current_config=current, + param_name="general_settings", + db_param_value={"b": 3, "c": 4, "d": 5}, + ) + assert out == {"general_settings": {"a": 1, "b": 3, "c": 4, "d": 5}} + + +def test_ProxyConfig__update_config_fields_invalid_param_raises(): + pc = ProxyConfig() + with pytest.raises(Exception): + # Missing required arg. + pc._update_config_fields(current_config={}, param_name="general_settings") # type: ignore[call-arg] diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py b/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py index ad6b4016461..7ef29b71bf0 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py @@ -1 +1,370 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py Anthropic-beta-headers reload routes (PR3). + +Routes covered: +- POST /reload/anthropic_beta_headers +- POST /schedule/anthropic_beta_headers_reload +- DELETE /schedule/anthropic_beta_headers_reload +- GET /schedule/anthropic_beta_headers_reload/status +""" + +from __future__ import annotations + +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import VOLATILE_KEYS, normalize + +# These routes return a "timestamp" ISO string that isn't in the default +# volatile-keys set — extend the set locally so dict-equality assertions +# can ignore it. +_VOLATILE = VOLATILE_KEYS | frozenset({"timestamp"}) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_prisma_with_config( + config_record=None, +): + """Build a MagicMock prisma_client with a ``db.litellm_config`` namespace. + + The conftest's ``mock_prisma`` fixture stubs ``litellm_configtable`` but + the anthropic-beta routes use ``prisma_client.db.litellm_config`` — + a different attribute. Build one here so each test gets isolated state. + """ + config = MagicMock() + config.find_unique = AsyncMock(return_value=config_record) + config.upsert = AsyncMock() + config.delete = AsyncMock() + + db = MagicMock() + db.litellm_config = config + + client = MagicMock() + client.db = db + return client + + +def _install_prisma(monkeypatch, prisma): + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", prisma) + + +def _stub_reload_beta_headers(monkeypatch, return_value=None): + """Replace ``litellm.anthropic_beta_headers_manager.reload_beta_headers_config`` + with a deterministic stub so the route never hits the network.""" + if return_value is None: + return_value = { + "anthropic": {"beta_headers": ["foo"]}, + "openai": {"beta_headers": ["bar"]}, + "provider_aliases": {"a": "b"}, + "description": "test", + } + import litellm.anthropic_beta_headers_manager as mgr + + stub = MagicMock(return_value=return_value) + monkeypatch.setattr(mgr, "reload_beta_headers_config", stub) + return stub + + +# --------------------------------------------------------------------------- +# POST /reload/anthropic_beta_headers +# --------------------------------------------------------------------------- + + +def test_reload_anthropic_beta_headers_admin_success(client, auth_as, monkeypatch): + """Admin can trigger immediate reload — handler returns providers count and + a success status. Pins the response dict shape.""" + from litellm.proxy._types import LitellmUserRoles + + _stub_reload_beta_headers(monkeypatch) + prisma = _make_prisma_with_config(config_record=None) + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/reload/anthropic_beta_headers") + + assert response.status_code == 200 + body = response.json() + # Two non-alias keys: "anthropic", "openai" + assert normalize(body, _VOLATILE) == { + "message": "Anthropic beta headers configuration reloaded successfully! 2 providers updated.", + "status": "success", + "providers_count": 2, + "timestamp": "", + } + # And the upsert was actually invoked (force_reload write). + prisma.db.litellm_config.upsert.assert_awaited_once() + + +def test_reload_anthropic_beta_headers_preserves_existing_interval( + client, auth_as, monkeypatch +): + """When an existing reload config has an interval set, the force-reload + write must preserve that interval (the route reads it back then upserts + with the same number). This pins the read-then-write behaviour.""" + from litellm.proxy._types import LitellmUserRoles + + _stub_reload_beta_headers(monkeypatch) + existing = SimpleNamespace( + param_name="anthropic_beta_headers_reload_config", + param_value={"interval_hours": 12, "force_reload": False}, + ) + prisma = _make_prisma_with_config(config_record=existing) + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/reload/anthropic_beta_headers") + + assert response.status_code == 200 + # The update branch's interval_hours was sourced from the existing record. + call_kwargs = prisma.db.litellm_config.upsert.await_args.kwargs + data = call_kwargs["data"] + update_payload = data["update"]["param_value"] + parsed = ( + json.loads(update_payload) + if isinstance(update_payload, str) + else update_payload + ) + assert parsed["interval_hours"] == 12 + assert parsed["force_reload"] is True + + +def test_reload_anthropic_beta_headers_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post("/reload/anthropic_beta_headers") + + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +def test_reload_anthropic_beta_headers_no_db_returns_500(client, auth_as, monkeypatch): + """When prisma_client is None the handler raises 500 with a clear message.""" + from litellm.proxy._types import LitellmUserRoles + + _install_prisma(monkeypatch, None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/reload/anthropic_beta_headers") + + assert response.status_code == 500 + assert "Database connection not available" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# POST /schedule/anthropic_beta_headers_reload +# --------------------------------------------------------------------------- + + +def test_schedule_anthropic_beta_headers_reload_admin_success( + client, auth_as, monkeypatch +): + """Happy path: admin schedules every N hours — response echoes interval.""" + from litellm.proxy._types import LitellmUserRoles + + prisma = _make_prisma_with_config() + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/schedule/anthropic_beta_headers_reload", params={"hours": 6} + ) + + assert response.status_code == 200 + assert normalize(response.json(), _VOLATILE) == { + "message": "Anthropic beta headers reload scheduled for every 6 hours", + "status": "success", + "interval_hours": 6, + "timestamp": "", + } + prisma.db.litellm_config.upsert.assert_awaited_once() + + +def test_schedule_anthropic_beta_headers_reload_zero_hours_400( + client, auth_as, monkeypatch +): + """``hours <= 0`` is rejected with 400 and a descriptive message.""" + from litellm.proxy._types import LitellmUserRoles + + prisma = _make_prisma_with_config() + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/schedule/anthropic_beta_headers_reload", params={"hours": 0} + ) + + assert response.status_code == 400 + assert "Hours must be greater than 0" in response.json().get("detail", "") + + +def test_schedule_anthropic_beta_headers_reload_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/schedule/anthropic_beta_headers_reload", params={"hours": 6} + ) + + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +def test_schedule_anthropic_beta_headers_reload_missing_hours_422(client, auth_as): + """``hours`` is a required query param — omitting it is a FastAPI 422.""" + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/schedule/anthropic_beta_headers_reload") + + assert response.status_code == 422 + assert "detail" in response.json() + + +# --------------------------------------------------------------------------- +# DELETE /schedule/anthropic_beta_headers_reload +# --------------------------------------------------------------------------- + + +def test_cancel_anthropic_beta_headers_reload_admin_success( + client, auth_as, monkeypatch +): + """Admin cancel: deletes the LiteLLM_Config row and returns success dict.""" + from litellm.proxy._types import LitellmUserRoles + + prisma = _make_prisma_with_config() + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.delete("/schedule/anthropic_beta_headers_reload") + + assert response.status_code == 200 + assert normalize(response.json(), _VOLATILE) == { + "message": "Anthropic beta headers reload schedule cancelled", + "status": "success", + "timestamp": "", + } + prisma.db.litellm_config.delete.assert_awaited_once_with( + where={"param_name": "anthropic_beta_headers_reload_config"} + ) + + +def test_cancel_anthropic_beta_headers_reload_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.delete("/schedule/anthropic_beta_headers_reload") + + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +def test_cancel_anthropic_beta_headers_reload_no_db_returns_500( + client, auth_as, monkeypatch +): + from litellm.proxy._types import LitellmUserRoles + + _install_prisma(monkeypatch, None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.delete("/schedule/anthropic_beta_headers_reload") + + assert response.status_code == 500 + assert "Database connection not available" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# GET /schedule/anthropic_beta_headers_reload/status +# --------------------------------------------------------------------------- + + +def test_get_anthropic_beta_headers_reload_status_scheduled( + client, auth_as, monkeypatch +): + """When a config row with ``interval_hours`` is present, ``scheduled`` is True + and ``interval_hours`` echoes the DB value. Pins the full response shape.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + record = SimpleNamespace( + param_name="anthropic_beta_headers_reload_config", + param_value={"interval_hours": 6, "force_reload": False}, + ) + prisma = _make_prisma_with_config(config_record=record) + _install_prisma(monkeypatch, prisma) + # No prior reload — next_run stays None. + monkeypatch.setattr(ps, "last_anthropic_beta_headers_reload", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/anthropic_beta_headers_reload/status") + + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": True, + "interval_hours": 6, + "last_run": None, + "next_run": None, + } + + +def test_get_anthropic_beta_headers_reload_status_not_scheduled_no_db( + client, auth_as, monkeypatch +): + """No DB connection: handler returns the unscheduled-status dict (not 500).""" + from litellm.proxy._types import LitellmUserRoles + + _install_prisma(monkeypatch, None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/anthropic_beta_headers_reload/status") + + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None, + } + + +def test_get_anthropic_beta_headers_reload_status_no_interval_unscheduled( + client, auth_as, monkeypatch +): + """Config row present but ``interval_hours`` is None → unscheduled response.""" + from litellm.proxy._types import LitellmUserRoles + + record = SimpleNamespace( + param_name="anthropic_beta_headers_reload_config", + param_value={"interval_hours": None, "force_reload": True}, + ) + prisma = _make_prisma_with_config(config_record=record) + _install_prisma(monkeypatch, prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/anthropic_beta_headers_reload/status") + + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None, + } + + +def test_get_anthropic_beta_headers_reload_status_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/schedule/anthropic_beta_headers_reload/status") + + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_assistants.py b/tests/test_litellm/proxy/proxy_server/test_routes_assistants.py index ad6b4016461..fd1f672dce9 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_assistants.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_assistants.py @@ -1 +1,180 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` assistants routes. + +Pins (PR2): + - GET /v1/assistants + - GET /assistants + - POST /v1/assistants + - POST /assistants + - DELETE /v1/assistants/{assistant_id:path} + - DELETE /assistants/{assistant_id:path} +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +GET_RESPONSE = { + "object": "list", + "data": [ + { + "id": "asst_1", + "object": "assistant", + "name": "Test Assistant", + "model": "gpt-4", + } + ], + "first_id": "asst_1", + "last_id": "asst_1", + "has_more": False, +} + + +CREATE_RESPONSE = { + "id": "asst_new", + "object": "assistant", + "name": "New", + "model": "gpt-4", + "created_at": 0, +} + + +DELETE_RESPONSE = {"id": "asst_1", "object": "assistant.deleted", "deleted": True} + + +@pytest.fixture +def patched_assistants(monkeypatch): + router = MagicMock() + router.aget_assistants = AsyncMock(return_value=dict(GET_RESPONSE)) + router.acreate_assistants = AsyncMock(return_value=dict(CREATE_RESPONSE)) + router.adelete_assistant = AsyncMock(return_value=dict(DELETE_RESPONSE)) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + post_call_failure_hook=AsyncMock(), update_request_status=AsyncMock() + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + return router + + +@pytest.fixture +def no_router(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + post_call_failure_hook=AsyncMock(), update_request_status=AsyncMock() + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + yield + + +# --------------------------------------------------------------------------- +# GET /v1/assistants, GET /assistants +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("path", ["/v1/assistants", "/assistants"]) +def test_get_assistants_happy_path(client, auth_as, patched_assistants, path): + """Pins ``GET /v1/assistants`` and ``GET /assistants``.""" + with auth_as(): + response = client.get(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "object": "list", + "data": [ + { + "id": "", + "object": "assistant", + "name": "Test Assistant", + "model": "gpt-4", + } + ], + "first_id": "asst_1", + "last_id": "asst_1", + "has_more": False, + } + + +@pytest.mark.parametrize("path", ["/v1/assistants", "/assistants"]) +def test_get_assistants_no_router_error(client, auth_as, no_router, path): + """Pins ``GET /v1/assistants`` and ``GET /assistants`` (error: no llm_router).""" + with auth_as(): + response = client.get(path) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# POST /v1/assistants, POST /assistants +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("path", ["/v1/assistants", "/assistants"]) +def test_create_assistant_happy_path(client, auth_as, patched_assistants, path): + """Pins ``POST /v1/assistants`` and ``POST /assistants``.""" + payload = {"model": "gpt-4", "name": "New"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "assistant", + "name": "New", + "model": "gpt-4", + "created_at": "", + } + + +@pytest.mark.parametrize("path", ["/v1/assistants", "/assistants"]) +def test_create_assistant_no_router_error(client, auth_as, no_router, path): + """Pins ``POST /v1/assistants`` and ``POST /assistants`` (error: no llm_router).""" + with auth_as(): + response = client.post(path, json={"model": "gpt-4"}) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# DELETE /v1/assistants/{assistant_id:path}, DELETE /assistants/{assistant_id:path} +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("path", ["/v1/assistants/asst_1", "/assistants/asst_1"]) +def test_delete_assistant_happy_path(client, auth_as, patched_assistants, path): + """Pins ``DELETE /v1/assistants/{assistant_id:path}`` and ``DELETE /assistants/{assistant_id:path}``.""" + with auth_as(): + response = client.delete(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "assistant.deleted", + "deleted": True, + } + + +@pytest.mark.parametrize("path", ["/v1/assistants/asst_1", "/assistants/asst_1"]) +def test_delete_assistant_no_router_error(client, auth_as, no_router, path): + """Pins ``DELETE /v1/assistants/{assistant_id:path}`` / ``DELETE /assistants/{assistant_id:path}`` (error).""" + with auth_as(): + response = client.delete(path) + assert response.status_code == 500 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py index ad6b4016461..d88bcf136e9 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py @@ -1 +1,193 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` audio routes. + +Pins (PR2): + - POST /v1/audio/speech + - POST /audio/speech + - POST /v1/audio/transcriptions + - POST /audio/transcriptions +""" + +from __future__ import annotations + +import io +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + + +@pytest.fixture +def patched_speech(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + pre_call_hook=AsyncMock(side_effect=lambda **kw: kw["data"]), + post_call_failure_hook=AsyncMock(), + update_request_status=AsyncMock(), + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + + class _FakeBinaryResp: + async def aiter_bytes(self, chunk_size: int = 8192): + async def _gen(): + yield b"\x00\x01\x02" + + return _gen() + + async def _llm_call(): + return _FakeBinaryResp() + + async def _fake_route_request(*args, **kwargs): + return _llm_call() + + monkeypatch.setattr(proxy_server, "route_request", _fake_route_request) + yield + + +@pytest.fixture +def patched_speech_error(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + pre_call_hook=AsyncMock(side_effect=lambda **kw: kw["data"]), + post_call_failure_hook=AsyncMock(), + update_request_status=AsyncMock(), + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + + async def _raise(*args, **kwargs): + raise ValueError("speech boom") + + monkeypatch.setattr(proxy_server, "route_request", _raise) + yield + + +@pytest.fixture +def patched_transcription(monkeypatch): + router = MagicMock() + router.model_names = ["whisper-1"] + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + pre_call_hook=AsyncMock(side_effect=lambda **kw: kw["data"]), + post_call_failure_hook=AsyncMock(), + post_call_response_headers_hook=AsyncMock(return_value={}), + update_request_status=AsyncMock(), + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + monkeypatch.setattr( + proxy_server, "check_file_size_under_limit", lambda **kwargs: True + ) + + async def _form_data(request): + from starlette.datastructures import FormData, UploadFile + + upload = UploadFile( + filename="audio.mp3", + file=io.BytesIO(b"\x00\x01\x02"), + ) + return FormData([("file", upload), ("model", "whisper-1")]) + + monkeypatch.setattr(proxy_server, "get_form_data", _form_data) + + async def _llm_call(): + return {"text": "hello world"} + + async def _fake_route_request(*args, **kwargs): + return _llm_call() + + monkeypatch.setattr(proxy_server, "route_request", _fake_route_request) + yield + + +@pytest.fixture +def patched_transcription_error(monkeypatch, patched_transcription): + async def _raise(*args, **kwargs): + raise ValueError("transcription boom") + + monkeypatch.setattr(proxy_server, "route_request", _raise) + yield + + +@pytest.mark.parametrize("path", ["/v1/audio/speech", "/audio/speech"]) +def test_audio_speech_happy_path(client, auth_as, patched_speech, path): + """Pins ``POST /v1/audio/speech`` and ``POST /audio/speech`` (happy).""" + payload = {"model": "tts-1", "input": "Hi", "voice": "alloy"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + response_summary = { + "status_code": response.status_code, + "content_type": response.headers.get("content-type", ""), + "body_bytes": response.content, + } + assert response_summary == { + "status_code": 200, + "content_type": "audio/mpeg", + "body_bytes": b"\x00\x01\x02", + } + + +@pytest.mark.parametrize("path", ["/v1/audio/speech", "/audio/speech"]) +def test_audio_speech_error(client, auth_as, patched_speech_error, path): + """Pins ``POST /v1/audio/speech`` and ``POST /audio/speech`` (error).""" + payload = {"model": "tts-1", "input": "Hi", "voice": "alloy"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 500 + assert len(response.content) > 0 + + +@pytest.mark.parametrize("path", ["/v1/audio/transcriptions", "/audio/transcriptions"]) +def test_audio_transcription_happy_path(client, auth_as, patched_transcription, path): + """Pins ``POST /v1/audio/transcriptions`` / ``POST /audio/transcriptions`` (happy).""" + files = {"file": ("audio.mp3", b"\x00\x01\x02", "audio/mpeg")} + data = {"model": "whisper-1"} + with auth_as(): + response = client.post(path, files=files, data=data) + assert response.status_code == 200 + body = response.json() + assert body == {"text": "hello world"} + response_summary = { + "status_code": response.status_code, + "text_field": body["text"], + "media_type_hint": response.headers.get("content-type", "").split(";")[0], + } + assert response_summary == { + "status_code": 200, + "text_field": "hello world", + "media_type_hint": "application/json", + } + + +@pytest.mark.parametrize("path", ["/v1/audio/transcriptions", "/audio/transcriptions"]) +def test_audio_transcription_error(client, auth_as, patched_transcription_error, path): + """Pins ``POST /v1/audio/transcriptions`` / ``POST /audio/transcriptions`` (error).""" + files = {"file": ("audio.mp3", b"\x00\x01\x02", "audio/mpeg")} + data = {"model": "whisper-1"} + with auth_as(): + response = client.post(path, files=files, data=data) + assert response.status_code == 500 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_chat_completions.py b/tests/test_litellm/proxy/proxy_server/test_routes_chat_completions.py index ad6b4016461..b186bb5ef5e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_chat_completions.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_chat_completions.py @@ -1 +1,134 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` chat-completions routes. + +Pins (PR2): + - POST /v1/chat/completions + - POST /chat/completions + - POST /engines/{model:path}/chat/completions + - POST /openai/deployments/{model:path}/chat/completions +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import common_request_processing, proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +HAPPY_RESPONSE = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-4", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "Hello from mock"}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} + + +@pytest.fixture +def patched_chat(monkeypatch): + """Stub chat-completions pipeline at ProxyBaseLLMRequestProcessing.""" + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + async def _fake_process(self, *args, **kwargs): + return dict(HAPPY_RESPONSE) + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _fake_process, + ) + yield + + +@pytest.fixture +def patched_chat_error(monkeypatch): + """Variant that makes the pipeline raise -> 400 via _handle_llm_api_exception.""" + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + from litellm.proxy._types import ProxyException + + async def _raise(self, *args, **kwargs): + raise ValueError("boom") + + async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj): + return ProxyException( + message="boom", type="bad_request_error", param="model", code=400 + ) + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _raise, + ) + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "_handle_llm_api_exception", + _handler, + ) + yield + + +_CHAT_PATHS = [ + "/v1/chat/completions", + "/chat/completions", + "/engines/gpt-4/chat/completions", + "/openai/deployments/gpt-4/chat/completions", +] + + +@pytest.mark.parametrize("path", _CHAT_PATHS) +def test_chat_completion_happy_path(client, auth_as, patched_chat, path): + """Pins all four ``POST .../chat/completions`` aliases (happy path). + + Covers ``POST /v1/chat/completions``, ``POST /chat/completions``, + ``POST /engines/{model:path}/chat/completions``, and + ``POST /openai/deployments/{model:path}/chat/completions``. + """ + payload = {"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "chat.completion", + "created": "", + "model": "gpt-4", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "Hello from mock"}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + + +@pytest.mark.parametrize("path", _CHAT_PATHS) +def test_chat_completion_pipeline_error(client, auth_as, patched_chat_error, path): + """Pins all four ``POST .../chat/completions`` aliases (error: 400). + + Covers ``POST /v1/chat/completions``, ``POST /chat/completions``, + ``POST /engines/{model:path}/chat/completions``, and + ``POST /openai/deployments/{model:path}/chat/completions``. + """ + payload = {"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 400 + assert "error" in response.json() or response.text != "" diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_completions.py b/tests/test_litellm/proxy/proxy_server/test_routes_completions.py index ad6b4016461..b5c60c23c02 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_completions.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_completions.py @@ -1 +1,126 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` text-completions routes. + +Pins (PR2): + - POST /v1/completions + - POST /completions + - POST /engines/{model:path}/completions + - POST /openai/deployments/{model:path}/completions +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import common_request_processing, proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +HAPPY_RESPONSE = { + "id": "cmpl-test", + "object": "text_completion", + "created": 0, + "model": "gpt-3.5-turbo-instruct", + "choices": [ + { + "index": 0, + "text": "Hello from mock", + "finish_reason": "stop", + "logprobs": None, + } + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5}, +} + + +@pytest.fixture +def patched_completion(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + async def _fake_process(self, *args, **kwargs): + return dict(HAPPY_RESPONSE) + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _fake_process, + ) + yield + + +@pytest.fixture +def completion_pipeline_raises(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + async def _raise(self, *args, **kwargs): + raise ValueError("boom") + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _raise, + ) + yield + + +_COMPLETION_PATHS = [ + "/v1/completions", + "/completions", + "/engines/gpt-3.5-turbo-instruct/completions", + "/openai/deployments/gpt-3.5-turbo-instruct/completions", +] + + +@pytest.mark.parametrize("path", _COMPLETION_PATHS) +def test_completion_happy_path(client, auth_as, patched_completion, path): + """Pins all four ``POST .../completions`` aliases (happy path). + + Covers ``POST /v1/completions``, ``POST /completions``, + ``POST /engines/{model:path}/completions``, and + ``POST /openai/deployments/{model:path}/completions``. + """ + payload = { + "model": "gpt-3.5-turbo-instruct", + "prompt": "Once upon", + "max_tokens": 5, + } + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "text_completion", + "created": "", + "model": "gpt-3.5-turbo-instruct", + "choices": [ + { + "index": 0, + "text": "Hello from mock", + "finish_reason": "stop", + "logprobs": None, + } + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5}, + } + + +@pytest.mark.parametrize("path", _COMPLETION_PATHS) +def test_completion_pipeline_error(client, auth_as, completion_pipeline_raises, path): + """Pins all four ``POST .../completions`` aliases (error path). + + Covers ``POST /v1/completions``, ``POST /completions``, + ``POST /engines/{model:path}/completions``, and + ``POST /openai/deployments/{model:path}/completions``. + """ + payload = {"model": "gpt-3.5-turbo-instruct", "prompt": "boom"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 500 + assert response.headers.get("content-type", "").startswith("application/json") diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index ad6b4016461..e89ada5bdef 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -1 +1,591 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py control-plane config routes (PR3). + +Routes covered: +- POST /config/update +- POST /config/field/update +- GET /config/field/info +- GET /config/list +- POST /config/field/delete +- POST /config/callback/delete +- GET /get/config/callbacks +- GET /config/yaml +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import VOLATILE_KEYS, normalize + + +def _install_litellm_config(mock_prisma: MagicMock) -> MagicMock: + """Ensure mock_prisma.db.litellm_config exists with async methods (the + conftest only stubs ``litellm_configtable`` — this is a different table).""" + table = MagicMock() + table.find_unique = AsyncMock(return_value=None) + table.find_first = AsyncMock(return_value=None) + table.find_many = AsyncMock(return_value=[]) + table.create = AsyncMock() + table.update = AsyncMock() + table.upsert = AsyncMock(return_value=None) + table.delete = AsyncMock() + mock_prisma.db.litellm_config = table + return table + + +# --------------------------------------------------------------------------- +# POST /config/update +# --------------------------------------------------------------------------- + + +def test_config_update_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """POST /config/update with admin role merges + upserts general_settings + and returns the canonical success message.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + fake_proxy_config = MagicMock() + fake_proxy_config.add_deployment = AsyncMock() + monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/update", + json={"general_settings": {"alerting": ["slack"]}}, + ) + assert response.status_code == 200 + assert normalize(response.json()) == {"message": "Config updated successfully"} + + +def test_config_update_non_admin_forbidden(client, auth_as, mock_prisma, monkeypatch): + """POST /config/update by a non-admin caller is rejected; the error + surfaces as a ProxyException with the admin-only message.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/config/update", + json={"general_settings": {"alerting": ["slack"]}}, + ) + assert response.status_code != 200 + body = response.json() + # ProxyException wraps the 403 detail string in its `message` field. + assert "admin" in str(body).lower() or "auth" in str(body).lower() + + +def test_config_update_no_db_error(client, auth_as, monkeypatch): + """POST /config/update with prisma_client=None returns a 'No DB Connected' + style error (the route raises Exception which the handler maps to 400).""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/update", + json={"general_settings": {"alerting": ["slack"]}}, + ) + assert response.status_code != 200 + assert ( + "db" in str(response.json()).lower() + or "connect" in str(response.json()).lower() + ) + + +# --------------------------------------------------------------------------- +# POST /config/field/update +# --------------------------------------------------------------------------- + + +def test_config_field_update_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """POST /config/field/update for a known field upserts the DB row and + returns the upsert response (we pin it to a specific shape).""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + table.find_first = AsyncMock(return_value=None) + upsert_row = { + "param_name": "general_settings", + "param_value": {"max_parallel_requests": 5}, + "id": "row-1", + } + table.upsert = AsyncMock(return_value=upsert_row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/update", + json={ + "field_name": "max_parallel_requests", + "field_value": 5, + "config_type": "general_settings", + }, + ) + assert response.status_code == 200 + assert normalize(response.json()) == { + "param_name": "general_settings", + "param_value": {"max_parallel_requests": 5}, + "id": "", + } + + +def test_config_field_update_non_admin_rejected( + client, auth_as, mock_prisma, monkeypatch +): + """Non-admin cannot update config fields — returns 400 with not-allowed + detail (handler uses 400 for the auth gate, not 403).""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/config/field/update", + json={ + "field_name": "max_parallel_requests", + "field_value": 5, + "config_type": "general_settings", + }, + ) + assert response.status_code == 400 + assert "error" in response.json().get("detail", {}) + + +def test_config_field_update_invalid_field(client, auth_as, mock_prisma, monkeypatch): + """Unknown field_name is rejected with 400 + 'Invalid field=' detail.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/update", + json={ + "field_name": "not_a_real_field_xyz", + "field_value": 1, + "config_type": "general_settings", + }, + ) + assert response.status_code == 400 + assert "Invalid field" in response.json().get("detail", {}).get("error", "") + + +# --------------------------------------------------------------------------- +# GET /config/field/info +# --------------------------------------------------------------------------- + + +def test_config_field_info_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """Admin gets back ConfigFieldInfo with the stored value pulled from DB.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + row = MagicMock() + row.param_value = {"max_parallel_requests": 7} + table.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get( + "/config/field/info", params={"field_name": "max_parallel_requests"} + ) + assert response.status_code == 200 + assert normalize(response.json()) == { + "field_name": "max_parallel_requests", + "field_value": 7, + } + + +def test_config_field_info_non_admin_rejected( + client, auth_as, mock_prisma, monkeypatch +): + """Non-admin (INTERNAL_USER) is denied — admin-view gate fires.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get( + "/config/field/info", params={"field_name": "max_parallel_requests"} + ) + assert response.status_code == 400 + assert "error" in response.json().get("detail", {}) + + +def test_config_field_info_field_not_in_db(client, auth_as, mock_prisma, monkeypatch): + """When the field is missing from the DB row, returns 400 'not in DB'.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + row = MagicMock() + row.param_value = {"some_other_field": "value"} + table.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get( + "/config/field/info", params={"field_name": "max_parallel_requests"} + ) + assert response.status_code == 400 + assert "not in DB" in response.json().get("detail", {}).get("error", "") + + +# --------------------------------------------------------------------------- +# GET /config/list +# --------------------------------------------------------------------------- + + +def test_config_list_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """Admin gets a non-empty list of ConfigList rows for general_settings + (one entry per known allowed_arg). Each row has the documented schema.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + row = MagicMock() + row.param_value = {"max_parallel_requests": 3} + table.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get( + "/config/list", params={"config_type": "general_settings"} + ) + assert response.status_code == 200 + body = response.json() + assert isinstance(body, list) + assert len(body) > 0 + sample = body[0] + shape = { + "has_field_name": "field_name" in sample, + "has_field_type": "field_type" in sample, + "has_field_value": "field_value" in sample, + "has_stored_in_db": "stored_in_db" in sample, + } + assert shape == { + "has_field_name": True, + "has_field_type": True, + "has_field_value": True, + "has_stored_in_db": True, + } + + +def test_config_list_non_admin_rejected(client, auth_as, mock_prisma, monkeypatch): + """Non-admin gets a 400 with the role embedded in the error message.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get( + "/config/list", params={"config_type": "general_settings"} + ) + assert response.status_code == 400 + assert "role" in response.json().get("detail", {}).get("error", "").lower() + + +def test_config_list_no_db_error(client, auth_as, monkeypatch): + """No DB → 400 with db_not_connected error.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get( + "/config/list", params={"config_type": "general_settings"} + ) + assert response.status_code == 400 + assert "error" in response.json().get("detail", {}) + + +# --------------------------------------------------------------------------- +# POST /config/field/delete +# --------------------------------------------------------------------------- + + +def test_config_field_delete_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """Admin can delete a stored general_settings field — returns the upsert row.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + existing = MagicMock() + existing.param_value = {"max_parallel_requests": 5, "other": "value"} + table.find_first = AsyncMock(return_value=existing) + table.upsert = AsyncMock( + return_value={ + "param_name": "general_settings", + "param_value": {"other": "value"}, + "id": "row-1", + } + ) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/delete", + json={ + "config_type": "general_settings", + "field_name": "max_parallel_requests", + }, + ) + assert response.status_code == 200 + assert normalize(response.json()) == { + "param_name": "general_settings", + "param_value": {"other": "value"}, + "id": "", + } + + +def test_config_field_delete_non_admin_rejected( + client, auth_as, mock_prisma, monkeypatch +): + """Non-admin caller hits the 400 not-allowed branch with role in detail.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/config/field/delete", + json={ + "config_type": "general_settings", + "field_name": "max_parallel_requests", + }, + ) + assert response.status_code == 400 + assert "role" in response.json().get("detail", {}).get("error", "").lower() + + +def test_config_field_delete_field_not_in_config( + client, auth_as, mock_prisma, monkeypatch +): + """If there is no general_settings row at all, returns 400 'not in config'.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + table.find_first = AsyncMock(return_value=None) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/delete", + json={ + "config_type": "general_settings", + "field_name": "max_parallel_requests", + }, + ) + assert response.status_code == 400 + assert "not in config" in response.json().get("detail", {}).get("error", "") + + +# --------------------------------------------------------------------------- +# POST /config/callback/delete +# --------------------------------------------------------------------------- + + +def test_config_callback_delete_happy_admin(client, auth_as, mock_prisma, monkeypatch): + """Admin deletes a configured success callback — handler returns the + success message + remaining callbacks + a timestamp.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "store_model_in_db", True) + + fake_proxy_config = MagicMock() + fake_proxy_config.get_config = AsyncMock( + return_value={"litellm_settings": {"success_callback": ["langfuse", "slack"]}} + ) + fake_proxy_config.save_config = AsyncMock() + fake_proxy_config.add_deployment = AsyncMock() + monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/callback/delete", json={"callback_name": "langfuse"} + ) + assert response.status_code == 200 + # `deleted_at` is an ISO timestamp generated at request time — extend + # the volatile set just for this assertion so dict-equality still works. + volatile = VOLATILE_KEYS | {"deleted_at"} + assert normalize(response.json(), volatile) == { + "message": "Successfully deleted callback: langfuse", + "removed_callback": "langfuse", + "remaining_callbacks": ["slack"], + "deleted_at": "", + } + + +def test_config_callback_delete_non_admin_rejected( + client, auth_as, mock_prisma, monkeypatch +): + """Non-admin caller is rejected with 400 not-allowed.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "store_model_in_db", True) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/config/callback/delete", json={"callback_name": "langfuse"} + ) + assert response.status_code == 400 + assert "role" in response.json().get("detail", {}).get("error", "").lower() + + +def test_config_callback_delete_not_found(client, auth_as, mock_prisma, monkeypatch): + """Callback missing from current config returns 404.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "store_model_in_db", True) + + fake_proxy_config = MagicMock() + fake_proxy_config.get_config = AsyncMock( + return_value={"litellm_settings": {"success_callback": ["slack"]}} + ) + monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/callback/delete", json={"callback_name": "langfuse"} + ) + # The handler re-raises HTTPException(404) verbatim (only generic + # `Exception` becomes a 500 ProxyException), so pin 404 strictly. + assert response.status_code == 404 + assert ( + "langfuse" in str(response.json()).lower() + or "not found" in str(response.json()).lower() + ) + + +# --------------------------------------------------------------------------- +# GET /get/config/callbacks +# --------------------------------------------------------------------------- + + +def test_get_config_callbacks_happy(client, auth_as, mock_prisma, monkeypatch): + """GET /get/config/callbacks returns the 5 pinned top-level keys: + status, callbacks, alerts, router_settings, available_callbacks.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "llm_router", None) + + fake_proxy_config = MagicMock() + fake_proxy_config.get_config = AsyncMock( + return_value={ + "litellm_settings": {"success_callback": []}, + "general_settings": {}, + "environment_variables": {}, + } + ) + monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/get/config/callbacks") + assert response.status_code == 200 + body = response.json() + shape = { + "status": body.get("status"), + "has_callbacks": "callbacks" in body, + "has_alerts": "alerts" in body, + "has_router_settings": "router_settings" in body, + "has_available_callbacks": "available_callbacks" in body, + } + assert shape == { + "status": "success", + "has_callbacks": True, + "has_alerts": True, + "has_router_settings": True, + "has_available_callbacks": True, + } + + +def test_get_config_callbacks_internal_error(client, auth_as, mock_prisma, monkeypatch): + """If proxy_config.get_config() raises, the handler wraps the failure in + a ProxyException → non-2xx response with an error body.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + fake_proxy_config = MagicMock() + fake_proxy_config.get_config = AsyncMock(side_effect=RuntimeError("boom")) + monkeypatch.setattr(ps, "proxy_config", fake_proxy_config) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/get/config/callbacks") + assert response.status_code >= 400 + assert ( + "boom" in str(response.json()).lower() + or "error" in str(response.json()).lower() + ) + + +# --------------------------------------------------------------------------- +# GET /config/yaml +# --------------------------------------------------------------------------- + + +def test_config_yaml_returns_demo_payload(client, auth_as): + """GET /config/yaml is documented as a mock endpoint. It declares + ConfigYAML as the body parameter, so a GET with an empty JSON body is + accepted and returns the canonical demo dict.""" + with auth_as(): + response = client.request("GET", "/config/yaml", json={}) + shape = { + "status": response.status_code, + "media_type_yaml": response.headers.get("content-type", "").startswith( + "application/json" + ), + "has_body": len(response.content) > 0, + } + assert shape == { + "status": 200, + "media_type_yaml": True, + "has_body": True, + } + assert response.json() == {"hello": "world"} + + +def test_config_yaml_invalid_method(client): + """POST against the GET-only /config/yaml is rejected (error path).""" + response = client.post("/config/yaml", json={}) + assert response.status_code == 405 + # Method-not-allowed responses still return a JSON-ish body via the + # FastAPI default handler — assert the body is not the success payload. + assert response.content != b'{"hello":"world"}' diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py index ad6b4016461..98249cb5ad5 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py @@ -1 +1,121 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` embeddings routes. + +Pins (PR2): + - POST /v1/embeddings + - POST /embeddings + - POST /engines/{model:path}/embeddings + - POST /openai/deployments/{model:path}/embeddings +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import common_request_processing, proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +HAPPY_RESPONSE = { + "object": "list", + "model": "text-embedding-ada-002", + "data": [{"embedding": [0.0, 0.1, 0.2], "index": 0, "object": "embedding"}], + "usage": {"prompt_tokens": 1, "total_tokens": 1}, +} + + +@pytest.fixture +def patched_embedding(monkeypatch): + router = MagicMock() + router.model_names = ["text-embedding-ada-002"] + router.get_deployment_by_model_group_name = MagicMock(return_value=None) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + async def _fake_process(self, *args, **kwargs): + return dict(HAPPY_RESPONSE) + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _fake_process, + ) + yield + + +@pytest.fixture +def embedding_pipeline_raises(monkeypatch): + router = MagicMock() + router.model_names = [] + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()) + ) + + from litellm.proxy._types import ProxyException + + async def _raise(self, *args, **kwargs): + raise ValueError("boom") + + async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None): + return ProxyException( + message="boom", type="bad_request_error", param="model", code=400 + ) + + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + _raise, + ) + monkeypatch.setattr( + common_request_processing.ProxyBaseLLMRequestProcessing, + "_handle_llm_api_exception", + _handler, + ) + yield + + +_EMBED_PATHS = [ + "/v1/embeddings", + "/embeddings", + "/engines/text-embedding-ada-002/embeddings", + "/openai/deployments/text-embedding-ada-002/embeddings", +] + + +@pytest.mark.parametrize("path", _EMBED_PATHS) +def test_embeddings_happy_path(client, auth_as, patched_embedding, path): + """Pins all four ``POST .../embeddings`` aliases (happy path). + + Covers ``POST /v1/embeddings``, ``POST /embeddings``, + ``POST /engines/{model:path}/embeddings``, and + ``POST /openai/deployments/{model:path}/embeddings``. + """ + payload = {"model": "text-embedding-ada-002", "input": "hello"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "object": "list", + "model": "text-embedding-ada-002", + "data": [{"embedding": [0.0, 0.1, 0.2], "index": 0, "object": "embedding"}], + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + + +@pytest.mark.parametrize("path", _EMBED_PATHS) +def test_embeddings_pipeline_error(client, auth_as, embedding_pipeline_raises, path): + """Pins all four ``POST .../embeddings`` aliases (error path). + + Covers ``POST /v1/embeddings``, ``POST /embeddings``, + ``POST /engines/{model:path}/embeddings``, and + ``POST /openai/deployments/{model:path}/embeddings``. + """ + payload = {"model": "text-embedding-ada-002", "input": "boom"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 400 + assert response.content # non-empty error body diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py b/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py index ad6b4016461..5b54a63d8a2 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_invitation.py @@ -1 +1,387 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py invitation routes (PR3). + +Routes covered: +- POST /invitation/new +- GET /invitation/info +- POST /invitation/update +- POST /invitation/delete +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import VOLATILE_KEYS, normalize + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_invitation( + invitation_id: str = "inv-abc", + user_id: str = "user-target", + created_by: str = "test-user-id", + is_accepted: bool = False, + accepted_at=None, +): + """Build an invitation row with the fields ``InvitationModel`` requires. + + FastAPI serializes the returned object against ``response_model=InvitationModel``, + so the object must expose ``id, user_id, is_accepted, accepted_at, expires_at, + created_at, created_by, updated_at, updated_by`` either as attributes or + dict keys. + """ + now = datetime.now(timezone.utc) + return SimpleNamespace( + id=invitation_id, + user_id=user_id, + is_accepted=is_accepted, + accepted_at=accepted_at, + expires_at=now + timedelta(days=7), + created_at=now, + created_by=created_by, + updated_at=now, + updated_by=created_by, + ) + + +# --------------------------------------------------------------------------- +# POST /invitation/new +# --------------------------------------------------------------------------- + + +def test_invitation_new_admin_happy(client, auth_as, monkeypatch, mock_prisma): + """Proxy admin → create_invitation_for_user returns invitation → 200.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.management_helpers import user_invitation + + invitation = _make_invitation(user_id="user-target") + + async def _fake_create_invitation(data, user_api_key_dict): + return invitation + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr( + user_invitation, "create_invitation_for_user", _fake_create_invitation + ) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/invitation/new", json={"user_id": "user-target"}) + + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "user_id": "user-target", + "is_accepted": False, + "accepted_at": None, + "expires_at": "", + "created_at": "", + "created_by": "test-user-id", + "updated_at": "", + "updated_by": "test-user-id", + } + + +def test_invitation_new_non_admin_forbidden(client, auth_as, monkeypatch, mock_prisma): + """Internal user without team/org admin privileges → 400 not-allowed.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.management_endpoints import common_utils + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + async def _no_privileges(**kwargs): + return False + + # Patch at the proxy_server import site (used by the route). + monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) + monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _no_privileges) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post("/invitation/new", json={"user_id": "user-target"}) + + assert response.status_code == 400 + err = response.json().get("error", response.json()) + err_text = str(err) + assert "role=" in err_text or "not allowed" in err_text.lower() + + +def test_invitation_new_db_not_connected_400(client, auth_as, monkeypatch): + """prisma_client is None → 400 db_not_connected_error.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/invitation/new", json={"user_id": "user-target"}) + + assert response.status_code == 400 + body = response.json() + err_text = str(body) + # The handler wraps via handle_exception_on_proxy, so the error body + # may take either the {"error": {...}} or {"detail": {...}} shape. + assert "No connected db" in err_text or "db" in err_text.lower() + + +def test_invitation_new_missing_user_id_422(client, auth_as, monkeypatch, mock_prisma): + """Body missing the required ``user_id`` field → FastAPI 422.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/invitation/new", json={}) + + assert response.status_code == 422 + body = response.json() + assert isinstance(body.get("detail"), list) + assert any("user_id" in str(item) for item in body["detail"]) + + +# --------------------------------------------------------------------------- +# GET /invitation/info +# --------------------------------------------------------------------------- + + +def test_invitation_info_admin_happy(client, auth_as, monkeypatch, mock_prisma): + """Admin requesting an existing invitation id → returns the invitation.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + invitation = _make_invitation(invitation_id="inv-xyz", user_id="user-target") + mock_prisma.db.litellm_invitationlink.find_unique.return_value = invitation + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/invitation/info", params={"invitation_id": "inv-xyz"}) + + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "user_id": "user-target", + "is_accepted": False, + "accepted_at": None, + "expires_at": "", + "created_at": "", + "created_by": "test-user-id", + "updated_at": "", + "updated_by": "test-user-id", + } + + +def test_invitation_info_not_admin_forbidden(client, auth_as, monkeypatch, mock_prisma): + """Non-admin viewer (no admin-view privileges) → 400 not-allowed.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + # _user_has_admin_view is referenced from proxy_server's import. + monkeypatch.setattr(ps, "_user_has_admin_view", lambda u: False) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/invitation/info", params={"invitation_id": "inv-xyz"}) + + assert response.status_code == 400 + err_text = str(response.json()) + assert "role=" in err_text or "not allowed" in err_text.lower() + + +def test_invitation_info_not_found_400(client, auth_as, monkeypatch, mock_prisma): + """Admin requesting an unknown invitation id → 400 does-not-exist.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + mock_prisma.db.litellm_invitationlink.find_unique.return_value = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get( + "/invitation/info", params={"invitation_id": "does-not-exist"} + ) + + assert response.status_code == 400 + assert response.json() == { + "detail": {"error": "Invitation id does not exist in the database."} + } + + +# --------------------------------------------------------------------------- +# POST /invitation/update +# --------------------------------------------------------------------------- + + +def test_invitation_update_happy(client, auth_as, monkeypatch, mock_prisma): + """Authenticated user → invitation marked accepted → returns updated row.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + accepted = _make_invitation( + invitation_id="inv-1", + user_id="user-target", + is_accepted=True, + accepted_at=datetime.now(timezone.utc), + ) + mock_prisma.db.litellm_invitationlink.update.return_value = accepted + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/invitation/update", + json={"invitation_id": "inv-1", "is_accepted": True}, + ) + + assert response.status_code == 200 + # ``accepted_at`` is a fresh timestamp on each run — extend volatile set. + extended = VOLATILE_KEYS | {"accepted_at"} + assert normalize(response.json(), extended) == { + "id": "", + "user_id": "user-target", + "is_accepted": True, + "accepted_at": "", + "expires_at": "", + "created_at": "", + "created_by": "test-user-id", + "updated_at": "", + "updated_by": "test-user-id", + } + + +def test_invitation_update_unknown_id_400(client, auth_as, monkeypatch, mock_prisma): + """Update against an invitation id the DB returns None for → 400.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + mock_prisma.db.litellm_invitationlink.update.return_value = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/invitation/update", + json={"invitation_id": "ghost", "is_accepted": True}, + ) + + assert response.status_code == 400 + assert response.json() == { + "detail": {"error": "Invitation id does not exist in the database."} + } + + +def test_invitation_update_no_user_id_500(client, auth_as, monkeypatch, mock_prisma): + """If the auth principal lacks a user_id, handler returns 500.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN, user_id=None): + response = client.post( + "/invitation/update", + json={"invitation_id": "inv-1", "is_accepted": True}, + ) + + assert response.status_code == 500 + err_text = str(response.json()) + assert "Unable to identify user id" in err_text + + +# --------------------------------------------------------------------------- +# POST /invitation/delete +# --------------------------------------------------------------------------- + + +def test_invitation_delete_admin_happy(client, auth_as, monkeypatch, mock_prisma): + """Proxy admin deletes by invitation_id → 200 with deleted row.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + deleted = _make_invitation(invitation_id="inv-del", user_id="user-target") + mock_prisma.db.litellm_invitationlink.delete.return_value = deleted + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/invitation/delete", json={"invitation_id": "inv-del"} + ) + + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "user_id": "user-target", + "is_accepted": False, + "accepted_at": None, + "expires_at": "", + "created_at": "", + "created_by": "test-user-id", + "updated_at": "", + "updated_by": "test-user-id", + } + + +def test_invitation_delete_non_admin_forbidden( + client, auth_as, monkeypatch, mock_prisma +): + """Non-admin user without elevated privileges → 400 not-allowed.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + async def _no_privileges(**kwargs): + return False + + monkeypatch.setattr(ps, "_user_has_admin_privileges", _no_privileges) + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post( + "/invitation/delete", json={"invitation_id": "inv-del"} + ) + + assert response.status_code == 400 + err_text = str(response.json()) + assert "role=" in err_text or "not allowed" in err_text.lower() + + +def test_invitation_delete_unknown_id_400(client, auth_as, monkeypatch, mock_prisma): + """Delete returns None (no row) → 400 does-not-exist.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + mock_prisma.db.litellm_invitationlink.delete.return_value = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/invitation/delete", json={"invitation_id": "ghost"} + ) + + assert response.status_code == 400 + assert response.json() == { + "detail": {"error": "Invitation id does not exist in the database."} + } + + +def test_invitation_delete_db_not_connected_400(client, auth_as, monkeypatch): + """prisma_client is None → 400 db_not_connected_error.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/invitation/delete", json={"invitation_id": "inv-del"} + ) + + assert response.status_code == 400 + err_text = str(response.json()) + assert "No connected db" in err_text or "db" in err_text.lower() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index ad6b4016461..6af1d6653e1 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -1 +1,387 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py login/SSO routes (PR3). + +Routes covered: +- GET /fallback/login +- POST /login +- POST /v2/login +- POST /v3/login +- POST /v3/login/exchange +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import normalize + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None: + """Patch authenticate_user + create_ui_token_object at their import paths. + + Both /login, /v2/login and /v3/login do a *local* (in-function) import of + these helpers, so we patch the module they live in. + """ + from litellm.proxy import proxy_server as ps + + async def _fake_auth(username, password, master_key, prisma_client): + if raise_on_auth: + raise Exception("boom-auth-failure") + fake = MagicMock() + fake.user_id = "u-1" + fake.user_email = "test@example.invalid" + fake.user_role = "proxy_admin" + fake.key = "sk-fake-ui-key" + return fake + + def _fake_token_object(login_result, general_settings, premium_user): + return { + "user_id": "u-1", + "user_email": "test@example.invalid", + "user_role": "proxy_admin", + "premium_user": premium_user, + "key": "sk-fake-ui-key", + } + + monkeypatch.setattr( + "litellm.proxy.auth.login_utils.authenticate_user", _fake_auth + ) + monkeypatch.setattr( + "litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object + ) + monkeypatch.setattr(ps, "master_key", "sk-test-master") + monkeypatch.setattr(ps, "general_settings", {}) + monkeypatch.setattr(ps, "premium_user", False) + + +# --------------------------------------------------------------------------- +# GET /fallback/login +# --------------------------------------------------------------------------- + + +def test_fallback_login_returns_html_form(client, monkeypatch): + """Pin: GET /fallback/login returns an HTML login form with status 200.""" + monkeypatch.delenv("UI_USERNAME", raising=False) + response = client.get("/fallback/login") + body_lower = response.text.lower() + shape = { + "status": response.status_code, + "content_type_html": response.headers.get("content-type", "").startswith( + "text/html" + ), + "has_form": " TestClient returns 500 with body + assert response.status_code == 500 + # Body must be non-empty so a future refactor that drops the error body + # would trip this gate. + assert len(response.content) > 0 + assert response.headers.get("content-type") is not None + + +# --------------------------------------------------------------------------- +# POST /v2/login +# --------------------------------------------------------------------------- + + +def test_v2_login_success_returns_token_and_redirect(client, monkeypatch): + """Pin: POST /v2/login returns JSON {redirect_url, token} + sets token cookie.""" + _install_login_mocks(monkeypatch) + response = client.post( + "/v2/login", + json={"username": "admin", "password": "password"}, + ) + assert response.status_code == 200 + assert normalize( + response.json(), volatile=frozenset({"token", "redirect_url"}) + ) == {"redirect_url": "", "token": ""} + body = response.json() + set_cookie = response.headers.get("set-cookie", "") + shape = { + "redirect_url_has_ui": "/ui/" in body.get("redirect_url", ""), + "redirect_url_has_login_success": "login=success" + in body.get("redirect_url", ""), + "token_in_body": bool(body.get("token")), + "token_cookie_set": "token=" in set_cookie, + } + assert shape == { + "redirect_url_has_ui": True, + "redirect_url_has_login_success": True, + "token_in_body": True, + "token_cookie_set": True, + } + + +def test_v2_login_authenticate_failure_500(client, monkeypatch): + """Error path: authenticate_user raising -> ProxyException -> 500 with structured error.""" + _install_login_mocks(monkeypatch, raise_on_auth=True) + response = client.post( + "/v2/login", + json={"username": "admin", "password": "wrong"}, + ) + assert response.status_code == 500 + body = response.json() + # Non-status assertion: response shape should carry an error + assert "error" in body or "detail" in body + assert isinstance(body, dict) + + +# --------------------------------------------------------------------------- +# POST /v3/login +# --------------------------------------------------------------------------- + + +def test_v3_login_without_control_plane_url_404(client, monkeypatch): + """Pin: /v3/login is gated on general_settings['control_plane_url'] — 404 when absent.""" + _install_login_mocks(monkeypatch) + # _install_login_mocks sets general_settings to {} — re-affirm + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.post( + "/v3/login", + json={"username": "admin", "password": "password"}, + ) + assert response.status_code == 404 + body = response.json() + # Detail carries the structured ProxyException error + detail = body.get("detail", {}) + if isinstance(detail, dict): + message = detail.get("error", {}) + if isinstance(message, dict): + message_str = message.get("message", "") + else: + message_str = str(message) + else: + message_str = str(detail) + assert "control_plane_url" in str(body) + + +def test_v3_login_success_returns_code(client, monkeypatch): + """Pin: /v3/login with control_plane_url returns {code, expires_in}.""" + from litellm.proxy import proxy_server as ps + + _install_login_mocks(monkeypatch) + monkeypatch.setattr( + ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} + ) + # Force the local (non-redis) cache path + monkeypatch.setattr(ps, "redis_usage_cache", None) + fake_cache = MagicMock() + fake_cache.async_set_cache = AsyncMock() + monkeypatch.setattr(ps, "user_api_key_cache", fake_cache) + + response = client.post( + "/v3/login", + json={"username": "admin", "password": "password"}, + ) + assert response.status_code == 200 + body = response.json() + # Strong assertion via normalize with extended volatile set ("code" is volatile) + assert normalize( + body, volatile=frozenset({"code", "expires_in"}) + ) == {"code": "", "expires_in": ""} + shape = { + "has_code": isinstance(body.get("code"), str) and len(body["code"]) > 0, + "expires_in_60": body.get("expires_in") == 60, + "cache_set_called": fake_cache.async_set_cache.await_count == 1, + } + assert shape == { + "has_code": True, + "expires_in_60": True, + "cache_set_called": True, + } + + +def test_v3_login_authenticate_failure_500(client, monkeypatch): + """Error path: with control_plane_url set, authenticate_user raises -> 500.""" + from litellm.proxy import proxy_server as ps + + _install_login_mocks(monkeypatch, raise_on_auth=True) + monkeypatch.setattr( + ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} + ) + + response = client.post( + "/v3/login", + json={"username": "admin", "password": "wrong"}, + ) + assert response.status_code == 500 + body = response.json() + assert isinstance(body, dict) + assert "error" in body or "detail" in body + + +# --------------------------------------------------------------------------- +# POST /v3/login/exchange +# --------------------------------------------------------------------------- + + +def test_v3_login_exchange_without_control_plane_url_404(client, monkeypatch): + """Pin: /v3/login/exchange gated on control_plane_url — 404 when absent.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.post("/v3/login/exchange", json={"code": "abc"}) + assert response.status_code == 404 + body = response.json() + assert "control_plane_url" in str(body) + assert isinstance(body, dict) + + +def test_v3_login_exchange_missing_code_400(client, monkeypatch): + """Error path: missing 'code' in body -> 400 with 'Missing' message.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr( + ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} + ) + + response = client.post("/v3/login/exchange", json={}) + assert response.status_code == 400 + body = response.json() + assert isinstance(body, dict) + assert "Missing" in str(body) or "code" in str(body) + + +def test_v3_login_exchange_invalid_code_401(client, monkeypatch): + """Error path: code that isn't in cache -> 401 'Invalid or expired'.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr( + ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} + ) + monkeypatch.setattr(ps, "redis_usage_cache", None) + fake_cache = MagicMock() + fake_cache.async_get_cache = AsyncMock(return_value=None) + fake_cache.async_delete_cache = AsyncMock() + monkeypatch.setattr(ps, "user_api_key_cache", fake_cache) + + response = client.post("/v3/login/exchange", json={"code": "nope"}) + assert response.status_code == 401 + body = response.json() + assert isinstance(body, dict) + assert "Invalid" in str(body) or "expired" in str(body) + + +def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatch): + """Pin: valid code -> JSON {token, redirect_url} + token cookie + cache deleted (single-use).""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr( + ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"} + ) + monkeypatch.setattr(ps, "redis_usage_cache", None) + + cached_payload = { + "token": "jwt-token-xyz", + "redirect_url": "https://litellm.example.invalid/ui/?login=success", + } + fake_cache = MagicMock() + fake_cache.async_get_cache = AsyncMock(return_value=cached_payload) + fake_cache.async_delete_cache = AsyncMock() + monkeypatch.setattr(ps, "user_api_key_cache", fake_cache) + + response = client.post("/v3/login/exchange", json={"code": "valid-code"}) + assert response.status_code == 200 + assert normalize( + response.json(), volatile=frozenset({"token", "redirect_url"}) + ) == {"token": "", "redirect_url": ""} + body = response.json() + set_cookie = response.headers.get("set-cookie", "") + shape = { + "token": body.get("token"), + "redirect_url": body.get("redirect_url"), + "token_cookie_set": "token=" in set_cookie, + "cache_deleted_once": fake_cache.async_delete_cache.await_count == 1, + } + assert shape == { + "token": "jwt-token-xyz", + "redirect_url": "https://litellm.example.invalid/ui/?login=success", + "token_cookie_set": True, + "cache_deleted_once": True, + } diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_misc.py b/tests/test_litellm/proxy/proxy_server/test_routes_misc.py index ad6b4016461..0c45e31afd2 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_misc.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_misc.py @@ -1 +1,230 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py misc routes (PR3). + +Routes covered: +- GET / +- GET /routes +- GET /adaptive_router/state +- GET /get_logo_url +- GET /get_image +- GET /get_favicon +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import normalize + + +# --------------------------------------------------------------------------- +# GET / +# --------------------------------------------------------------------------- + + +def test_home_returns_200_with_body(client, auth_as): + """GET / serves either the home string or the Swagger UI fallback — + both return 200 with a non-empty body. This pins the contract: root + always answers and never errors.""" + with auth_as(): + response = client.get("/") + shape = { + "status": response.status_code, + "has_body": len(response.content) > 0, + "has_content_type": bool(response.headers.get("content-type")), + } + assert shape == {"status": 200, "has_body": True, "has_content_type": True} + + +def test_home_invalid_method_405(client): + """GET / handler is GET-only; DELETE returns 405 (error path).""" + response = client.delete("/") + assert response.status_code == 405 + assert len(response.content) > 0 and response.headers.get("content-type") + + +# --------------------------------------------------------------------------- +# GET /routes +# --------------------------------------------------------------------------- + + +def test_get_routes_returns_routes_list(client, auth_as): + with auth_as(): + response = client.get("/routes") + assert response.status_code == 200 + body = response.json() + assert isinstance(body, dict) + assert "routes" in body + assert isinstance(body["routes"], list) + assert len(body["routes"]) > 0 + sample = body["routes"][0] + shape = { + "has_path": "path" in sample, + "has_methods": "methods" in sample, + "has_endpoint": "endpoint" in sample, + } + assert shape == { + "has_path": True, + "has_methods": True, + "has_endpoint": True, + } + + +def test_get_routes_invalid_method_405(client): + """POST against the GET-only /routes endpoint is rejected (error path).""" + response = client.post("/routes") + assert response.status_code == 405 + body = response.json() if response.headers.get("content-type", "").startswith( + "application/json" + ) else {} + assert isinstance(body, dict) + + +# --------------------------------------------------------------------------- +# GET /adaptive_router/state +# --------------------------------------------------------------------------- + + +def test_adaptive_router_state_returns_snapshots(client, auth_as, monkeypatch): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + fake_router = MagicMock() + snap = {"router_name": "ar-1", "queue_depth": 0, "posteriors": []} + bandit = MagicMock() + bandit.get_state_snapshot = AsyncMock(return_value=snap) + fake_router.adaptive_routers = {"ar-1": bandit} + monkeypatch.setattr(ps, "llm_router", fake_router) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/adaptive_router/state") + assert response.status_code == 200 + assert normalize(response.json()) == { + "routers": [ + {"router_name": "ar-1", "queue_depth": 0, "posteriors": []}, + ] + } + + +def test_adaptive_router_state_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/adaptive_router/state") + assert response.status_code == 403 + assert "error" in response.json().get("detail", {}) + + +def test_adaptive_router_state_not_configured_404(client, auth_as, monkeypatch): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + fake_router = MagicMock() + fake_router.adaptive_routers = {} + monkeypatch.setattr(ps, "llm_router", fake_router) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/adaptive_router/state") + assert response.status_code == 404 + assert "adaptive_router" in response.json().get("detail", {}).get("error", "") + + +# --------------------------------------------------------------------------- +# GET /get_logo_url +# --------------------------------------------------------------------------- + + +def test_get_logo_url_returns_http_url_when_set(client, monkeypatch): + monkeypatch.setenv("UI_LOGO_PATH", "https://example.invalid/logo.png") + response = client.get("/get_logo_url") + assert response.status_code == 200 + assert normalize(response.json()) == {"logo_url": "https://example.invalid/logo.png"} + + +def test_get_logo_url_blank_when_local_path(client, monkeypatch): + """Local filesystem paths must NOT be disclosed via this endpoint.""" + monkeypatch.setenv("UI_LOGO_PATH", "/var/lib/litellm/internal-secret-logo.png") + response = client.get("/get_logo_url") + assert response.status_code == 200 + assert normalize(response.json()) == {"logo_url": ""} + + +def test_get_logo_url_blank_when_unset(client, monkeypatch): + monkeypatch.delenv("UI_LOGO_PATH", raising=False) + response = client.get("/get_logo_url") + assert response.status_code == 200 + assert normalize(response.json()) == {"logo_url": ""} + + +def test_get_logo_url_invalid_scheme_blank(client, monkeypatch): + """file:// and other non-HTTP schemes are not disclosed (error/edge path).""" + monkeypatch.setenv("UI_LOGO_PATH", "file:///etc/passwd") + response = client.get("/get_logo_url") + assert response.status_code == 200 + assert normalize(response.json()) == {"logo_url": ""} + + +# --------------------------------------------------------------------------- +# GET /get_image +# --------------------------------------------------------------------------- + + +def test_get_image_returns_default_logo(client, monkeypatch): + monkeypatch.delenv("UI_LOGO_PATH", raising=False) + response = client.get("/get_image") + assert response.status_code == 200 + media_type = response.headers.get("content-type", "").split(";")[0] + shape = { + "status": response.status_code, + "media_type_image": media_type.startswith("image/"), + "has_body": len(response.content) > 0, + } + assert shape == {"status": 200, "media_type_image": True, "has_body": True} + + +def test_get_image_redirects_remote_url(client, monkeypatch): + """Remote logo URLs are served via redirect — the proxy never fetches them server-side.""" + monkeypatch.setenv("UI_LOGO_PATH", "https://example.invalid/logo.png") + response = client.get("/get_image", follow_redirects=False) + assert response.status_code in (302, 303, 307, 308) + assert response.headers.get("location") == "https://example.invalid/logo.png" + + +def test_get_image_invalid_local_path_falls_back(client, monkeypatch): + """Non-existent UI_LOGO_PATH (error path) falls back to default logo, still 200.""" + monkeypatch.setenv("UI_LOGO_PATH", "/nonexistent/path/to/logo.png") + response = client.get("/get_image") + assert response.status_code == 200 + shape = { + "status": response.status_code, + "media_type_image": response.headers.get("content-type", "").startswith( + "image/" + ), + "has_body": len(response.content) > 0, + } + assert shape == {"status": 200, "media_type_image": True, "has_body": True} + + +# --------------------------------------------------------------------------- +# GET /get_favicon +# --------------------------------------------------------------------------- + + +def test_get_favicon_returns_file(client): + response = client.get("/get_favicon") + assert response.status_code == 200 + shape = { + "status": response.status_code, + "has_body": len(response.content) > 0, + "content_type_set": bool(response.headers.get("content-type")), + } + assert shape == {"status": 200, "has_body": True, "content_type_set": True} + + +def test_get_favicon_invalid_custom_path_falls_back(client, monkeypatch): + """Bad UI_FAVICON_PATH (error/edge path) falls back to default — still 200.""" + monkeypatch.setenv("UI_FAVICON_PATH", "/nonexistent/favicon.ico") + response = client.get("/get_favicon") + assert response.status_code == 200 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py index ad6b4016461..16e410f1b1e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py @@ -1 +1,371 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py model cost map routes (PR3). + +Routes covered: +- POST /reload/model_cost_map +- POST /schedule/model_cost_map_reload +- DELETE /schedule/model_cost_map_reload +- GET /schedule/model_cost_map_reload/status +- GET /model/cost_map/source +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from .conftest import VOLATILE_KEYS, normalize + +# Some response bodies include a "timestamp" — extend the volatile set so +# dict-equality assertions remain stable. +_VOLATILE = VOLATILE_KEYS | frozenset({"timestamp"}) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _attach_litellm_config(mock_prisma): + """Attach a litellm_config table mock (not in conftest's _PRISMA_TABLES).""" + table = MagicMock() + table.find_unique = AsyncMock(return_value=None) + table.find_first = AsyncMock(return_value=None) + table.find_many = AsyncMock(return_value=[]) + table.upsert = AsyncMock() + table.create = AsyncMock() + table.update = AsyncMock() + table.delete = AsyncMock() + table.delete_many = AsyncMock() + mock_prisma.db.litellm_config = table + return table + + +# --------------------------------------------------------------------------- +# POST /reload/model_cost_map +# --------------------------------------------------------------------------- + + +def test_reload_model_cost_map_happy(client, auth_as, monkeypatch, mock_prisma): + """Admin can trigger a manual reload; handler returns model count + status.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _attach_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + fake_cost_map = {"gpt-4": {"input_cost": 0.03}, "gpt-3.5": {"input_cost": 0.002}} + monkeypatch.setattr( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map", + lambda url=None: fake_cost_map, + ) + monkeypatch.setattr("litellm.add_known_models", lambda model_cost_map=None: None) + monkeypatch.setattr("litellm.model_cost", {}, raising=False) + monkeypatch.setattr( + "litellm.proxy.proxy_server._invalidate_model_cost_lowercase_map", + lambda: None, + raising=False, + ) + + async def _fake_invalidate(name): + return None + + monkeypatch.setattr(ps, "invalidate_config_param", _fake_invalidate) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/reload/model_cost_map") + assert response.status_code == 200 + body = normalize(response.json(), volatile=_VOLATILE) + assert body == { + "message": "Price data reloaded successfully! 2 models updated.", + "status": "success", + "models_count": 2, + "timestamp": "", + } + assert table.upsert.await_count == 1 + + +def test_reload_model_cost_map_not_admin_forbidden(client, auth_as): + """Non-admin caller gets 403 with a role-specific detail.""" + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post("/reload/model_cost_map") + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +def test_reload_model_cost_map_no_db_500(client, auth_as, monkeypatch): + """Admin path but prisma_client is None — handler raises 500.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/reload/model_cost_map") + assert response.status_code == 500 + assert "Database connection not available" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# POST /schedule/model_cost_map_reload +# --------------------------------------------------------------------------- + + +def test_schedule_model_cost_map_reload_happy( + client, auth_as, monkeypatch, mock_prisma +): + """Admin schedules a reload — handler upserts config and echoes interval.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _attach_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + async def _fake_invalidate(name): + return None + + monkeypatch.setattr(ps, "invalidate_config_param", _fake_invalidate) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/schedule/model_cost_map_reload?hours=6") + assert response.status_code == 200 + body = normalize(response.json(), volatile=_VOLATILE) + assert body == { + "message": "Model cost map reload scheduled for every 6 hours", + "status": "success", + "interval_hours": 6, + "timestamp": "", + } + assert table.upsert.await_count == 1 + + +def test_schedule_model_cost_map_reload_invalid_hours( + client, auth_as, monkeypatch, mock_prisma +): + """hours <= 0 is rejected with 400.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + _attach_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/schedule/model_cost_map_reload?hours=0") + assert response.status_code == 400 + assert "Hours must be greater than 0" in response.json().get("detail", "") + + +def test_schedule_model_cost_map_reload_not_admin_forbidden(client, auth_as): + """Non-admin caller blocked with 403.""" + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.post("/schedule/model_cost_map_reload?hours=6") + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# DELETE /schedule/model_cost_map_reload +# --------------------------------------------------------------------------- + + +def test_cancel_model_cost_map_reload_happy(client, auth_as, monkeypatch, mock_prisma): + """Admin cancellation deletes config row and returns success body.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _attach_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + async def _fake_invalidate(name): + return None + + monkeypatch.setattr(ps, "invalidate_config_param", _fake_invalidate) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.delete("/schedule/model_cost_map_reload") + assert response.status_code == 200 + body = normalize(response.json(), volatile=_VOLATILE) + assert body == { + "message": "Model cost map reload schedule cancelled", + "status": "success", + "timestamp": "", + } + assert table.delete.await_count == 1 + + +def test_cancel_model_cost_map_reload_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.delete("/schedule/model_cost_map_reload") + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +def test_cancel_model_cost_map_reload_no_db_500(client, auth_as, monkeypatch): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.delete("/schedule/model_cost_map_reload") + assert response.status_code == 500 + assert "Database connection not available" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# GET /schedule/model_cost_map_reload/status +# --------------------------------------------------------------------------- + + +def test_get_model_cost_map_reload_status_no_db_not_scheduled( + client, auth_as, monkeypatch +): + """No prisma client → returns the not-scheduled shape (4 keys, all-null).""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + monkeypatch.setattr(ps, "prisma_client", None) + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/model_cost_map_reload/status") + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None, + } + + +def test_get_model_cost_map_reload_status_scheduled( + client, auth_as, monkeypatch, mock_prisma +): + """A valid config row → scheduled=True and the interval is echoed.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _attach_litellm_config(mock_prisma) + config_row = MagicMock() + config_row.param_value = {"interval_hours": 12, "force_reload": False} + table.find_unique = AsyncMock(return_value=config_row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "last_model_cost_map_reload", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/model_cost_map_reload/status") + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": True, + "interval_hours": 12, + "last_run": None, + "next_run": None, + } + + +def test_get_model_cost_map_reload_status_no_config_not_scheduled( + client, auth_as, monkeypatch, mock_prisma +): + """Config row exists but interval_hours=None → not scheduled.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _attach_litellm_config(mock_prisma) + config_row = MagicMock() + config_row.param_value = {"interval_hours": None, "force_reload": True} + table.find_unique = AsyncMock(return_value=config_row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "last_model_cost_map_reload", None) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/schedule/model_cost_map_reload/status") + assert response.status_code == 200 + assert normalize(response.json()) == { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None, + } + + +def test_get_model_cost_map_reload_status_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/schedule/model_cost_map_reload/status") + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") + + +# --------------------------------------------------------------------------- +# GET /model/cost_map/source +# --------------------------------------------------------------------------- + + +def test_get_model_cost_map_source_happy(client, auth_as, monkeypatch): + """Admin gets the source-info dict, augmented with the current model_count.""" + from litellm.proxy._types import LitellmUserRoles + + fake_info = { + "source": "remote", + "url": "https://example.invalid/cost_map.json", + "is_env_forced": False, + "fallback_reason": None, + } + monkeypatch.setattr( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map_source_info", + lambda: fake_info, + ) + monkeypatch.setattr("litellm.model_cost", {"a": 1, "b": 2, "c": 3}, raising=False) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/model/cost_map/source") + assert response.status_code == 200 + assert normalize(response.json()) == { + "source": "remote", + "url": "https://example.invalid/cost_map.json", + "is_env_forced": False, + "fallback_reason": None, + "model_count": 3, + } + + +def test_get_model_cost_map_source_admin_view_only_allowed( + client, auth_as, monkeypatch +): + """PROXY_ADMIN_VIEW_ONLY can read source info — pins the read-only ACL.""" + from litellm.proxy._types import LitellmUserRoles + + fake_info = { + "source": "local", + "url": None, + "is_env_forced": True, + "fallback_reason": None, + } + monkeypatch.setattr( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map_source_info", + lambda: fake_info, + ) + monkeypatch.setattr("litellm.model_cost", {"a": 1}, raising=False) + + with auth_as(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY): + response = client.get("/model/cost_map/source") + assert response.status_code == 200 + assert normalize(response.json()) == { + "source": "local", + "url": None, + "is_env_forced": True, + "fallback_reason": None, + "model_count": 1, + } + + +def test_get_model_cost_map_source_not_admin_forbidden(client, auth_as): + from litellm.proxy._types import LitellmUserRoles + + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/model/cost_map/source") + assert response.status_code == 403 + assert "Admin role required" in response.json().get("detail", "") diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index ad6b4016461..98259824378 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -1 +1,149 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` model-info routes. + +Pins (PR2): + - GET /v2/model/info + - GET /v1/model/info + - GET /model/info + - GET /model_group/info +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +# --------------------------------------------------------------------------- +# GET /v2/model/info +# --------------------------------------------------------------------------- + + +@pytest.fixture +def empty_router(monkeypatch): + router = MagicMock() + router.model_list = [] + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", []) + yield router + + +@pytest.fixture +def null_router(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(proxy_server, "llm_model_list", None) + yield + + +def test_v2_model_info_empty_router_happy_path(client, auth_as, empty_router): + """Pins ``GET /v2/model/info`` (empty router branch returns deterministic shape).""" + with auth_as(): + response = client.get("/v2/model/info") + assert response.status_code == 200 + assert normalize(response.json()) == { + "data": [], + "total_count": 0, + "current_page": 1, + "total_pages": 0, + "size": 50, + } + + +def test_v2_model_info_invalid_page_returns_422(client, auth_as, empty_router): + """Pins ``GET /v2/model/info`` (error: invalid page parameter).""" + with auth_as(): + response = client.get("/v2/model/info", params={"page": 0}) + assert response.status_code == 422 + assert "detail" in response.json() + + +# --------------------------------------------------------------------------- +# GET /v1/model/info, GET /model/info +# --------------------------------------------------------------------------- + + +@pytest.fixture +def configured_router(monkeypatch): + deployment = MagicMock() + deployment.model_dump = MagicMock( + return_value={ + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "abc", "db_model": False}, + } + ) + router = MagicMock() + router.get_deployment = MagicMock(return_value=deployment) + router.get_model_names = MagicMock(return_value=["gpt-4"]) + router.get_model_access_groups = MagicMock(return_value={}) + router.get_model_list = MagicMock(return_value=[]) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": "gpt-4"}]) + monkeypatch.setattr(proxy_server, "user_model", None) + monkeypatch.setattr(proxy_server, "_get_proxy_model_info", lambda model: model) + yield router + + +@pytest.mark.parametrize("path", ["/v1/model/info", "/model/info"]) +def test_v1_model_info_specific_id_happy(client, auth_as, configured_router, path): + """Pins ``GET /v1/model/info`` and ``GET /model/info`` (happy: specific id). + + Includes ``litellm_model_id`` so the early-return branch produces a + deterministic ``{"data": []}`` body without touching + the full model-info enrichment pipeline. + """ + with auth_as(): + response = client.get(path, params={"litellm_model_id": "abc"}) + assert response.status_code == 200 + body = normalize(response.json()) + assert body == { + "data": [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "", "db_model": False}, + } + ] + } + + +@pytest.mark.parametrize("path", ["/v1/model/info", "/model/info"]) +def test_v1_model_info_no_model_list_error(client, auth_as, null_router, path): + """Pins ``GET /v1/model/info`` and ``GET /model/info`` (error: no model list).""" + with auth_as(): + response = client.get(path) + assert response.status_code == 500 + assert "LLM Model List not loaded" in response.text + + +# --------------------------------------------------------------------------- +# GET /model_group/info +# --------------------------------------------------------------------------- + + +def test_model_group_info_no_models_happy(client, auth_as, null_router): + """Pins ``GET /model_group/info`` (happy: empty list when no models).""" + with auth_as(): + response = client.get("/model_group/info") + assert response.status_code == 200 + summary = { + "status_code": response.status_code, + "body": normalize(response.json()), + "object_kind": "model_group_info", + } + assert summary == { + "status_code": 200, + "body": {"data": []}, + "object_kind": "model_group_info", + } + + +def test_model_group_info_invalid_method(client, auth_as, null_router): + """Pins ``GET /model_group/info`` (error: method not allowed).""" + with auth_as(): + response = client.post("/model_group/info", json={}) + assert response.status_code == 405 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py index ad6b4016461..246e2cbba54 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py @@ -1 +1,228 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` model-metrics routes. + +Pins (PR2): + - GET /model/streaming_metrics + - GET /model/metrics + - GET /model/metrics/slow_responses + - GET /model/metrics/exceptions + - GET /model/settings + - GET /alerting/settings +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles + +from .conftest import normalize # type: ignore[import-not-found] + +# --------------------------------------------------------------------------- +# Shared fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture +def prisma_with_query_raw(monkeypatch): + pc = MagicMock() + pc.db.query_raw = AsyncMock(return_value=[]) + monkeypatch.setattr(proxy_server, "prisma_client", pc) + return pc + + +@pytest.fixture +def no_prisma(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", None) + yield + + +# --------------------------------------------------------------------------- +# GET /model/streaming_metrics +# --------------------------------------------------------------------------- + + +def test_model_streaming_metrics_happy(client, auth_as, prisma_with_query_raw): + """Pins ``GET /model/streaming_metrics`` (happy: empty data list). + + Drives the deterministic branch where ``query_raw`` returns an empty + list; the handler should return the empty payload unchanged so the + pin can rely on the exact response shape. + """ + with auth_as(): + response = client.get( + "/model/streaming_metrics", params={"_selected_model_group": "gpt-4"} + ) + assert response.status_code == 200 + assert normalize(response.json()) == {"data": [], "all_api_bases": []} + + +def test_model_streaming_metrics_no_prisma_error(client, auth_as, no_prisma): + """Pins ``GET /model/streaming_metrics`` (error: prisma not initialized).""" + with auth_as(): + response = client.get("/model/streaming_metrics") + assert response.status_code == 500 + assert response.content + + +# --------------------------------------------------------------------------- +# GET /model/metrics +# --------------------------------------------------------------------------- + + +def test_model_metrics_happy(client, auth_as, prisma_with_query_raw): + """Pins ``GET /model/metrics`` (happy: empty result).""" + with auth_as(): + response = client.get("/model/metrics") + assert response.status_code == 200 + assert normalize(response.json()) == {"data": [], "all_api_bases": []} + + +def test_model_metrics_no_prisma_error(client, auth_as, no_prisma): + """Pins ``GET /model/metrics`` (error: prisma not initialized).""" + with auth_as(): + response = client.get("/model/metrics") + assert response.status_code == 500 + assert response.content + + +# --------------------------------------------------------------------------- +# GET /model/metrics/slow_responses +# --------------------------------------------------------------------------- + + +def test_model_metrics_slow_responses_happy( + client, auth_as, prisma_with_query_raw, monkeypatch +): + """Pins ``GET /model/metrics/slow_responses`` (happy: empty list).""" + logging_obj = MagicMock() + logging_obj.slack_alerting_instance.alerting_threshold = 30 + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging_obj) + with auth_as(): + response = client.get("/model/metrics/slow_responses") + assert response.status_code == 200 + assert normalize(response.json()) == [] + + +def test_model_metrics_slow_responses_no_prisma(client, auth_as, no_prisma): + """Pins ``GET /model/metrics/slow_responses`` (error: prisma not initialized).""" + with auth_as(): + response = client.get("/model/metrics/slow_responses") + assert response.status_code == 500 + assert response.content + + +# --------------------------------------------------------------------------- +# GET /model/metrics/exceptions +# --------------------------------------------------------------------------- + + +def test_model_metrics_exceptions_happy(client, auth_as, prisma_with_query_raw): + """Pins ``GET /model/metrics/exceptions`` (happy: empty).""" + with auth_as(): + response = client.get("/model/metrics/exceptions") + assert response.status_code == 200 + assert normalize(response.json()) == {"data": [], "exception_types": []} + + +def test_model_metrics_exceptions_no_prisma(client, auth_as, no_prisma): + """Pins ``GET /model/metrics/exceptions`` (error: prisma not initialized).""" + with auth_as(): + response = client.get("/model/metrics/exceptions") + assert response.status_code == 500 + assert response.content + + +# --------------------------------------------------------------------------- +# GET /model/settings +# --------------------------------------------------------------------------- + + +def test_model_settings_happy(client, auth_as, monkeypatch): + """Pins ``GET /model/settings`` (happy).""" + monkeypatch.setattr(litellm, "provider_list", ["openai"]) + monkeypatch.setattr( + litellm, + "get_provider_fields", + lambda custom_llm_provider: [], + ) + with auth_as(): + response = client.get("/model/settings") + assert response.status_code == 200 + body = response.json() + assert body == [{"name": "openai", "fields": []}] + summary = { + "status_code": response.status_code, + "first_entry_name": body[0]["name"], + "body_length": len(body), + } + assert summary == { + "status_code": 200, + "first_entry_name": "openai", + "body_length": 1, + } + + +def test_model_settings_method_not_allowed(client, auth_as): + """Pins ``GET /model/settings`` (error: wrong method).""" + with auth_as(): + response = client.post("/model/settings", json={}) + assert response.status_code == 405 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# GET /alerting/settings +# --------------------------------------------------------------------------- + + +def test_alerting_settings_no_db_error(client, auth_as, no_prisma): + """Pins ``GET /alerting/settings`` (error: db not connected).""" + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + assert response.status_code == 400 + assert "error" in response.text or "detail" in response.text + + +def test_alerting_settings_non_admin_error(client, auth_as, monkeypatch): + """Pins ``GET /alerting/settings`` (error: non-admin forbidden).""" + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + with auth_as(LitellmUserRoles.INTERNAL_USER): + response = client.get("/alerting/settings") + assert response.status_code == 400 + assert "internal_user" in response.text.lower() or "error" in response.text + + +def test_alerting_settings_happy(client, auth_as, monkeypatch): + """Pins ``GET /alerting/settings`` (happy: returns list of ConfigList entries).""" + pc = MagicMock() + pc.db.litellm_config.find_first = AsyncMock(return_value=None) + monkeypatch.setattr(proxy_server, "prisma_client", pc) + + logging_obj = MagicMock() + args_model = MagicMock() + args_model.model_dump = MagicMock(return_value={}) + logging_obj.slack_alerting_instance.alerting_args = args_model + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging_obj) + monkeypatch.setattr(proxy_server, "general_settings", {}) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/alerting/settings") + assert response.status_code == 200 + body = response.json() + assert body[0]["field_name"] == "slack_alerting" + summary = { + "status_code": response.status_code, + "first_field_name": body[0]["field_name"], + "first_field_value": body[0]["field_value"], + "first_field_type": body[0]["field_type"], + } + assert summary == { + "status_code": 200, + "first_field_name": "slack_alerting", + "first_field_value": False, + "first_field_type": "Boolean", + } diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_models.py b/tests/test_litellm/proxy/proxy_server/test_routes_models.py index ad6b4016461..381835fbc14 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_models.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_models.py @@ -1 +1,132 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` model routes. + +Pins (PR2): + - GET /v1/models + - GET /models + - GET /v1/models/{model_id} + - GET /models/{model_id} +""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +import litellm +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + + +def _stub_model_info_response( + model_id: str = "gpt-4", provider: str = "openai" +) -> dict: + return { + "id": model_id, + "object": "model", + "created": 0, + "owned_by": provider, + } + + +@pytest.fixture +def patched_models(monkeypatch): + """Stub router + utility helpers used by the /models routes.""" + from litellm.proxy import utils as proxy_utils + + router = MagicMock() + router.get_fully_blocked_model_names = MagicMock(return_value=set()) + router.get_model_names = MagicMock(return_value=["gpt-4", "claude-sonnet"]) + router.get_model_access_groups = MagicMock(return_value={}) + + deployment = MagicMock() + deployment.litellm_params.model = "gpt-4" + router.get_deployment_by_model_group_name = MagicMock(return_value=deployment) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + + async def _fake_get_available_models_for_user(**kwargs): + return ["gpt-4", "claude-sonnet"] + + monkeypatch.setattr( + proxy_utils, + "get_available_models_for_user", + _fake_get_available_models_for_user, + ) + + def _fake_create_model_info_response(model_id, provider="openai", **kwargs): + return _stub_model_info_response(model_id=model_id, provider=provider) + + monkeypatch.setattr( + proxy_utils, "create_model_info_response", _fake_create_model_info_response + ) + + monkeypatch.setattr(proxy_utils, "validate_model_access", lambda **kwargs: None) + + monkeypatch.setattr( + litellm, + "get_llm_provider", + lambda model: (model, "openai", None, None), + ) + + return router + + +@pytest.mark.parametrize("path", ["/v1/models", "/models"]) +def test_get_models_happy_path(client, auth_as, patched_models, path): + """Pins: ``GET /v1/models``, ``GET /models``.""" + with auth_as(): + response = client.get(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "data": [ + { + "id": "", + "object": "model", + "created": "", + "owned_by": "openai", + }, + { + "id": "", + "object": "model", + "created": "", + "owned_by": "openai", + }, + ], + "object": "list", + } + + +@pytest.mark.parametrize("path", ["/v1/models", "/models"]) +def test_get_models_invalid_scope_returns_400(client, auth_as, patched_models, path): + """Pins: ``GET /v1/models``, ``GET /models`` (error path: invalid scope).""" + with auth_as(): + response = client.get(path, params={"scope": "not-a-real-scope"}) + assert response.status_code == 400 + assert "Invalid scope parameter" in str(response.json()) + + +@pytest.mark.parametrize("path", ["/v1/models/gpt-4", "/models/gpt-4"]) +def test_get_model_by_id_happy_path(client, auth_as, patched_models, path): + """Pins: ``GET /v1/models/{model_id}``, ``GET /models/{model_id}``.""" + with auth_as(): + response = client.get(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "model", + "created": "", + "owned_by": "openai", + } + + +@pytest.mark.parametrize("path", ["/v1/models/missing", "/models/missing"]) +def test_get_model_by_id_not_found(client, auth_as, patched_models, path): + """Pins: ``GET /v1/models/{model_id}``, ``GET /models/{model_id}`` (error: 404).""" + patched_models.get_deployment_by_model_group_name = MagicMock(return_value=None) + with auth_as(): + response = client.get(path) + assert response.status_code == 404 + assert "not found" in response.text.lower() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_moderations.py b/tests/test_litellm/proxy/proxy_server/test_routes_moderations.py index ad6b4016461..4553a5e7cf4 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_moderations.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_moderations.py @@ -1 +1,111 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` moderations routes. + +Pins (PR2): + - POST /v1/moderations + - POST /moderations +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +HAPPY_RESPONSE = { + "id": "modr-test", + "model": "text-moderation-stable", + "results": [ + { + "flagged": False, + "categories": {"violence": False}, + "category_scores": {"violence": 0.01}, + } + ], +} + + +@pytest.fixture +def patched_moderation(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + pre_call_hook=AsyncMock(side_effect=lambda **kw: kw["data"]), + post_call_failure_hook=AsyncMock(), + update_request_status=AsyncMock(), + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + + async def _fake_llm_call(): + return dict(HAPPY_RESPONSE) + + async def _fake_route_request(*args, **kwargs): + return _fake_llm_call() + + monkeypatch.setattr(proxy_server, "route_request", _fake_route_request) + yield + + +@pytest.fixture +def moderation_pipeline_raises(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + pre_call_hook=AsyncMock(side_effect=lambda **kw: kw["data"]), + post_call_failure_hook=AsyncMock(), + update_request_status=AsyncMock(), + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + + async def _raise(*args, **kwargs): + raise ValueError("boom") + + monkeypatch.setattr(proxy_server, "route_request", _raise) + yield + + +@pytest.mark.parametrize("path", ["/v1/moderations", "/moderations"]) +def test_moderation_happy_path(client, auth_as, patched_moderation, path): + """Pins ``POST /v1/moderations`` and ``POST /moderations`` (happy).""" + payload = {"model": "text-moderation-stable", "input": "Sample text"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "model": "text-moderation-stable", + "results": [ + { + "flagged": False, + "categories": {"violence": False}, + "category_scores": {"violence": 0.01}, + } + ], + } + + +@pytest.mark.parametrize("path", ["/v1/moderations", "/moderations"]) +def test_moderation_error(client, auth_as, moderation_pipeline_raises, path): + """Pins ``POST /v1/moderations`` and ``POST /moderations`` (error).""" + payload = {"model": "text-moderation-stable", "input": "Sample text"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 500 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py index ad6b4016461..35ae9a3568e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py @@ -1 +1,350 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Pin tests for proxy_server.py onboarding routes (PR3). + +Routes covered: +- GET /onboarding/get_token +- POST /onboarding/claim_token +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import jwt +import pytest + +from .conftest import normalize + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_invite( + invite_id: str = "inv-123", + user_id: str = "user-abc", + expires_at: datetime | None = None, + is_accepted: bool = False, + accepted_at=None, +): + """Build a fake invitation object with the attributes the handler reads.""" + if expires_at is None: + expires_at = datetime.now(timezone.utc) + timedelta(days=1) + return SimpleNamespace( + id=invite_id, + user_id=user_id, + expires_at=expires_at, + is_accepted=is_accepted, + accepted_at=accepted_at, + ) + + +def _make_user_obj( + user_id: str = "user-abc", + user_email: str = "alice@example.com", + user_role: str = "internal_user", +): + return SimpleNamespace( + user_id=user_id, + user_email=user_email, + user_role=user_role, + password=None, + ) + + +def _install_tx_context(mock_prisma): + """Wire ``async with prisma_client.db.tx() as tx`` to return ``mock_prisma.db``. + + The handler runs the update inside a transaction; have ``tx`` yield a + namespace that exposes the same tables as the outer client so its + ``update_many`` / ``update`` calls hit our mocks. + """ + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=mock_prisma.db) + tx_cm.__aexit__ = AsyncMock(return_value=None) + mock_prisma.db.tx = MagicMock(return_value=tx_cm) + + +# --------------------------------------------------------------------------- +# GET /onboarding/get_token +# --------------------------------------------------------------------------- + + +def test_onboarding_get_token_happy(client, monkeypatch, mock_prisma): + """Valid invite link → returns dict with login_url, token, user_email.""" + from litellm.proxy import proxy_server as ps + + invite = _make_invite() + user_obj = _make_user_obj() + mock_prisma.db.litellm_invitationlink.find_unique.return_value = invite + mock_prisma.db.litellm_usertable.find_unique.return_value = user_obj + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + monkeypatch.setattr(ps, "premium_user", False) + + response = client.get("/onboarding/get_token", params={"invite_link": "inv-123"}) + assert response.status_code == 200 + body = response.json() + assert set(body.keys()) == {"login_url", "token", "user_email"} + assert body["user_email"] == "alice@example.com" + assert "ui/onboarding" in body["login_url"] + assert "token=" in body["login_url"] + # The JWT in body["token"] must decode with the master_key. + decoded = jwt.decode(body["token"], "sk-master-test", algorithms=["HS256"]) + assert normalize( + { + "user_id": decoded["user_id"], + "user_email": decoded["user_email"], + "login_method": decoded["login_method"], + "premium_user": decoded["premium_user"], + } + ) == { + "user_id": "user-abc", + "user_email": "alice@example.com", + "login_method": "username_password", + "premium_user": False, + } + + +def test_onboarding_get_token_master_key_missing_500(client, monkeypatch, mock_prisma): + """No master_key configured → 500 with the master_key error payload.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", None) + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.get("/onboarding/get_token", params={"invite_link": "inv-123"}) + assert response.status_code == 500 + body = response.json() + # ProxyException serializes to {"error": {"message": ..., "type": ..., "param": ..., "code": ...}} + err_blob = body.get("error", body) + assert "Master Key not set" in str(err_blob) + + +def test_onboarding_get_token_invalid_invite_link_401( + client, monkeypatch, mock_prisma +): + """Unknown invite link → 401 with the not-in-db error message.""" + from litellm.proxy import proxy_server as ps + + mock_prisma.db.litellm_invitationlink.find_unique.return_value = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.get( + "/onboarding/get_token", params={"invite_link": "does-not-exist"} + ) + assert response.status_code == 401 + assert response.json() == { + "detail": {"error": "Invitation link does not exist in db."} + } + + +def test_onboarding_get_token_expired_invite_401(client, monkeypatch, mock_prisma): + """Invite whose expires_at is in the past → 401 expired.""" + from litellm.proxy import proxy_server as ps + + expired_invite = _make_invite( + expires_at=datetime.now(timezone.utc) - timedelta(days=2) + ) + mock_prisma.db.litellm_invitationlink.find_unique.return_value = expired_invite + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.get("/onboarding/get_token", params={"invite_link": "inv-123"}) + assert response.status_code == 401 + assert response.json().get("detail", {}).get("error") == "Invitation link has expired." + + +def test_onboarding_get_token_missing_query_param_422(client, monkeypatch, mock_prisma): + """No ``invite_link`` query param → FastAPI 422 with a non-empty detail array.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.get("/onboarding/get_token") + assert response.status_code == 422 + body = response.json() + assert isinstance(body.get("detail"), list) + assert len(body["detail"]) >= 1 + + +# --------------------------------------------------------------------------- +# POST /onboarding/claim_token +# --------------------------------------------------------------------------- + + +def _make_onboarding_jwt( + master_key: str, + invitation_link: str = "inv-123", + user_id: str = "user-abc", + token_type: str = "litellm_onboarding", +) -> str: + return jwt.encode( + { + "token_type": token_type, + "invitation_link": invitation_link, + "user_id": user_id, + "exp": datetime.now(timezone.utc) + timedelta(minutes=15), + }, + master_key, + algorithm="HS256", + ) + + +def test_claim_onboarding_link_happy(client, monkeypatch, mock_prisma): + """Valid claim → returns login_url, token, user_email, user.""" + from litellm.proxy import proxy_server as ps + + invite = _make_invite() + user_obj = _make_user_obj() + mock_prisma.db.litellm_invitationlink.find_unique.return_value = invite + mock_prisma.db.litellm_invitationlink.update_many.return_value = 1 + mock_prisma.db.litellm_invitationlink.update.return_value = invite + mock_prisma.db.litellm_usertable.update.return_value = user_obj + _install_tx_context(mock_prisma) + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + monkeypatch.setattr(ps, "premium_user", False) + + # Avoid hitting generate_key_helper_fn (touches DB / many globals); patch + # the helper directly so we focus on the route's own behavior. + async def _fake_session_token(user_obj): + return "session-jwt-token" + + monkeypatch.setattr( + ps, "_generate_onboarding_ui_session_token", _fake_session_token + ) + + onboarding_jwt = _make_onboarding_jwt("sk-master-test") + response = client.post( + "/onboarding/claim_token", + json={ + "invitation_link": "inv-123", + "user_id": "user-abc", + "password": "hunter2", + }, + headers={"Authorization": f"Bearer {onboarding_jwt}"}, + ) + assert response.status_code == 200 + body = response.json() + assert set(body.keys()) == {"login_url", "token", "user_email", "user"} + assert body["token"] == "session-jwt-token" + assert body["user_email"] == "alice@example.com" + assert body["login_url"].endswith("/ui/?login=success") + + +def test_claim_onboarding_link_invalid_invite_401(client, monkeypatch, mock_prisma): + """Unknown invite link → 401 with not-in-db error.""" + from litellm.proxy import proxy_server as ps + + mock_prisma.db.litellm_invitationlink.find_unique.return_value = None + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.post( + "/onboarding/claim_token", + json={ + "invitation_link": "missing", + "user_id": "user-abc", + "password": "hunter2", + }, + headers={"Authorization": "Bearer irrelevant"}, + ) + assert response.status_code == 401 + assert response.json() == { + "detail": {"error": "Invitation link does not exist in db."} + } + + +def test_claim_onboarding_link_user_id_mismatch_401( + client, monkeypatch, mock_prisma +): + """Invitation belongs to a different user_id → 401 with mismatch error.""" + from litellm.proxy import proxy_server as ps + + invite = _make_invite(user_id="user-real-owner") + mock_prisma.db.litellm_invitationlink.find_unique.return_value = invite + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + response = client.post( + "/onboarding/claim_token", + json={ + "invitation_link": "inv-123", + "user_id": "user-attacker", + "password": "hunter2", + }, + headers={"Authorization": "Bearer irrelevant"}, + ) + assert response.status_code == 401 + err = response.json().get("detail", {}).get("error", "") + assert "Invalid invitation link" in err + assert "user-attacker" in err + + +def test_claim_onboarding_link_missing_field_422(client, monkeypatch, mock_prisma): + """Missing required body field → FastAPI 422 with detail listing the missing field.""" + from litellm.proxy import proxy_server as ps + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + # Missing "password" + response = client.post( + "/onboarding/claim_token", + json={"invitation_link": "inv-123", "user_id": "user-abc"}, + ) + assert response.status_code == 422 + body = response.json() + assert isinstance(body.get("detail"), list) + # The missing field should be referenced in the validation error. + assert any("password" in str(item) for item in body["detail"]) + + +def test_claim_onboarding_link_bad_onboarding_jwt_401( + client, monkeypatch, mock_prisma +): + """Onboarding JWT decodes but token_type / invitation_link don't match → 401.""" + from litellm.proxy import proxy_server as ps + + invite = _make_invite() + mock_prisma.db.litellm_invitationlink.find_unique.return_value = invite + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps, "master_key", "sk-master-test") + monkeypatch.setattr(ps, "general_settings", {}) + + # Wrong token_type — handler rejects. + bogus_jwt = _make_onboarding_jwt( + "sk-master-test", + token_type="not_onboarding", + ) + response = client.post( + "/onboarding/claim_token", + json={ + "invitation_link": "inv-123", + "user_id": "user-abc", + "password": "hunter2", + }, + headers={"Authorization": f"Bearer {bogus_jwt}"}, + ) + assert response.status_code == 401 + assert ( + response.json().get("detail", {}).get("error") + == "Invalid onboarding session for invitation link." + ) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_queue.py b/tests/test_litellm/proxy/proxy_server/test_routes_queue.py index ad6b4016461..27cc6300711 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_queue.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_queue.py @@ -1 +1,91 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` queue routes. + +Pins (PR2): + - POST /queue/chat/completions +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +HAPPY_RESPONSE = { + "id": "chatcmpl-queue", + "object": "chat.completion", + "created": 0, + "model": "gpt-4", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "queued reply"}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + "priority": 0, +} + + +@pytest.fixture +def patched_queue(monkeypatch): + router = MagicMock() + router.schedule_acompletion = AsyncMock(return_value=dict(HAPPY_RESPONSE)) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock(post_call_failure_hook=AsyncMock()), + ) + return router + + +@pytest.fixture +def queue_no_router(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock(post_call_failure_hook=AsyncMock()), + ) + yield + + +def test_queue_chat_completions_happy(client, auth_as, patched_queue): + """Pins ``POST /queue/chat/completions`` (happy).""" + payload = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "priority": 0, + } + with auth_as(): + response = client.post("/queue/chat/completions", json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "chat.completion", + "created": "", + "model": "gpt-4", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "queued reply"}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + "priority": 0, + } + + +def test_queue_chat_completions_no_router_error(client, auth_as, queue_no_router): + """Pins ``POST /queue/chat/completions`` (error: no llm_router).""" + payload = {"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]} + with auth_as(): + response = client.post("/queue/chat/completions", json=payload) + assert response.status_code == 500 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_threads.py b/tests/test_litellm/proxy/proxy_server/test_routes_threads.py index ad6b4016461..493315f041d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_threads.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_threads.py @@ -1 +1,274 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` threads routes. + +Pins (PR2): + - POST /v1/threads + - POST /threads + - GET /v1/threads/{thread_id} + - GET /threads/{thread_id} + - POST /v1/threads/{thread_id}/messages + - POST /threads/{thread_id}/messages + - GET /v1/threads/{thread_id}/messages + - GET /threads/{thread_id}/messages + - POST /v1/threads/{thread_id}/runs + - POST /threads/{thread_id}/runs +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +CREATE_THREAD = {"id": "thr_1", "object": "thread", "created_at": 0, "metadata": {}} +GET_THREAD = { + "id": "thr_1", + "object": "thread", + "created_at": 0, + "tool_resources": {}, +} +ADD_MESSAGE = { + "id": "msg_1", + "object": "thread.message", + "thread_id": "thr_1", + "role": "user", + "content": [], +} +GET_MESSAGES = { + "object": "list", + "data": [ + { + "id": "msg_1", + "object": "thread.message", + "thread_id": "thr_1", + "role": "user", + "content": [], + } + ], + "first_id": "msg_1", + "last_id": "msg_1", + "has_more": False, +} +RUN_THREAD = { + "id": "run_1", + "object": "thread.run", + "thread_id": "thr_1", + "assistant_id": "asst_1", + "status": "queued", +} + + +@pytest.fixture +def patched_threads(monkeypatch): + router = MagicMock() + router.acreate_thread = AsyncMock(return_value=dict(CREATE_THREAD)) + router.aget_thread = AsyncMock(return_value=dict(GET_THREAD)) + router.a_add_message = AsyncMock(return_value=dict(ADD_MESSAGE)) + router.aget_messages = AsyncMock(return_value=dict(GET_MESSAGES)) + router.arun_thread = AsyncMock(return_value=dict(RUN_THREAD)) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + post_call_failure_hook=AsyncMock(), update_request_status=AsyncMock() + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + return router + + +@pytest.fixture +def no_router(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr( + proxy_server, + "proxy_logging_obj", + MagicMock( + post_call_failure_hook=AsyncMock(), update_request_status=AsyncMock() + ), + ) + + async def _add_data(data, **kwargs): + return data + + monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data) + yield + + +# --------------------------------------------------------------------------- +# POST /v1/threads, POST /threads +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("path", ["/v1/threads", "/threads"]) +def test_create_thread_happy(client, auth_as, patched_threads, path): + """Pins ``POST /v1/threads`` and ``POST /threads``.""" + with auth_as(): + response = client.post(path, json={}) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "thread", + "created_at": "", + "metadata": {}, + } + + +@pytest.mark.parametrize("path", ["/v1/threads", "/threads"]) +def test_create_thread_error(client, auth_as, no_router, path): + """Pins ``POST /v1/threads`` / ``POST /threads`` (error: no llm_router).""" + with auth_as(): + response = client.post(path, json={}) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# GET /v1/threads/{thread_id}, GET /threads/{thread_id} +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("path", ["/v1/threads/thr_1", "/threads/thr_1"]) +def test_get_thread_happy(client, auth_as, patched_threads, path): + """Pins ``GET /v1/threads/{thread_id}`` and ``GET /threads/{thread_id}``.""" + with auth_as(): + response = client.get(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "thread", + "created_at": "", + "tool_resources": {}, + } + + +@pytest.mark.parametrize("path", ["/v1/threads/thr_1", "/threads/thr_1"]) +def test_get_thread_error(client, auth_as, no_router, path): + """Pins ``GET /v1/threads/{thread_id}`` / ``GET /threads/{thread_id}`` (error).""" + with auth_as(): + response = client.get(path) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# POST /v1/threads/{thread_id}/messages, POST /threads/{thread_id}/messages +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/messages", "/threads/thr_1/messages"], +) +def test_add_message_happy(client, auth_as, patched_threads, path): + """Pins ``POST /v1/threads/{thread_id}/messages`` and ``POST /threads/{thread_id}/messages``.""" + payload = {"role": "user", "content": "hi"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "thread.message", + "thread_id": "thr_1", + "role": "user", + "content": [], + } + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/messages", "/threads/thr_1/messages"], +) +def test_add_message_error(client, auth_as, no_router, path): + """Pins ``POST /v1/threads/{thread_id}/messages`` / ``POST /threads/{thread_id}/messages`` (error).""" + with auth_as(): + response = client.post(path, json={"role": "user", "content": "hi"}) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# GET /v1/threads/{thread_id}/messages, GET /threads/{thread_id}/messages +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/messages", "/threads/thr_1/messages"], +) +def test_get_messages_happy(client, auth_as, patched_threads, path): + """Pins ``GET /v1/threads/{thread_id}/messages`` and ``GET /threads/{thread_id}/messages``.""" + with auth_as(): + response = client.get(path) + assert response.status_code == 200 + assert normalize(response.json()) == { + "object": "list", + "data": [ + { + "id": "", + "object": "thread.message", + "thread_id": "thr_1", + "role": "user", + "content": [], + } + ], + "first_id": "msg_1", + "last_id": "msg_1", + "has_more": False, + } + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/messages", "/threads/thr_1/messages"], +) +def test_get_messages_error(client, auth_as, no_router, path): + """Pins ``GET /v1/threads/{thread_id}/messages`` / ``GET /threads/{thread_id}/messages`` (error).""" + with auth_as(): + response = client.get(path) + assert response.status_code == 500 + assert len(response.content) > 0 + + +# --------------------------------------------------------------------------- +# POST /v1/threads/{thread_id}/runs, POST /threads/{thread_id}/runs +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/runs", "/threads/thr_1/runs"], +) +def test_run_thread_happy(client, auth_as, patched_threads, path): + """Pins ``POST /v1/threads/{thread_id}/runs`` and ``POST /threads/{thread_id}/runs``.""" + payload = {"assistant_id": "asst_1"} + with auth_as(): + response = client.post(path, json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "id": "", + "object": "thread.run", + "thread_id": "thr_1", + "assistant_id": "asst_1", + "status": "queued", + } + + +@pytest.mark.parametrize( + "path", + ["/v1/threads/thr_1/runs", "/threads/thr_1/runs"], +) +def test_run_thread_error(client, auth_as, no_router, path): + """Pins ``POST /v1/threads/{thread_id}/runs`` / ``POST /threads/{thread_id}/runs`` (error).""" + with auth_as(): + response = client.post(path, json={"assistant_id": "asst_1"}) + assert response.status_code == 500 + assert len(response.content) > 0 diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py index ad6b4016461..c6070437d35 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py @@ -1 +1,160 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for ``proxy_server.py`` llm-utils routes. + +Pins (PR2): + - POST /utils/token_counter + - GET /utils/supported_openai_params + - POST /utils/transform_request +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.proxy import proxy_server + +from .conftest import normalize # type: ignore[import-not-found] + +# --------------------------------------------------------------------------- +# POST /utils/token_counter +# --------------------------------------------------------------------------- + + +@pytest.fixture +def patched_token_counter(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(litellm, "disable_token_counter", False, raising=False) + monkeypatch.setattr( + litellm.utils, + "_select_tokenizer", + lambda model, custom_tokenizer=None: { + "type": "openai_tokenizer", + "tokenizer": None, + }, + ) + monkeypatch.setattr(litellm, "token_counter", lambda **kwargs: 7) + yield + + +def test_token_counter_happy_path(client, auth_as, patched_token_counter): + """Pins ``POST /utils/token_counter``.""" + payload = {"model": "gpt-4", "prompt": "Hi there"} + with auth_as(): + response = client.post("/utils/token_counter", json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "total_tokens": 7, + "request_model": "gpt-4", + "model_used": "gpt-4", + "tokenizer_type": "openai_tokenizer", + "original_response": None, + "error": False, + "error_message": None, + "status_code": None, + } + + +def test_token_counter_missing_input_returns_400( + client, auth_as, patched_token_counter +): + """Pins ``POST /utils/token_counter`` (error: missing input).""" + with auth_as(): + response = client.post("/utils/token_counter", json={"model": "gpt-4"}) + assert response.status_code == 400 + assert "prompt or messages or contents" in response.text + + +# --------------------------------------------------------------------------- +# GET /utils/supported_openai_params +# --------------------------------------------------------------------------- + + +@pytest.fixture +def patched_supported_params(monkeypatch): + monkeypatch.setattr( + litellm, + "get_llm_provider", + lambda model: (model, "openai", None, None), + ) + monkeypatch.setattr( + litellm, + "get_supported_openai_params", + lambda model, custom_llm_provider=None: ["max_tokens", "temperature", "top_p"], + ) + yield + + +def test_supported_openai_params_happy_path(client, auth_as, patched_supported_params): + """Pins ``GET /utils/supported_openai_params``.""" + with auth_as(): + response = client.get( + "/utils/supported_openai_params", params={"model": "gpt-4"} + ) + assert response.status_code == 200 + assert normalize(response.json()) == { + "supported_openai_params": ["max_tokens", "temperature", "top_p"], + } + + +def test_supported_openai_params_invalid_model(client, auth_as, monkeypatch): + """Pins ``GET /utils/supported_openai_params`` (error: unknown model).""" + + def _raise(model): + raise Exception("unknown") + + monkeypatch.setattr(litellm, "get_llm_provider", _raise) + with auth_as(): + response = client.get("/utils/supported_openai_params", params={"model": "??"}) + assert response.status_code == 400 + assert "Could not map model" in response.text + + +# --------------------------------------------------------------------------- +# POST /utils/transform_request +# --------------------------------------------------------------------------- + + +@pytest.fixture +def patched_transform(monkeypatch): + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr(proxy_server, "is_request_body_safe", lambda **kwargs: True) + + def _fake_return_raw_request(endpoint, kwargs): + return { + "raw_request_api_base": "https://api.openai.com/v1/chat/completions", + "raw_request_body": kwargs, + "raw_request_headers": {"Authorization": "Bearer redacted"}, + } + + monkeypatch.setattr("litellm.utils.return_raw_request", _fake_return_raw_request) + yield + + +def test_transform_request_happy_path(client, auth_as, patched_transform): + """Pins ``POST /utils/transform_request``.""" + payload = {"call_type": "completion", "request_body": {"model": "gpt-4"}} + with auth_as(): + response = client.post("/utils/transform_request", json=payload) + assert response.status_code == 200 + assert normalize(response.json()) == { + "raw_request_api_base": "https://api.openai.com/v1/chat/completions", + "raw_request_body": {"model": "gpt-4"}, + "raw_request_headers": {"Authorization": "Bearer redacted"}, + } + + +def test_transform_request_unsafe_body(client, auth_as, monkeypatch): + """Pins ``POST /utils/transform_request`` (error: unsafe body).""" + monkeypatch.setattr(proxy_server, "llm_router", None) + + def _raise(**kwargs): + raise ValueError("unsafe model") + + monkeypatch.setattr(proxy_server, "is_request_body_safe", _raise) + payload = {"call_type": "completion", "request_body": {"model": "evil"}} + with auth_as(): + response = client.post("/utils/transform_request", json=payload) + assert response.status_code == 400 + assert "unsafe" in response.text or "error" in response.text diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index ad6b4016461..ec8b06d9c97 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -1 +1,817 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for spend-counter helpers in proxy_server. + +Pins covered: +- ``get_current_spend`` +- ``increment_spend_counters`` +- ``_reconcile_budget_reservation_for_counter_update`` +- ``_increment_end_user_and_tag_spend_counters`` +- ``_increment_org_spend_counter`` +- ``_init_and_increment_unreserved_spend_counter`` +- ``_init_and_increment_spend_counter`` +- ``_init_and_increment_window_spend_counter`` +- ``_ensure_spend_counter_initialized`` +- ``_get_source_cache_base_spend`` +- ``_ensure_window_spend_counter_initialized`` +- ``_is_spend_counter_cache_warm`` +- ``_increment_spend_counter_cache`` +- ``_invalidate_spend_counter`` +- ``update_cache`` +""" + +from __future__ import annotations + +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.proxy_server as ps + +from .conftest import normalize + + +def _make_spend_counter_cache( + *, + redis_get_value=None, + redis_get_side_effect=None, + redis_increment_value=None, + redis_increment_side_effect=None, + in_memory_value=None, + with_redis: bool = True, +): + cache = MagicMock() + cache.in_memory_cache = MagicMock() + cache.in_memory_cache.get_cache = MagicMock(return_value=in_memory_value) + cache.in_memory_cache.set_cache = MagicMock() + cache.in_memory_cache.delete_cache = MagicMock() + if with_redis: + cache.redis_cache = MagicMock() + cache.redis_cache.async_get_cache = AsyncMock( + return_value=redis_get_value, side_effect=redis_get_side_effect + ) + cache.redis_cache.async_increment = AsyncMock( + return_value=redis_increment_value, + side_effect=redis_increment_side_effect, + ) + cache.redis_cache.async_delete_cache = AsyncMock() + else: + cache.redis_cache = None + cache.async_increment_cache = AsyncMock(return_value=redis_increment_value) + cache.async_get_cache = AsyncMock(return_value=None) + cache.async_set_cache = AsyncMock() + cache.async_delete_cache = AsyncMock() + cache.async_set_cache_pipeline = AsyncMock() + return cache + + +def _make_user_api_key_cache(get_value=None, get_side_effect=None): + cache = MagicMock() + cache.async_get_cache = AsyncMock( + return_value=get_value, side_effect=get_side_effect + ) + cache.async_set_cache_pipeline = AsyncMock() + return cache + + +# --------------------------------------------------------------------------- +# get_current_spend +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_get_current_spend_reads_redis_first(monkeypatch): + fake_cache = _make_spend_counter_cache(redis_get_value=42.5) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=0.0) + + observed = { + "value": result, + "redis_called": fake_cache.redis_cache.async_get_cache.called, + "in_memory_called": fake_cache.in_memory_cache.get_cache.called, + } + assert normalize(observed) == { + "value": 42.5, + "redis_called": True, + "in_memory_called": False, + } + + +@pytest.mark.asyncio +async def test_get_current_spend_redis_error_falls_back_to_in_memory(monkeypatch): + fake_cache = _make_spend_counter_cache( + redis_get_side_effect=RuntimeError("redis down"), + in_memory_value=17.0, + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + result = await ps.get_current_spend( + counter_key="spend:key:abc", fallback_spend=99.0 + ) + assert result == 17.0 + + +# --------------------------------------------------------------------------- +# increment_spend_counters +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_increment_spend_counters_increments_all_buckets(monkeypatch): + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=5.0 + ) + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + + async def _fake_coalesced(**kwargs): + return None + + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced) + ) + + await ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + response_cost=5.0, + ) + + observed = { + "redis_increment_called": fake_cache.redis_cache.async_increment.called, + "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "user_cache_used": fake_user_cache.async_get_cache.called, + } + assert normalize(observed) == { + "redis_increment_called": True, + "increment_calls": 4, + "user_cache_used": True, + } + + +@pytest.mark.asyncio +async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation( + monkeypatch, +): + fake_cache = _make_spend_counter_cache() + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + reservation = {"finalized": False} + + await ps.increment_spend_counters( + token="hashed-tok", + team_id="t1", + user_id="u1", + response_cost=0, + budget_reservation=reservation, + ) + + assert reservation == {"finalized": True} + assert fake_cache.redis_cache.async_increment.called is False + + +# --------------------------------------------------------------------------- +# _reconcile_budget_reservation_for_counter_update +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none(): + result = await ps._reconcile_budget_reservation_for_counter_update( + budget_reservation=None, response_cost=1.0 + ) + assert result == set() + + +@pytest.mark.asyncio +async def test_reconcile_budget_reservation_for_counter_update_failure_invalidates( + monkeypatch, +): + """Reservation reconcile raising must invalidate reserved counters but + not propagate the exception.""" + import litellm.proxy.spend_tracking.budget_reservation as br + + monkeypatch.setattr( + br, + "get_reserved_counter_keys", + MagicMock(return_value={"spend:key:abc"}), + ) + monkeypatch.setattr( + br, + "reconcile_budget_reservation", + AsyncMock(side_effect=RuntimeError("boom")), + ) + fake_invalidate = AsyncMock() + monkeypatch.setattr(br, "invalidate_budget_reservation_counters", fake_invalidate) + + result = await ps._reconcile_budget_reservation_for_counter_update( + budget_reservation={"foo": "bar"}, response_cost=1.0 + ) + + assert result == {"spend:key:abc"} + assert fake_invalidate.called is True + + +# --------------------------------------------------------------------------- +# _increment_end_user_and_tag_spend_counters +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_increment_end_user_and_tag_spend_counters_increments_each_unique_tag( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=3.0 + ) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps._increment_end_user_and_tag_spend_counters( + end_user_id="eu1", + tags=["a", "b", "a", "", None], + response_cost=3.0, + reserved_counter_keys=set(), + ) + + observed = { + "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count, + "called": fake_cache.redis_cache.async_increment.called, + } + assert normalize(observed) == { + "increment_calls": 3, + "in_memory_set_calls": 3, + "called": True, + } + + +@pytest.mark.asyncio +async def test_increment_end_user_and_tag_spend_counters_no_end_user_no_tags_invalid_input_noop( + monkeypatch, +): + fake_cache = _make_spend_counter_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._increment_end_user_and_tag_spend_counters( + end_user_id=None, + tags=None, + response_cost=1.0, + reserved_counter_keys=set(), + ) + + assert fake_cache.redis_cache.async_increment.called is False + + +# --------------------------------------------------------------------------- +# _increment_org_spend_counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_increment_org_spend_counter_increments_when_org_present(monkeypatch): + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=10.0 + ) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps._increment_org_spend_counter( + org_id="org-1", + response_cost=10.0, + reserved_counter_keys=set(), + ) + + observed = { + "increment_called": fake_cache.redis_cache.async_increment.called, + "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs[ + "key" + ], + } + assert normalize(observed) == { + "increment_called": True, + "increment_calls": 1, + "counter_key_arg": "spend:org:org-1", + } + + +@pytest.mark.asyncio +async def test_increment_org_spend_counter_no_org_is_noop_invalid_id(monkeypatch): + fake_cache = _make_spend_counter_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._increment_org_spend_counter( + org_id=None, + response_cost=1.0, + reserved_counter_keys=set(), + ) + + assert fake_cache.redis_cache.async_increment.called is False + + +# --------------------------------------------------------------------------- +# _init_and_increment_unreserved_spend_counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_init_and_increment_unreserved_spend_counter_skips_reserved_keys( + monkeypatch, +): + fake_cache = _make_spend_counter_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._init_and_increment_unreserved_spend_counter( + counter_key="spend:tag:x", + source_cache_key="tag:x", + increment=1.0, + reserved_counter_keys={"spend:tag:x"}, + ) + + assert fake_cache.redis_cache.async_increment.called is False + + +@pytest.mark.asyncio +async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_reserved( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=2.0 + ) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps._init_and_increment_unreserved_spend_counter( + counter_key="spend:tag:y", + source_cache_key="tag:y", + increment=2.0, + reserved_counter_keys=set(), + ) + + observed = { + "increment_called": fake_cache.redis_cache.async_increment.called, + "redis_get_called": fake_cache.redis_cache.async_get_cache.called, + "reseed_consulted": True, + } + assert observed == { + "increment_called": True, + "redis_get_called": True, + "reseed_consulted": True, + } + + +# --------------------------------------------------------------------------- +# _init_and_increment_spend_counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypatch): + fake_cache = _make_spend_counter_cache( + redis_get_value=11.0, redis_increment_value=14.0 + ) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + reseed = AsyncMock(return_value=None) + monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed) + + await ps._init_and_increment_spend_counter( + counter_key="spend:key:k", + source_cache_key="k", + increment=3.0, + ) + + observed = { + "reseed_called": reseed.called, + "increment_called": fake_cache.redis_cache.async_increment.called, + "in_memory_seeded_from_redis": fake_cache.in_memory_cache.set_cache.called, + } + assert normalize(observed) == { + "reseed_called": False, + "increment_called": True, + "in_memory_seeded_from_redis": True, + } + + +# --------------------------------------------------------------------------- +# _init_and_increment_window_spend_counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_init_and_increment_window_spend_counter_increments_when_initialized( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_get_value=0.0, redis_increment_value=5.0 + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, + "coalesced_window", + AsyncMock(return_value=0.0), + ) + + await ps._init_and_increment_window_spend_counter( + counter_key="spend:key:k:window:1d", + entity_type="Key", + entity_id="k", + window_start=datetime(2024, 1, 1), + increment=5.0, + ) + + observed = { + "redis_increment_called": fake_cache.redis_cache.async_increment.called, + "increment_calls": fake_cache.redis_cache.async_increment.call_count, + "in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count, + } + assert normalize(observed) == { + "redis_increment_called": True, + "increment_calls": 1, + "in_memory_set_calls": 2, + } + + +@pytest.mark.asyncio +async def test_init_and_increment_window_spend_counter_missing_window_start_invalid_skips( + monkeypatch, +): + fake_cache = _make_spend_counter_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._init_and_increment_window_spend_counter( + counter_key="spend:key:k:window:1d", + entity_type="Key", + entity_id="k", + window_start=None, + increment=5.0, + ) + + assert fake_cache.redis_cache.async_increment.called is False + + +# --------------------------------------------------------------------------- +# _ensure_spend_counter_initialized +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ensure_spend_counter_initialized_warm_skips_reseed_and_source( + monkeypatch, +): + fake_cache = _make_spend_counter_cache(redis_get_value=20.0) + fake_user_cache = _make_user_api_key_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + reseed = AsyncMock(return_value=None) + monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed) + + await ps._ensure_spend_counter_initialized( + counter_key="spend:user:u", + source_cache_key="u", + ) + + observed = { + "warm_check_redis": fake_cache.redis_cache.async_get_cache.called, + "reseed_called": reseed.called, + "source_cache_called": fake_user_cache.async_get_cache.called, + } + assert normalize(observed) == { + "warm_check_redis": True, + "reseed_called": False, + "source_cache_called": False, + } + + +@pytest.mark.asyncio +async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_get_value=None, redis_increment_value=7.0 + ) + fake_user_cache = _make_user_api_key_cache(get_value={"spend": 7.0}) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None) + ) + + await ps._ensure_spend_counter_initialized( + counter_key="spend:user:u", + source_cache_key="u", + ) + + observed = { + "source_cache_called": fake_user_cache.async_get_cache.called, + "seed_increment_called": fake_cache.redis_cache.async_increment.called, + "warm_check_done": fake_cache.redis_cache.async_get_cache.called, + } + assert normalize(observed) == { + "source_cache_called": True, + "seed_increment_called": True, + "warm_check_done": True, + } + + +# --------------------------------------------------------------------------- +# _get_source_cache_base_spend +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_get_source_cache_base_spend_reads_first_hit_from_list(monkeypatch): + fake_user_cache = MagicMock() + + async def _get(key, **kwargs): + if key == "miss": + return None + if key == "hit-obj": + obj = MagicMock() + obj.spend = 12.0 + return obj + return None + + fake_user_cache.async_get_cache = AsyncMock(side_effect=_get) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + + result = await ps._get_source_cache_base_spend( + source_cache_key=["miss", "hit-obj", "miss2"] + ) + + observed = { + "result": result, + "calls": fake_user_cache.async_get_cache.call_count, + "stopped_after_hit": fake_user_cache.async_get_cache.call_count == 2, + } + assert normalize(observed) == { + "result": 12.0, + "calls": 2, + "stopped_after_hit": True, + } + + +@pytest.mark.asyncio +async def test_get_source_cache_base_spend_no_hits_returns_zero_fallback(monkeypatch): + """All cache lookups miss — function falls back to 0.0 (no error).""" + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + + result = await ps._get_source_cache_base_spend(source_cache_key="missing-key") + assert result == 0.0 + + +# --------------------------------------------------------------------------- +# _ensure_window_spend_counter_initialized +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ensure_window_spend_counter_initialized_warm_returns_true(monkeypatch): + fake_cache = _make_spend_counter_cache(redis_get_value=3.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", None) + window_reseed = AsyncMock(return_value=0.0) + monkeypatch.setattr(ps.SpendCounterReseed, "coalesced_window", window_reseed) + + initialized = await ps._ensure_window_spend_counter_initialized( + counter_key="spend:key:k:window:1d", + entity_type="Key", + entity_id="k", + window_start=datetime(2024, 1, 1), + ) + + observed = { + "initialized": initialized, + "reseed_called": window_reseed.called, + "redis_get_called": fake_cache.redis_cache.async_get_cache.called, + } + assert normalize(observed) == { + "initialized": True, + "reseed_called": False, + "redis_get_called": True, + } + + +@pytest.mark.asyncio +async def test_ensure_window_spend_counter_initialized_db_failure_invalid_returns_false( + monkeypatch, +): + fake_cache = _make_spend_counter_cache(redis_get_value=None) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr( + ps.SpendCounterReseed, + "coalesced_window", + AsyncMock(return_value=None), + ) + + initialized = await ps._ensure_window_spend_counter_initialized( + counter_key="spend:key:k:window:1d", + entity_type="Key", + entity_id="k", + window_start=datetime(2024, 1, 1), + ) + + assert initialized is False + + +# --------------------------------------------------------------------------- +# _is_spend_counter_cache_warm +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_is_spend_counter_cache_warm_redis_hit_seeds_in_memory(monkeypatch): + fake_cache = _make_spend_counter_cache(redis_get_value=99.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + result = await ps._is_spend_counter_cache_warm(counter_key="spend:user:u") + + observed = { + "result": result, + "redis_get_called": fake_cache.redis_cache.async_get_cache.called, + "in_memory_set_called": fake_cache.in_memory_cache.set_cache.called, + } + assert normalize(observed) == { + "result": True, + "redis_get_called": True, + "in_memory_set_called": True, + } + + +@pytest.mark.asyncio +async def test_is_spend_counter_cache_warm_redis_error_falls_back_to_in_memory( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_get_side_effect=RuntimeError("redis err"), + in_memory_value=None, + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + result = await ps._is_spend_counter_cache_warm(counter_key="spend:user:u") + assert result is False + + +# --------------------------------------------------------------------------- +# _increment_spend_counter_cache +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkeypatch): + fake_cache = _make_spend_counter_cache(redis_increment_value=44.0) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + result = await ps._increment_spend_counter_cache( + counter_key="spend:key:k", increment=4.0 + ) + + observed = { + "result": result, + "redis_increment_called": fake_cache.redis_cache.async_increment.called, + "in_memory_set_called": fake_cache.in_memory_cache.set_cache.called, + } + assert normalize(observed) == { + "result": 44.0, + "redis_increment_called": True, + "in_memory_set_called": True, + } + + +@pytest.mark.asyncio +async def test_increment_spend_counter_cache_redis_error_raises_and_invalidates( + monkeypatch, +): + fake_cache = _make_spend_counter_cache( + redis_increment_side_effect=RuntimeError("incr fail") + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + with pytest.raises(RuntimeError): + await ps._increment_spend_counter_cache( + counter_key="spend:key:k", increment=1.0 + ) + + assert fake_cache.in_memory_cache.delete_cache.called is True + assert fake_cache.redis_cache.async_delete_cache.called is True + + +# --------------------------------------------------------------------------- +# _invalidate_spend_counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch): + fake_cache = _make_spend_counter_cache() + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._invalidate_spend_counter(counter_key="spend:key:k") + + observed = { + "in_memory_delete_called": fake_cache.in_memory_cache.delete_cache.called, + "redis_delete_called": fake_cache.redis_cache.async_delete_cache.called, + "delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs[ + "key" + ], + } + assert normalize(observed) == { + "in_memory_delete_called": True, + "redis_delete_called": True, + "delete_args_key": "spend:key:k", + } + + +@pytest.mark.asyncio +async def test_invalidate_spend_counter_swallows_redis_failure_no_raise(monkeypatch): + fake_cache = _make_spend_counter_cache() + fake_cache.redis_cache.async_delete_cache = AsyncMock( + side_effect=RuntimeError("redis down") + ) + monkeypatch.setattr(ps, "spend_counter_cache", fake_cache) + + await ps._invalidate_spend_counter(counter_key="spend:key:k") + + assert fake_cache.in_memory_cache.delete_cache.called is True + + +# --------------------------------------------------------------------------- +# update_cache +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_update_cache_no_cached_entities_schedules_pipeline_flush(monkeypatch): + fake_user_cache = _make_user_api_key_cache(get_value=None) + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + + await ps.update_cache( + token=None, + user_id="u1", + end_user_id="eu1", + team_id="t1", + response_cost=1.0, + parent_otel_span=None, + tags=["x"], + ) + + observed = { + "lookups": fake_user_cache.async_get_cache.call_count, + "got_user": True, + "got_team": True, + } + assert normalize(observed) == { + "lookups": 4, + "got_user": True, + "got_team": True, + } + + +@pytest.mark.asyncio +async def test_update_cache_user_cache_failure_invalid_state_is_swallowed(monkeypatch): + """An inner _update_user_cache raising must not propagate — update_cache + catches and logs, the public coroutine still completes normally.""" + fake_user_cache = MagicMock() + fake_user_cache.async_get_cache = AsyncMock(side_effect=RuntimeError("cache down")) + fake_user_cache.async_set_cache_pipeline = AsyncMock() + monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache) + + result = await ps.update_cache( + token=None, + user_id="u1", + end_user_id=None, + team_id=None, + response_cost=1.0, + parent_otel_span=None, + tags=None, + ) + + assert result is None diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index ad6b4016461..33de1ede917 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -1 +1,555 @@ -"""Placeholder. Filled by a follow-up PR per the Notion plan.""" +"""Behavior pins for the proxy_server streaming helpers. + +Pins covered: +- ``data_generator`` +- ``async_assistants_data_generator`` +- ``_get_client_requested_model_for_streaming`` +- ``_restamp_streaming_chunk_model`` +- ``_fast_serialize_simple_model_response_stream`` +- ``_serialize_streaming_chunk`` +- ``_apply_streaming_chunk_hooks`` +- ``_format_streaming_sse_chunk`` +- ``async_data_generator`` +- ``select_data_generator`` +""" + +from __future__ import annotations + +import json +from typing import Any, AsyncIterator +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.proxy_server as ps +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.proxy_server import ( + _apply_streaming_chunk_hooks, + _fast_serialize_simple_model_response_stream, + _format_streaming_sse_chunk, + _get_client_requested_model_for_streaming, + _restamp_streaming_chunk_model, + _serialize_streaming_chunk, + async_assistants_data_generator, + async_data_generator, + data_generator, + select_data_generator, +) +from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage + +from .conftest import normalize + + +def _user_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test-key", user_id="u") + + +def _simple_chunk(model: str = "gpt-4", content: str = "hi") -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-test", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(content=content, role="assistant"), + ) + ], + created=0, + model=model, + object="chat.completion.chunk", + ) + + +async def _async_iter(items): + for it in items: + yield it + + +async def _async_iter_raises(exc: Exception): + # yield once then raise — exercises the mid-stream failure branch + yield _simple_chunk(content="partial") + raise exc + + +# --------------------------------------------------------------------------- +# data_generator +# --------------------------------------------------------------------------- + + +def test_data_generator_yields_sse_lines_for_dict_chunks(): + class DictChunk: + def __init__(self, payload): + self._payload = payload + + def dict(self): + return self._payload + + chunks = [ + DictChunk({"id": "1", "object": "chat.completion.chunk", "model": "gpt-4"}), + DictChunk({"id": "2", "object": "chat.completion.chunk", "model": "gpt-4"}), + ] + out = list(data_generator(chunks)) + + assert len(out) == 2 + payloads = [json.loads(line.removeprefix("data: ").rstrip("\n\n")) for line in out] + assert normalize(payloads[0]) == { + "id": "", + "object": "chat.completion.chunk", + "model": "gpt-4", + } + assert payloads[1]["model"] == "gpt-4" + + +def test_data_generator_fallback_when_dict_raises_exception(): + class BadChunk: + def dict(self): + raise RuntimeError("cannot serialize") + + # When .dict() raises, the inner json.dumps(chunk) on a non-JSON-serializable + # instance also raises — the generator does not catch the second failure. + with pytest.raises((TypeError, RuntimeError)): + list(data_generator([BadChunk()])) + + +# --------------------------------------------------------------------------- +# async_assistants_data_generator +# --------------------------------------------------------------------------- + + +class _FakeAssistantsStream: + """Mimic the async-context-manager + async-iterable shape of the + assistants streaming object (e.g. AssistantEventHandler).""" + + def __init__(self, chunks): + self._chunks = chunks + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return False + + def __aiter__(self): + async def _gen(): + for c in self._chunks: + yield c + + return _gen() + + +@pytest.mark.asyncio +async def test_async_assistants_data_generator_yields_sse_and_done(monkeypatch): + chunk = _simple_chunk(content="hello") + + async def _passthrough_hook(*, user_api_key_dict, response, data, **kwargs): + return response + + monkeypatch.setattr( + ps.proxy_logging_obj, + "async_post_call_streaming_hook", + _passthrough_hook, + ) + + stream = _FakeAssistantsStream([chunk]) + out = [] + async for line in async_assistants_data_generator( + response=stream, + user_api_key_dict=_user_auth(), + request_data={}, + ): + out.append(line) + + assert out[-1] == "data: [DONE]\n\n" + body = json.loads(out[0].removeprefix("data: ").rstrip("\n\n")) + assert normalize(body) == { + "id": "", + "created": "", + "model": "gpt-4", + "object": "chat.completion.chunk", + "choices": [ + { + "index": 0, + "delta": {"content": "hello", "role": "assistant"}, + } + ], + } + + +@pytest.mark.asyncio +async def test_async_assistants_data_generator_hook_failure_yields_error_chunk( + monkeypatch, +): + async def _boom_hook(*args, **kwargs): + raise RuntimeError("hook exploded") + + async def _noop_failure(*args, **kwargs): + return None + + monkeypatch.setattr( + ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom_hook + ) + monkeypatch.setattr(ps.proxy_logging_obj, "post_call_failure_hook", _noop_failure) + + stream = _FakeAssistantsStream([_simple_chunk()]) + out = [] + async for line in async_assistants_data_generator( + response=stream, + user_api_key_dict=_user_auth(), + request_data={}, + ): + out.append(line) + + assert any("error" in line for line in out) + assert out[-1].startswith('data: {"error":') + + +# --------------------------------------------------------------------------- +# _get_client_requested_model_for_streaming +# --------------------------------------------------------------------------- + + +def test_get_client_requested_model_for_streaming_prefers_client_requested(): + request_data = { + "_litellm_client_requested_model": "gpt-4", + "model": "openai/internal-gpt-4", + "litellm_call_id": "abc", + } + result = _get_client_requested_model_for_streaming(request_data) + assert result == "gpt-4" + + snapshot = { + "result": result, + "client_field_preserved": request_data["_litellm_client_requested_model"], + "model_field_preserved": request_data["model"], + } + assert normalize(snapshot) == { + "result": "gpt-4", + "client_field_preserved": "gpt-4", + "model_field_preserved": "openai/internal-gpt-4", + } + + +def test_get_client_requested_model_for_streaming_falls_back_to_model_field(): + result = _get_client_requested_model_for_streaming({"model": "claude-sonnet"}) + assert result == "claude-sonnet" + + +def test_get_client_requested_model_for_streaming_missing_returns_empty_invalid(): + """When neither key is set or values are non-strings, the helper returns "" + rather than raising — callers depend on this to skip restamping.""" + assert _get_client_requested_model_for_streaming({}) == "" + assert _get_client_requested_model_for_streaming({"model": 123}) == "" + + +# --------------------------------------------------------------------------- +# _restamp_streaming_chunk_model +# --------------------------------------------------------------------------- + + +def test_restamp_streaming_chunk_model_overrides_model_on_basemodel(): + chunk = _simple_chunk(model="openai/internal-x") + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="gpt-4", + request_data={"litellm_call_id": "id-1"}, + model_mismatch_logged=False, + ) + snapshot = { + "model": new_chunk.model, + "logged": logged, + "same_object": new_chunk is chunk, + } + assert snapshot == {"model": "gpt-4", "logged": True, "same_object": True} + + +def test_restamp_streaming_chunk_model_overrides_model_on_dict(): + chunk = {"model": "internal", "choices": []} + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="gpt-4", + request_data={}, + model_mismatch_logged=True, + ) + assert new_chunk["model"] == "gpt-4" + assert logged is True + + +def test_restamp_streaming_chunk_model_invalid_chunk_type_unchanged(): + """For a non-BaseModel, non-dict chunk the helper returns it as-is + along with the original ``model_mismatch_logged`` flag.""" + chunk = "raw string chunk" + new_chunk, logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="gpt-4", + request_data={}, + model_mismatch_logged=False, + ) + assert new_chunk == "raw string chunk" + assert logged is False + + +# --------------------------------------------------------------------------- +# _fast_serialize_simple_model_response_stream +# --------------------------------------------------------------------------- + + +def test_fast_serialize_simple_model_response_stream_returns_bytes_payload(): + chunk = _simple_chunk() + result = _fast_serialize_simple_model_response_stream(chunk) + assert isinstance(result, bytes) + payload = json.loads(result) + assert normalize(payload) == { + "id": "", + "object": "chat.completion.chunk", + "created": "", + "model": "gpt-4", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "hi"}, + } + ], + } + + +def test_fast_serialize_simple_model_response_stream_with_usage_returns_none_invalid(): + """Fast path bails (returns None) when ``usage`` is populated — the slow + path is required to preserve usage fields. Returning None here is the + "I cannot handle this" sentinel, not a hard error.""" + chunk = _simple_chunk() + chunk.usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2) + assert _fast_serialize_simple_model_response_stream(chunk) is None + + +# --------------------------------------------------------------------------- +# _serialize_streaming_chunk +# --------------------------------------------------------------------------- + + +def test_serialize_streaming_chunk_simple_uses_fast_path_bytes(): + result = _serialize_streaming_chunk(_simple_chunk()) + assert isinstance(result, bytes) + payload = json.loads(result) + assert normalize(payload) == { + "id": "", + "object": "chat.completion.chunk", + "created": "", + "model": "gpt-4", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "hi"}, + } + ], + } + + +def test_serialize_streaming_chunk_invalid_input_raises_attribute_error(): + """The helper is typed as ``BaseModel`` — handing it a plain dict trips + the attribute-access path (no ``model_dump_json``).""" + with pytest.raises(AttributeError): + _serialize_streaming_chunk({"not": "a model"}) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# _apply_streaming_chunk_hooks +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): + chunk = _simple_chunk(content="abc") + + async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None): + return response + + monkeypatch.setattr( + ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough + ) + + new_chunk, new_str = await _apply_streaming_chunk_hooks( + chunk=chunk, + user_api_key_dict=_user_auth(), + request_data={}, + str_so_far="prior:", + ) + + observed = { + "chunk_is_basemodel": isinstance(new_chunk, ModelResponseStream), + "str_so_far": new_str, + "grew": len(new_str) > len("prior:"), + } + assert observed == { + "chunk_is_basemodel": True, + "str_so_far": "prior:abc", + "grew": True, + } + + +@pytest.mark.asyncio +async def test_apply_streaming_chunk_hooks_hook_raises_exception(monkeypatch): + async def _boom(*args, **kwargs): + raise RuntimeError("hook failed") + + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom) + + with pytest.raises(RuntimeError): + await _apply_streaming_chunk_hooks( + chunk=_simple_chunk(), + user_api_key_dict=_user_auth(), + request_data={}, + str_so_far="", + ) + + +# --------------------------------------------------------------------------- +# _format_streaming_sse_chunk +# --------------------------------------------------------------------------- + + +def test_format_streaming_sse_chunk_handles_bytes_and_str_shapes(): + bytes_out = _format_streaming_sse_chunk(b'{"a":1}') + str_out = _format_streaming_sse_chunk('{"a":1}') + + snapshot = { + "bytes_out": bytes_out, + "str_out": str_out, + "bytes_starts_with_data": bytes_out.startswith(b"data: "), + } + assert snapshot == { + "bytes_out": b'data: {"a":1}\n\n', + "str_out": 'data: {"a":1}\n\n', + "bytes_starts_with_data": True, + } + + +def test_format_streaming_sse_chunk_invalid_empty_string_still_wraps(): + """Edge case: empty string still gets the ``data: \\n\\n`` wrapping + — clients expect SSE shape even on empty payloads.""" + result = _format_streaming_sse_chunk("") + assert result == "data: \n\n" + + +# --------------------------------------------------------------------------- +# async_data_generator +# --------------------------------------------------------------------------- + + +def _patch_logging_flags(monkeypatch, needs_wrap=False, needs_per_chunk=False): + monkeypatch.setattr( + ps.proxy_logging_obj, + "needs_iterator_wrap", + lambda: needs_wrap, + ) + monkeypatch.setattr( + ps.proxy_logging_obj, + "needs_per_chunk_streaming_hook", + lambda: needs_per_chunk, + ) + # ``_fire_deferred_stream_logging`` is a classmethod — patch the + # underlying function so the no-wrap branch is a no-op rather than + # touching real logging globals. + monkeypatch.setattr( + ps.ProxyLogging, + "_fire_deferred_stream_logging", + staticmethod(lambda request_data: None), + ) + + +@pytest.mark.asyncio +async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch): + _patch_logging_flags(monkeypatch) + + response = _async_iter([_simple_chunk(content="hello")]) + out = [] + async for line in async_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ): + out.append(line) + + assert out[-1] == "data: [DONE]\n\n" + # First chunk is bytes (fast path) wrapped via _format_streaming_sse_chunk. + first = out[0] + assert isinstance(first, bytes) + payload = json.loads(first.removeprefix(b"data: ").rstrip(b"\n\n")) + assert normalize(payload) == { + "id": "", + "object": "chat.completion.chunk", + "created": "", + "model": "gpt-4", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "hello"}, + } + ], + } + + +@pytest.mark.asyncio +async def test_async_data_generator_mid_stream_exception_yields_error_payload( + monkeypatch, +): + _patch_logging_flags(monkeypatch) + + async def _noop_failure(*args, **kwargs): + return None + + monkeypatch.setattr(ps.proxy_logging_obj, "post_call_failure_hook", _noop_failure) + + response = _async_iter_raises(RuntimeError("upstream blew up")) + out = [] + async for line in async_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={}, + ): + out.append(line) + + # First entry is the successful "partial" chunk (bytes), last is the error. + assert any( + isinstance(item, str) and item.startswith('data: {"error":') for item in out + ) + + +# --------------------------------------------------------------------------- +# select_data_generator +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_select_data_generator_returns_async_generator(monkeypatch): + _patch_logging_flags(monkeypatch) + + response = _async_iter([_simple_chunk()]) + gen = select_data_generator( + response=response, + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ) + + # Drain to confirm it really is an async iterator emitting SSE shape. + collected = [] + async for line in gen: + collected.append(line) + + snapshot = { + "is_async_iterable": hasattr(gen, "__aiter__"), + "yielded_at_least_one": len(collected) >= 1, + "ends_with_done": collected[-1] == "data: [DONE]\n\n", + } + assert snapshot == { + "is_async_iterable": True, + "yielded_at_least_one": True, + "ends_with_done": True, + } + + +def test_select_data_generator_missing_required_kwarg_raises_type_error(): + """``select_data_generator`` requires all three keyword args — calling + without ``request_data`` raises TypeError at the wrapper, before any + streaming starts.""" + with pytest.raises(TypeError): + select_data_generator(response=_async_iter([]), user_api_key_dict=_user_auth()) # type: ignore[call-arg]