import os import sys from dotenv import load_dotenv load_dotenv() import io import os # this file is to test litellm/proxy sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import asyncio import logging import pytest from fastapi import Request from starlette.datastructures import URL, Headers, QueryParams import litellm from litellm.proxy._types import LiteLLMRoutes from litellm.proxy.auth.auth_utils import get_request_route from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.proxy_server import app # Configure logging logging.basicConfig( level=logging.DEBUG, # Set the desired logging level format="%(asctime)s - %(levelname)s - %(message)s", ) def test_routes_on_litellm_proxy(): """ Goal of this test: Test that we have all the critical OpenAI Routes on the Proxy server Fast API router this prevents accidentelly deleting /threads, or /batches etc """ # Force-load lazy features so the test sees the full route set. Continue # on per-feature import failure — the assertion below still catches # missing-route regressions. import importlib from litellm.proxy._lazy_features import LAZY_FEATURES registered_paths = [getattr(r, "path", "") for r in app.routes] for feat in LAZY_FEATURES: if any(rp.startswith(p) for p in feat.path_prefixes for rp in registered_paths): continue try: module = importlib.import_module(feat.module_path) feat.register_fn(app, module) except Exception as exc: print(f"warning: failed to force-load {feat.name}: {exc}") _all_routes = [] for route in app.routes: _path_as_str = str(route.path) if ":path" in _path_as_str: # remove the :path _path_as_str = _path_as_str.replace(":path", "") _all_routes.append(_path_as_str) print("ALL ROUTES on LiteLLM Proxy:", _all_routes) print("\n\n") print("ALL OPENAI ROUTES:", LiteLLMRoutes.openai_routes.value) for route in LiteLLMRoutes.openai_routes.value: # realtime routes - /realtime?model=gpt-4o if "realtime" in route: assert "/realtime" in _all_routes # wildcard patterns like /containers/* - check that base path exists elif RouteChecks._is_wildcard_pattern(pattern=route): # For wildcard patterns, check that the base path (without * and trailing /) exists base_path = route[:-1].rstrip( "/" ) # Remove the trailing * and any trailing / # Check if base path exists (e.g., /containers or /v1/containers) assert ( base_path in _all_routes ), f"Wildcard pattern {route} requires base path {base_path} to exist" else: assert route in _all_routes @pytest.mark.parametrize( "route,expected", [ # Test exact matches ("/chat/completions", True), ("/v1/chat/completions", True), ("/embeddings", True), ("/v1/models", True), ("/utils/token_counter", True), # Test routes with placeholders ("/engines/gpt-4/chat/completions", True), ("/openai/deployments/gpt-3.5-turbo/chat/completions", True), ("/threads/thread_49EIN5QF32s4mH20M7GFKdlZ", True), ("/v1/threads/thread_49EIN5QF32s4mH20M7GFKdlZ", True), ("/threads/thread_49EIN5QF32s4mH20M7GFKdlZ/messages", True), ("/v1/threads/thread_49EIN5QF32s4mH20M7GFKdlZ/runs", True), ("/v1/batches/123456", True), # Test non-OpenAI routes ("/some/random/route", False), ("/v2/chat/completions", False), ("/threads/invalid/format", False), ("/v1/non_existent_endpoint", False), # Bedrock Pass Through Routes ("/bedrock/model/cohere.command-r-v1:0/converse", True), ("/vertex-ai/model/text-embedding-004/embeddings", True), # LiteLLM native RAG routes ("/rag/ingest", True), ("/v1/rag/ingest", True), ("/rag/query", True), ("/v1/rag/query", True), ], ) def test_is_llm_api_route(route: str, expected: bool): assert RouteChecks.is_llm_api_route(route) == expected # Test-case for routes that are similar but should return False @pytest.mark.parametrize( "route", [ "/v1/threads/thread_id/invalid", "/threads/thread_id/invalid", "/v1/batches/123/invalid", "/engines/model/invalid/completions", ], ) def test_is_llm_api_route_similar_but_false(route: str): assert RouteChecks.is_llm_api_route(route) is False def test_anthropic_api_routes(): # allow non proxy admins to call anthropic api routes assert RouteChecks.is_llm_api_route(route="/v1/messages") is True def create_request(path: str, base_url: str = "http://testserver") -> Request: return Request( { "type": "http", "method": "GET", "scheme": "http", "server": ("testserver", 80), "path": path, "query_string": b"", "headers": Headers().raw, "client": ("testclient", 50000), "root_path": URL(base_url).path, } ) def test_get_request_route_with_base_url(): request = create_request( path="/genai/chat/completions", base_url="http://testserver/genai" ) result = get_request_route(request) assert result == "/chat/completions" def test_get_request_route_without_base_url(): request = create_request("/chat/completions") result = get_request_route(request) assert result == "/chat/completions" def test_get_request_route_with_nested_path(): request = create_request(path="/embeddings", base_url="http://testserver/ishaan") result = get_request_route(request) assert result == "/embeddings" def test_get_request_route_with_query_params(): request = create_request(path="/genai/test", base_url="http://testserver/genai") request.scope["query_string"] = b"param=value" result = get_request_route(request) assert result == "/test" def test_get_request_route_with_base_url_not_at_start(): request = create_request("/api/genai/test") result = get_request_route(request) assert result == "/api/genai/test" def _create_request_with_host_header(path: str, host_header: str) -> Request: return Request( { "type": "http", "method": "GET", "scheme": "http", "server": ("localhost", 4000), "path": path, "query_string": b"", "headers": [(b"host", host_header.encode())], "client": ("127.0.0.1", 50000), "root_path": "", } ) @pytest.mark.parametrize( "host_header", [ "localhost/?x=1", "localhost:4000/?x=1", "localhost/#test", "localhost:4000/#test", ], ) def test_get_request_route_not_bypassed_by_malformed_host(host_header: str): for protected_path in ["/health", "/user/new", "/key/generate", "/get/internal_user_settings"]: request = _create_request_with_host_header(path=protected_path, host_header=host_header) result = get_request_route(request) assert result == protected_path, ( f"Host: {host_header!r} caused route {protected_path!r} to resolve as {result!r}" )