mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(proxy): add admin-only /debug/asyncio-tasks/stacks endpoint
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
9a715df212
commit
259220e2d8
5 changed files with 191 additions and 37 deletions
|
|
@ -109,6 +109,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
"/debug/asyncio-tasks",
|
||||
"/watsonx",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ import sys
|
|||
import tracemalloc
|
||||
from collections import Counter
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, NamedTuple, Protocol, TypedDict
|
||||
from types import FrameType
|
||||
from typing import Any, Final, NamedTuple, Protocol, TypeAlias, TypedDict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from typing_extensions import ReadOnly
|
||||
|
|
@ -15,12 +16,83 @@ from typing_extensions import ReadOnly
|
|||
from litellm import get_secret_str
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PYTHON_GC_THRESHOLD
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
class _Frame(TypedDict):
|
||||
file: ReadOnly[str]
|
||||
line: ReadOnly[int]
|
||||
function: ReadOnly[str]
|
||||
|
||||
|
||||
class _TaskStackGroup(TypedDict):
|
||||
count: ReadOnly[int]
|
||||
coroutine: ReadOnly[str]
|
||||
task_names: ReadOnly[tuple[str, ...]]
|
||||
stack: ReadOnly[tuple[_Frame, ...]]
|
||||
|
||||
|
||||
class _TaskStackDump(TypedDict):
|
||||
worker_pid: ReadOnly[int]
|
||||
total_active_tasks: ReadOnly[int]
|
||||
groups: ReadOnly[tuple[_TaskStackGroup, ...]]
|
||||
|
||||
|
||||
_TaskStackKey: TypeAlias = tuple[tuple[str, int, str], ...]
|
||||
_TaskStackRecord: TypeAlias = tuple[_TaskStackKey, tuple[_Frame, ...], asyncio.Task[object]]
|
||||
|
||||
|
||||
def _frame_from_stack_frame(frame: FrameType) -> _Frame:
|
||||
result: Final[_Frame] = {
|
||||
"file": frame.f_code.co_filename,
|
||||
"line": frame.f_lineno,
|
||||
"function": frame.f_code.co_name,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def _task_stack(task: asyncio.Task[object], max_frames: int) -> tuple[_Frame, ...]:
|
||||
return tuple(_frame_from_stack_frame(frame) for frame in task.get_stack(limit=max_frames))
|
||||
|
||||
|
||||
def _task_stack_key(stack: tuple[_Frame, ...]) -> _TaskStackKey:
|
||||
return tuple((frame["file"], frame["line"], frame["function"]) for frame in stack)
|
||||
|
||||
|
||||
def _task_coroutine_name(task: asyncio.Task[object]) -> str:
|
||||
coroutine: Final = task.get_coro()
|
||||
coroutine_name: Final = getattr(coroutine, "__qualname__", None)
|
||||
return coroutine_name if isinstance(coroutine_name, str) else repr(coroutine)
|
||||
|
||||
|
||||
def _task_stack_group(stack_key: _TaskStackKey, records: tuple[_TaskStackRecord, ...]) -> _TaskStackGroup:
|
||||
matching_records: Final = tuple(record for record in records if record[0] == stack_key)
|
||||
sample_record: Final = matching_records[0]
|
||||
sample_tasks: Final = tuple(record[2] for record in matching_records)
|
||||
result: Final[_TaskStackGroup] = {
|
||||
"count": len(matching_records),
|
||||
"coroutine": _task_coroutine_name(sample_tasks[0]),
|
||||
"task_names": tuple(task.get_name() for task in sample_tasks[:5]),
|
||||
"stack": sample_record[1],
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def _group_task_stacks(tasks: tuple[asyncio.Task[object], ...], max_frames: int) -> tuple[_TaskStackGroup, ...]:
|
||||
records: Final = tuple(
|
||||
(stack_key, stack, task)
|
||||
for task in tasks
|
||||
for stack in (_task_stack(task, max_frames),)
|
||||
for stack_key in (_task_stack_key(stack),)
|
||||
)
|
||||
stack_keys: Final = tuple(dict.fromkeys(record[0] for record in records))
|
||||
groups: Final = tuple(_task_stack_group(stack_key, records) for stack_key in stack_keys)
|
||||
return tuple(sorted(groups, key=lambda group: group["count"], reverse=True))
|
||||
|
||||
|
||||
# Configure garbage collection thresholds from environment variables
|
||||
def configure_gc_thresholds():
|
||||
"""Configure Python garbage collection thresholds from environment variables."""
|
||||
|
|
@ -87,6 +159,27 @@ async def get_active_tasks_stats():
|
|||
}
|
||||
|
||||
|
||||
@router.get("/debug/asyncio-tasks/stacks", include_in_schema=False)
|
||||
async def get_active_task_stacks(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
max_frames: int = Query(default=40, ge=1, le=200),
|
||||
) -> _TaskStackDump:
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(status_code=403, detail="Only proxy admins can read asyncio task stacks")
|
||||
|
||||
max_tasks_to_check: Final = 5000
|
||||
active_tasks: Final = tuple(task for task in asyncio.all_tasks() if not task.done())
|
||||
current_task: Final = asyncio.current_task()
|
||||
tasks: Final = tuple(task for task in active_tasks if task is not current_task)[:max_tasks_to_check]
|
||||
groups: Final = _group_task_stacks(tasks, max_frames)
|
||||
result: Final[_TaskStackDump] = {
|
||||
"worker_pid": os.getpid(),
|
||||
"total_active_tasks": len(active_tasks),
|
||||
"groups": groups,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
if os.environ.get("LITELLM_PROFILE", "false").lower() == "true":
|
||||
try:
|
||||
import objgraph
|
||||
|
|
|
|||
68
tests/test_litellm/proxy/common_utils/test_debug_utils.py
Normal file
68
tests/test_litellm/proxy/common_utils/test_debug_utils.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
import asyncio
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.debug_utils import router as debug_router
|
||||
|
||||
|
||||
async def _park_for_test() -> None:
|
||||
await asyncio.sleep(30)
|
||||
|
||||
|
||||
def _test_app(user_role: LitellmUserRoles) -> FastAPI:
|
||||
app = FastAPI()
|
||||
app.include_router(debug_router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=user_role)
|
||||
return app
|
||||
|
||||
|
||||
def _client(app: FastAPI) -> httpx.AsyncClient:
|
||||
return httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_task_stacks_require_proxy_admin() -> None:
|
||||
app = _test_app(LitellmUserRoles.INTERNAL_USER)
|
||||
async with _client(app) as client:
|
||||
response = await client.get("/debug/asyncio-tasks/stacks")
|
||||
|
||||
assert response.status_code == 403
|
||||
assert response.json()["detail"] == "Only proxy admins can read asyncio task stacks"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_task_stacks_include_parked_task() -> None:
|
||||
app = _test_app(LitellmUserRoles.PROXY_ADMIN)
|
||||
parked_task = asyncio.create_task(_park_for_test(), name="parked-test-task")
|
||||
try:
|
||||
async with _client(app) as client:
|
||||
response = await client.get("/debug/asyncio-tasks/stacks")
|
||||
finally:
|
||||
parked_task.cancel()
|
||||
await asyncio.gather(parked_task, return_exceptions=True)
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
parked_group = next(group for group in body["groups"] if "_park_for_test" in group["coroutine"])
|
||||
assert any(frame["function"] == "_park_for_test" for frame in parked_group["stack"])
|
||||
assert any(frame["file"].endswith("test_debug_utils.py") for frame in parked_group["stack"])
|
||||
assert body["total_active_tasks"] >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_task_stacks_respect_max_frames() -> None:
|
||||
app = _test_app(LitellmUserRoles.PROXY_ADMIN)
|
||||
parked_task = asyncio.create_task(_park_for_test(), name="parked-test-task")
|
||||
try:
|
||||
async with _client(app) as client:
|
||||
response = await client.get("/debug/asyncio-tasks/stacks?max_frames=1")
|
||||
finally:
|
||||
parked_task.cancel()
|
||||
await asyncio.gather(parked_task, return_exceptions=True)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert all(len(group["stack"]) <= 1 for group in response.json()["groups"])
|
||||
|
|
@ -108,12 +108,8 @@ def test_gateway_plus_backend_covers_full_app():
|
|||
for r in app.router.routes
|
||||
if not isinstance(r, Mount) and getattr(r, "path", None) is not None
|
||||
}
|
||||
gateway_paths = _component_paths(
|
||||
app.router.routes, GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES
|
||||
)
|
||||
backend_paths = _component_paths(
|
||||
app.router.routes, BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES
|
||||
)
|
||||
gateway_paths = _component_paths(app.router.routes, GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES)
|
||||
backend_paths = _component_paths(app.router.routes, BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES)
|
||||
|
||||
uncovered = all_paths - (gateway_paths | backend_paths)
|
||||
|
||||
|
|
@ -126,16 +122,15 @@ def test_gateway_plus_backend_covers_full_app():
|
|||
|
||||
def test_backend_mount_paths_defined():
|
||||
"""BACKEND_MOUNT_PATHS constant must exist and be a frozenset."""
|
||||
assert isinstance(BACKEND_MOUNT_PATHS, frozenset), \
|
||||
assert isinstance(BACKEND_MOUNT_PATHS, frozenset), (
|
||||
f"BACKEND_MOUNT_PATHS must be a frozenset, got {type(BACKEND_MOUNT_PATHS)}"
|
||||
assert len(BACKEND_MOUNT_PATHS) > 0, \
|
||||
"BACKEND_MOUNT_PATHS must contain at least one Mount path"
|
||||
)
|
||||
assert len(BACKEND_MOUNT_PATHS) > 0, "BACKEND_MOUNT_PATHS must contain at least one Mount path"
|
||||
|
||||
|
||||
def test_swagger_mount_in_backend_allowlist():
|
||||
"""The /swagger Mount must be in BACKEND_MOUNT_PATHS."""
|
||||
assert "/swagger" in BACKEND_MOUNT_PATHS, \
|
||||
"/swagger Mount path must be in BACKEND_MOUNT_PATHS"
|
||||
assert "/swagger" in BACKEND_MOUNT_PATHS, "/swagger Mount path must be in BACKEND_MOUNT_PATHS"
|
||||
|
||||
|
||||
def test_backend_keeps_swagger_mount():
|
||||
|
|
@ -145,32 +140,31 @@ def test_backend_keeps_swagger_mount():
|
|||
for r in app.router.routes
|
||||
if isinstance(r, Mount) and getattr(r, "path", None) in BACKEND_MOUNT_PATHS
|
||||
}
|
||||
assert "/swagger" in backend_mounts, \
|
||||
assert "/swagger" in backend_mounts, (
|
||||
"/swagger Mount is expected on the proxy app and should be in BACKEND_MOUNT_PATHS"
|
||||
)
|
||||
|
||||
|
||||
def test_backend_drops_non_allowlisted_mounts():
|
||||
"""Verify that Mounts NOT in BACKEND_MOUNT_PATHS would be dropped from backend."""
|
||||
all_mounts = {
|
||||
getattr(r, "path")
|
||||
for r in app.router.routes
|
||||
if isinstance(r, Mount) and getattr(r, "path", None) is not None
|
||||
getattr(r, "path") for r in app.router.routes if isinstance(r, Mount) and getattr(r, "path", None) is not None
|
||||
}
|
||||
non_backend_mounts = all_mounts - BACKEND_MOUNT_PATHS
|
||||
|
||||
assert len(non_backend_mounts) > 0, \
|
||||
assert len(non_backend_mounts) > 0, (
|
||||
"Expected at least one non-backend Mount (e.g., /ui, /_next) to verify filtering logic"
|
||||
)
|
||||
for mount_path in non_backend_mounts:
|
||||
assert mount_path not in BACKEND_MOUNT_PATHS, \
|
||||
f"Mount {mount_path} should not be in BACKEND_MOUNT_PATHS"
|
||||
assert mount_path not in BACKEND_MOUNT_PATHS, f"Mount {mount_path} should not be in BACKEND_MOUNT_PATHS"
|
||||
|
||||
|
||||
def test_gateway_mount_paths_defined():
|
||||
"""GATEWAY_MOUNT_PATHS constant must exist and expose /metrics."""
|
||||
assert isinstance(GATEWAY_MOUNT_PATHS, frozenset), \
|
||||
assert isinstance(GATEWAY_MOUNT_PATHS, frozenset), (
|
||||
f"GATEWAY_MOUNT_PATHS must be a frozenset, got {type(GATEWAY_MOUNT_PATHS)}"
|
||||
assert "/metrics" in GATEWAY_MOUNT_PATHS, \
|
||||
"/metrics Mount path must be in GATEWAY_MOUNT_PATHS"
|
||||
)
|
||||
assert "/metrics" in GATEWAY_MOUNT_PATHS, "/metrics Mount path must be in GATEWAY_MOUNT_PATHS"
|
||||
|
||||
|
||||
def test_gateway_trim_keeps_metrics_mount():
|
||||
|
|
@ -185,15 +179,20 @@ def test_gateway_trim_keeps_metrics_mount():
|
|||
metrics_mount = Mount("/metrics", app=make_asgi_app())
|
||||
routes = [*app.router.routes, metrics_mount]
|
||||
trimmed = [r for r in routes if _is_gateway_route(r)]
|
||||
assert metrics_mount in trimmed, \
|
||||
"/metrics Mount must survive the gateway route trim"
|
||||
assert metrics_mount in trimmed, "/metrics Mount must survive the gateway route trim"
|
||||
|
||||
|
||||
def test_gateway_drops_ui_and_swagger_mounts():
|
||||
"""UI static and swagger Mounts must still be trimmed from the gateway."""
|
||||
for path in ("/ui", "/_next", "/litellm-asset-prefix/_next", "/swagger"):
|
||||
assert not _is_gateway_route(Mount(path, app=make_asgi_app())), \
|
||||
assert not _is_gateway_route(Mount(path, app=make_asgi_app())), (
|
||||
f"Mount {path} must not be served by the gateway"
|
||||
)
|
||||
|
||||
|
||||
def test_gateway_keeps_asyncio_task_stacks_route():
|
||||
route = next(route for route in app.router.routes if getattr(route, "path", None) == "/debug/asyncio-tasks/stacks")
|
||||
assert _is_gateway_route(route)
|
||||
|
||||
|
||||
def test_every_app_mount_is_assigned_to_a_component():
|
||||
|
|
|
|||
|
|
@ -9,11 +9,7 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
|||
|
||||
def _get_route_dependency_calls(router, path: str, method: str):
|
||||
for route in router.routes:
|
||||
if (
|
||||
isinstance(route, APIRoute)
|
||||
and route.path == path
|
||||
and method in route.methods
|
||||
):
|
||||
if isinstance(route, APIRoute) and route.path == path and method in route.methods:
|
||||
return [dependency.call for dependency in route.dependant.dependencies]
|
||||
raise AssertionError(f"Route {method} {path} not found")
|
||||
|
||||
|
|
@ -21,14 +17,11 @@ def _get_route_dependency_calls(router, path: str, method: str):
|
|||
def test_sensitive_debug_routes_require_auth_dependency():
|
||||
for path, method in (
|
||||
("/debug/asyncio-tasks", "GET"),
|
||||
("/debug/asyncio-tasks/stacks", "GET"),
|
||||
("/otel-spans", "GET"),
|
||||
):
|
||||
assert user_api_key_auth in _get_route_dependency_calls(
|
||||
debug_router, path, method
|
||||
)
|
||||
assert user_api_key_auth in _get_route_dependency_calls(debug_router, path, method)
|
||||
|
||||
|
||||
def test_provider_budgets_requires_auth_dependency():
|
||||
assert user_api_key_auth in _get_route_dependency_calls(
|
||||
spend_router, "/provider/budgets", "GET"
|
||||
)
|
||||
assert user_api_key_auth in _get_route_dependency_calls(spend_router, "/provider/budgets", "GET")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue