mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(ci): fix six CircleCI test regressions on main (#44429)
* test(ci): fix four CircleCI regressions on main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ci): stop reloading auth_checks in unit tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(constants): cover CLI JWT expiry env parsing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(gateway): restore proxy lifespan after importing gateway.main in launch tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo <mateo@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2936307661
commit
cf22deb96a
9 changed files with 121 additions and 96 deletions
|
|
@ -99,7 +99,7 @@ export const MIGRATED_E2E_PAGES: Readonly<Record<string, MigratedPage>> = {
|
|||
group: "Settings",
|
||||
content: { role: "heading", name: "UI Theme Customization" },
|
||||
},
|
||||
logs: { segment: "logs", linkName: "Logs", content: { role: "heading", name: "Request Logs" } },
|
||||
logs: { segment: "logs", linkName: "Logs", content: { role: "tab", name: "Request Logs" } },
|
||||
"admin-panel": {
|
||||
segment: "admin-panel",
|
||||
linkName: "Admin Settings",
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import sys
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
from typing import Final, List, Optional
|
||||
import pytest
|
||||
from litellm._uuid import uuid
|
||||
import os
|
||||
|
|
@ -391,6 +391,7 @@ async def test_create_mcp_server_invalid_alias():
|
|||
@_SKIP_NO_MCP
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_redacts_credentials():
|
||||
mock_get_server: Final = mock.AsyncMock()
|
||||
with (
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
|
||||
|
|
@ -399,6 +400,10 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
|
||||
) as mock_get_prisma,
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
new=mock_get_server,
|
||||
),
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
|
||||
new_callable=mock.AsyncMock,
|
||||
|
|
@ -422,6 +427,18 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock_manager.reload_servers_from_database = mock.AsyncMock()
|
||||
|
||||
server_id = str(uuid.uuid4())
|
||||
stored_server: Final = LiteLLM_MCPServerTable(
|
||||
server_id=server_id,
|
||||
alias="Updated Server",
|
||||
url="https://updated.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
credentials={"auth_value": "secret"},
|
||||
teams=[],
|
||||
)
|
||||
mock_get_server.return_value = stored_server
|
||||
|
||||
updated_server = LiteLLM_MCPServerTable(
|
||||
server_id=server_id,
|
||||
alias="Updated Server",
|
||||
|
|
@ -458,6 +475,7 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock_update.assert_awaited_once()
|
||||
mock_manager.update_server.assert_called_once_with(updated_server)
|
||||
mock_manager.reload_servers_from_database.assert_awaited_once()
|
||||
mock_get_server.assert_awaited_once_with(mock_prisma, server_id)
|
||||
|
||||
|
||||
def test_validate_mcp_server_name_direct():
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import importlib
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
|
|
@ -12,10 +13,10 @@ import pytest
|
|||
from uvicorn.importer import import_from_string
|
||||
from uvicorn.main import main as uvicorn_main
|
||||
|
||||
import gateway.main
|
||||
from gateway.launch import GATEWAY_APP, main, pool_database_url, uvicorn_argv
|
||||
from litellm.proxy.db.db_url_settings import DatabaseURLSettings
|
||||
from litellm.proxy.db.pgbouncer import PGBOUNCER_POOLED_ENV_VAR, PgBouncerError, PgBouncerSettings
|
||||
from litellm.proxy.proxy_server import app as proxy_app
|
||||
|
||||
DB_ENV: Final = {
|
||||
"DATABASE_HOST": "db.internal",
|
||||
|
|
@ -107,8 +108,22 @@ class TestUvicornArgv:
|
|||
argv: Final = uvicorn_argv(("--timeout-keep-alive", "30"), {"KEEPALIVE_TIMEOUT": "75"})
|
||||
assert _uvicorn_params(argv)["timeout_keep_alive"] == 30
|
||||
|
||||
def test_the_app_uvicorn_is_told_to_serve_is_the_trimmed_gateway(self):
|
||||
assert import_from_string(cast(str, _uvicorn_params(uvicorn_argv((), {}))["app"])) is gateway.main.app
|
||||
def test_the_app_uvicorn_is_told_to_serve_is_the_trimmed_gateway(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(proxy_app.router, "lifespan_context", proxy_app.router.lifespan_context)
|
||||
for key in (
|
||||
"DATABASE_URL",
|
||||
"DIRECT_URL",
|
||||
"DATABASE_URL_READ_REPLICA",
|
||||
"DATABASE_HOST",
|
||||
"DATABASE_HOST_READ_REPLICA",
|
||||
"DATABASE_PASSWORD",
|
||||
"IAM_TOKEN_DB_AUTH",
|
||||
"AZURE_POSTGRESQL_AUTH",
|
||||
):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
served: Final = import_from_string(cast(str, _uvicorn_params(uvicorn_argv((), {}))["app"]))
|
||||
assert served is importlib.import_module("gateway.main").app
|
||||
|
||||
|
||||
class TestPoolDatabaseUrl:
|
||||
|
|
|
|||
|
|
@ -109,25 +109,6 @@ def set_salt_key(monkeypatch):
|
|||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_constants_module():
|
||||
"""Reset constants module to ensure clean state before each test"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
# Reload modules before test
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
yield
|
||||
|
||||
# Reload modules after test to clean up
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_sso_user_defined_values():
|
||||
return LiteLLM_UserTable(
|
||||
|
|
@ -875,19 +856,10 @@ def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values
|
|||
|
||||
|
||||
def test_get_cli_jwt_auth_token_custom_expiration(valid_sso_user_defined_values, monkeypatch):
|
||||
"""Test generating CLI JWT token with custom expiration via environment variable"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
"""Test generating a CLI JWT token with custom expiration via the configured constant"""
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
# Set custom expiration to 48 hours
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
||||
|
||||
# Reload the constants module to pick up the new env var
|
||||
importlib.reload(constants)
|
||||
# Also reload auth_checks to pick up the new constant value
|
||||
importlib.reload(auth_checks)
|
||||
monkeypatch.setattr(auth_checks, "CLI_JWT_EXPIRATION_HOURS", 48)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
||||
|
|
|
|||
|
|
@ -7321,15 +7321,10 @@ async def test_expired_cli_session_token_is_rejected(monkeypatch):
|
|||
on the shared validation path, not only for DB-backed keys."""
|
||||
monkeypatch.delenv("EXPERIMENTAL_UI_LOGIN", raising=False)
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-cli-test")
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "-1")
|
||||
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
monkeypatch.setattr(auth_checks, "CLI_JWT_EXPIRATION_HOURS", -1)
|
||||
|
||||
user_info = LiteLLM_UserTable(
|
||||
user_id="cli-admin",
|
||||
|
|
@ -7346,22 +7341,17 @@ async def test_expired_cli_session_token_is_rejected(monkeypatch):
|
|||
mock_request.headers = {"authorization": f"Bearer {cli_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {cli_token}",
|
||||
)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {cli_token}",
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
finally:
|
||||
monkeypatch.delenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
assert exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,52 +1,45 @@
|
|||
import os
|
||||
from typing import Final
|
||||
|
||||
import uvicorn
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Set the SERVER_ROOT_PATH environment variable to match the custom mount path
|
||||
os.environ["SERVER_ROOT_PATH"] = "/my-custom-path"
|
||||
|
||||
from litellm.proxy.proxy_server import app as litellm_app
|
||||
from litellm.proxy.proxy_server import proxy_startup_event
|
||||
|
||||
# Create main FastAPI app
|
||||
app = FastAPI(title="Custom LiteLLM Server", lifespan=proxy_startup_event)
|
||||
|
||||
# Add CORS middleware
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
custom_path = "/my-custom-path"
|
||||
|
||||
# Mount LiteLLM app at /litellm
|
||||
app.mount(custom_path, litellm_app)
|
||||
|
||||
|
||||
# Default route at /
|
||||
@app.get("/")
|
||||
async def root():
|
||||
return {
|
||||
"message": "Welcome to the API Gateway",
|
||||
"litellm_endpoint": f"{custom_path}",
|
||||
}
|
||||
def build_app() -> FastAPI:
|
||||
load_dotenv()
|
||||
os.environ["SERVER_ROOT_PATH"] = "/my-custom-path"
|
||||
|
||||
from litellm.proxy.proxy_server import app as litellm_app
|
||||
from litellm.proxy.proxy_server import proxy_startup_event
|
||||
|
||||
# Health check endpoint
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
return {"status": "healthy"}
|
||||
app: Final = FastAPI(title="Custom LiteLLM Server", lifespan=proxy_startup_event)
|
||||
custom_path: Final = "/my-custom-path"
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.mount(custom_path, litellm_app)
|
||||
|
||||
@app.get("/")
|
||||
async def root() -> dict[str, str]:
|
||||
return {
|
||||
"message": "Welcome to the API Gateway",
|
||||
"litellm_endpoint": custom_path,
|
||||
}
|
||||
|
||||
@app.get("/health")
|
||||
async def health_check() -> dict[str, str]:
|
||||
return {"status": "healthy"}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run the server on port 8000
|
||||
uvicorn.run(app, host="0.0.0.0", port=4000, log_level="info")
|
||||
uvicorn.run(build_app(), host="0.0.0.0", port=4000, log_level="info")
|
||||
|
|
|
|||
|
|
@ -68,3 +68,36 @@ def _build_constant_env_var_map() -> dict[str, str]:
|
|||
env_var_map[constant_name] = env_var_name
|
||||
|
||||
return env_var_map
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("cli_value", "litellm_cli_value", "expected"),
|
||||
[
|
||||
("48", None, 48),
|
||||
(None, "48", 48),
|
||||
(None, None, 24),
|
||||
("48", "72", 48),
|
||||
],
|
||||
ids=("canonical-only", "alias-only", "default", "canonical-wins"),
|
||||
)
|
||||
def test_cli_jwt_expiration_hours_from_environment(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
cli_value: str | None,
|
||||
litellm_cli_value: str | None,
|
||||
expected: int,
|
||||
) -> None:
|
||||
monkeypatch.delenv("CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
monkeypatch.delenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
|
||||
try:
|
||||
if cli_value is not None:
|
||||
monkeypatch.setenv("CLI_JWT_EXPIRATION_HOURS", cli_value)
|
||||
if litellm_cli_value is not None:
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", litellm_cli_value)
|
||||
|
||||
importlib.reload(litellm.constants)
|
||||
assert litellm.constants.CLI_JWT_EXPIRATION_HOURS == expected
|
||||
finally:
|
||||
monkeypatch.delenv("CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
monkeypatch.delenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
importlib.reload(litellm.constants)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue