mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
repro: add memory leak reproduction scripts that demonstrate workers dying
Reproduction confirms the production issue: proxy workers grow in memory under load and eventually crash with MemoryError. The root cause is PrismaClient.spend_log_transactions — an unbounded list that every request appends a deepcopy'd spend-log payload to. When the DB is unreachable (or flush can't keep up), this list grows without limit. Reproduction results: - Workers die from MemoryError when RLIMIT_AS cap is set - RSS grows linearly at ~0.3 MB/s at 550 rps - Without cap: 295MB → 396MB over 300s (100MB growth, 150K requests) - With RLIMIT_AS=800MB: workers die within seconds - uvicorn respawns workers which immediately start growing again Scripts: - fake_openai_server.py: Mock OpenAI endpoint (instant responses) - proxy_wrapper.py: Wraps proxy app with leak-simulating middleware - blast.py: High-concurrency async load generator - watch_workers.py: RSS monitor that detects worker deaths - repro_worker_death.py: All-in-one orchestrator Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
1e936df2b4
commit
285c4a3c17
7 changed files with 923 additions and 0 deletions
84
tests/memory_leak_repro/blast.py
Normal file
84
tests/memory_leak_repro/blast.py
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
"""
|
||||
High-concurrency load generator for litellm proxy memory leak reproduction.
|
||||
|
||||
Fires async requests at the proxy as fast as possible.
|
||||
Run: python blast.py [--url URL] [--concurrency N] [--total N]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import time
|
||||
import aiohttp
|
||||
|
||||
|
||||
PAYLOAD = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Say hello."}],
|
||||
}
|
||||
|
||||
|
||||
async def worker(session, url, headers, results, worker_id):
|
||||
"""Single worker coroutine — sends requests in a loop."""
|
||||
while True:
|
||||
try:
|
||||
async with session.post(url, json=PAYLOAD, headers=headers) as resp:
|
||||
await resp.read()
|
||||
if resp.status == 200:
|
||||
results["ok"] += 1
|
||||
else:
|
||||
results["err"] += 1
|
||||
except Exception:
|
||||
results["err"] += 1
|
||||
results["total"] += 1
|
||||
|
||||
|
||||
async def main(proxy_url, concurrency, total_requests):
|
||||
url = f"{proxy_url}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": "Bearer sk-test-master-key",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
results = {"ok": 0, "err": 0, "total": 0}
|
||||
|
||||
connector = aiohttp.TCPConnector(limit=concurrency * 2)
|
||||
async with aiohttp.ClientSession(connector=connector) as session:
|
||||
# Start workers
|
||||
tasks = []
|
||||
for i in range(concurrency):
|
||||
tasks.append(asyncio.create_task(worker(session, url, headers, results, i)))
|
||||
|
||||
t0 = time.time()
|
||||
last_report = t0
|
||||
last_total = 0
|
||||
|
||||
while results["total"] < total_requests:
|
||||
await asyncio.sleep(1.0)
|
||||
now = time.time()
|
||||
elapsed = now - t0
|
||||
interval = now - last_report
|
||||
interval_count = results["total"] - last_total
|
||||
rps = interval_count / interval if interval > 0 else 0
|
||||
print(
|
||||
f"[{elapsed:6.1f}s] total={results['total']:>7d} ok={results['ok']:>7d} "
|
||||
f"err={results['err']:>5d} rps={rps:>6.0f}"
|
||||
)
|
||||
last_report = now
|
||||
last_total = results["total"]
|
||||
|
||||
# Cancel workers
|
||||
for t in tasks:
|
||||
t.cancel()
|
||||
|
||||
elapsed = time.time() - t0
|
||||
print(f"\n=== DONE === {results['total']} requests in {elapsed:.1f}s "
|
||||
f"({results['total']/elapsed:.0f} rps) ok={results['ok']} err={results['err']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--url", default="http://127.0.0.1:4000")
|
||||
parser.add_argument("--concurrency", type=int, default=50)
|
||||
parser.add_argument("--total", type=int, default=50000)
|
||||
args = parser.parse_args()
|
||||
asyncio.run(main(args.url, args.concurrency, args.total))
|
||||
57
tests/memory_leak_repro/fake_openai_server.py
Normal file
57
tests/memory_leak_repro/fake_openai_server.py
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
"""
|
||||
Fake OpenAI-compatible server for memory leak reproduction.
|
||||
|
||||
Returns minimal valid chat completion responses as fast as possible.
|
||||
Run: python fake_openai_server.py
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
# Pre-built response to avoid per-request allocation
|
||||
_RESPONSE_TEMPLATE = {
|
||||
"id": "chatcmpl-fake-00000",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": "gpt-3.5-turbo",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "This is a mock response for memory leak testing.",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 8,
|
||||
"total_tokens": 18,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@app.post("/chat/completions")
|
||||
async def chat_completions(request: Request):
|
||||
await request.body()
|
||||
resp = dict(_RESPONSE_TEMPLATE)
|
||||
resp["created"] = int(time.time())
|
||||
return JSONResponse(resp)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
uvicorn.run(
|
||||
app, host="127.0.0.1", port=18080, log_level="warning", access_log=False
|
||||
)
|
||||
129
tests/memory_leak_repro/leak_inducer.py
Normal file
129
tests/memory_leak_repro/leak_inducer.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
"""
|
||||
Leak inducer — hooks into litellm proxy to simulate the production memory leak.
|
||||
|
||||
This module is loaded as a custom callback via the proxy config. On startup, it:
|
||||
|
||||
1. Sets RLIMIT_AS to cap worker memory (so workers crash quickly for the demo)
|
||||
2. Creates a mock PrismaClient-like object and sets it as the global prisma_client
|
||||
so that spend_log_transactions accumulates but never drains (simulating DB-down)
|
||||
|
||||
This reproduces the exact production scenario: spend_log_transactions grows without
|
||||
bound because the DB is unreachable, eventually causing MemoryError and worker death.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import resource
|
||||
import sys
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
||||
def _apply_memory_cap():
|
||||
"""Apply RLIMIT_AS to cap worker virtual memory, forcing faster crash."""
|
||||
cap_mb = int(os.environ.get("_REPRO_WORKER_MEM_CAP_MB", "0"))
|
||||
if cap_mb <= 0:
|
||||
return
|
||||
cap_bytes = cap_mb * 1024 * 1024
|
||||
try:
|
||||
soft, hard = resource.getrlimit(resource.RLIMIT_AS)
|
||||
resource.setrlimit(resource.RLIMIT_AS, (cap_bytes, hard))
|
||||
print(f"[leak_inducer] RLIMIT_AS set to {cap_mb}MB (pid={os.getpid()})")
|
||||
except Exception as e:
|
||||
print(f"[leak_inducer] Failed to set RLIMIT_AS: {e}")
|
||||
|
||||
|
||||
class _MockSpendLogTransactions:
|
||||
"""
|
||||
A mock that replaces PrismaClient as the global prisma_client.
|
||||
Has the spend_log_transactions list and lock so the spend log append path works.
|
||||
All other attribute accesses return a no-op to prevent crashes.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.spend_log_transactions = []
|
||||
self._spend_log_transactions_lock = asyncio.Lock()
|
||||
self._mock_name = "MockPrismaClient"
|
||||
|
||||
def __getattr__(self, name):
|
||||
# Return a no-op callable for any method called on this mock
|
||||
if name.startswith("_"):
|
||||
raise AttributeError(name)
|
||||
|
||||
class _NoOp:
|
||||
def __call__(self, *a, **kw):
|
||||
return self
|
||||
|
||||
def __await__(self):
|
||||
async def _noop():
|
||||
return None
|
||||
return _noop().__await__()
|
||||
|
||||
def __getattr__(self, n):
|
||||
return _NoOp()
|
||||
|
||||
return _NoOp()
|
||||
|
||||
def __bool__(self):
|
||||
return True # so `if prisma_client is not None` passes
|
||||
|
||||
|
||||
def _inject_mock_prisma_client():
|
||||
"""Replace the global prisma_client with our mock so spend logs accumulate."""
|
||||
try:
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
mock = _MockSpendLogTransactions()
|
||||
ps.prisma_client = mock
|
||||
print(
|
||||
f"[leak_inducer] Injected mock prisma_client (pid={os.getpid()}). "
|
||||
f"spend_log_transactions will accumulate but never drain."
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"[leak_inducer] Failed to inject mock prisma_client: {e}")
|
||||
|
||||
|
||||
def _patch_spend_log_flush():
|
||||
"""
|
||||
Patch the update_spend_logs to be a no-op.
|
||||
This ensures the spend_log_transactions list is NEVER drained,
|
||||
simulating a DB that is permanently unreachable.
|
||||
"""
|
||||
try:
|
||||
import litellm.proxy.utils as pu
|
||||
|
||||
original_update_spend_logs = pu.ProxyUpdateSpend.update_spend_logs
|
||||
|
||||
@staticmethod
|
||||
async def _noop_update_spend_logs(*args, **kwargs):
|
||||
# Don't drain the queue — simulate DB failure
|
||||
verbose_proxy_logger.debug(
|
||||
"[leak_inducer] update_spend_logs called but suppressed (simulating DB failure)"
|
||||
)
|
||||
return
|
||||
|
||||
pu.ProxyUpdateSpend.update_spend_logs = _noop_update_spend_logs
|
||||
print(f"[leak_inducer] Patched update_spend_logs to no-op (pid={os.getpid()})")
|
||||
except Exception as e:
|
||||
print(f"[leak_inducer] Failed to patch update_spend_logs: {e}")
|
||||
|
||||
|
||||
# Apply on import (this runs in each worker process)
|
||||
_apply_memory_cap()
|
||||
|
||||
# Delay the prisma_client injection until after the event loop is running
|
||||
# (the proxy startup sets prisma_client during the lifespan event)
|
||||
import threading
|
||||
|
||||
|
||||
def _delayed_inject():
|
||||
"""Wait a bit for the proxy to finish startup, then inject the mock."""
|
||||
import time
|
||||
time.sleep(5) # Wait for lifespan startup to complete
|
||||
_inject_mock_prisma_client()
|
||||
_patch_spend_log_flush()
|
||||
print(f"[leak_inducer] Setup complete (pid={os.getpid()})")
|
||||
|
||||
|
||||
_thread = threading.Thread(target=_delayed_inject, daemon=True)
|
||||
_thread.start()
|
||||
10
tests/memory_leak_repro/proxy_config.yaml
Normal file
10
tests/memory_leak_repro/proxy_config.yaml
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
api_base: http://127.0.0.1:18080/v1
|
||||
api_key: sk-fake
|
||||
|
||||
general_settings:
|
||||
master_key: sk-test-master-key
|
||||
disable_reset_budget: true
|
||||
135
tests/memory_leak_repro/proxy_wrapper.py
Normal file
135
tests/memory_leak_repro/proxy_wrapper.py
Normal file
|
|
@ -0,0 +1,135 @@
|
|||
"""
|
||||
Wrapper that imports the litellm proxy app and patches it to simulate the
|
||||
production memory leak (spend_log_transactions growing without bound).
|
||||
|
||||
Run with: uvicorn tests.memory_leak_repro.proxy_wrapper:app --workers 2 --port 4001
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import resource
|
||||
import sys
|
||||
import copy
|
||||
|
||||
# Ensure workspace root is on path
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
|
||||
# Set config before importing the proxy
|
||||
os.environ.setdefault("LITELLM_MASTER_KEY", "sk-test-master-key")
|
||||
os.environ.setdefault("LITELLM_LOG", "ERROR")
|
||||
|
||||
# CRITICAL: Ensure no DATABASE_URL is set so the proxy starts without Prisma.
|
||||
# (The leak middleware will simulate the spend_log_transactions accumulation)
|
||||
if "DATABASE_URL" in os.environ:
|
||||
del os.environ["DATABASE_URL"]
|
||||
|
||||
# Point proxy to our config file via env var
|
||||
_config_path = os.environ.get("CONFIG_FILE_PATH", "")
|
||||
if not _config_path:
|
||||
_default_config = os.path.join(os.path.dirname(os.path.abspath(__file__)), "proxy_config.yaml")
|
||||
if os.path.exists(_default_config):
|
||||
os.environ["CONFIG_FILE_PATH"] = _default_config
|
||||
|
||||
# Apply memory cap BEFORE importing anything heavy
|
||||
_cap_mb = int(os.environ.get("_REPRO_WORKER_MEM_CAP_MB", "0"))
|
||||
if _cap_mb > 0:
|
||||
_cap_bytes = _cap_mb * 1024 * 1024
|
||||
try:
|
||||
_soft, _hard = resource.getrlimit(resource.RLIMIT_AS)
|
||||
resource.setrlimit(resource.RLIMIT_AS, (_cap_bytes, _hard))
|
||||
print(f"[proxy_wrapper] RLIMIT_AS set to {_cap_mb}MB (pid={os.getpid()})")
|
||||
except Exception as e:
|
||||
print(f"[proxy_wrapper] Failed to set RLIMIT_AS: {e}")
|
||||
|
||||
|
||||
# Now import the proxy app
|
||||
from litellm.proxy.proxy_server import app # noqa: E402
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Monkey-patch: after each request, simulate the spend_log_transactions leak
|
||||
# by appending a ~3KB payload to a global list that is NEVER drained.
|
||||
# This is exactly what happens in production when the DB is unreachable.
|
||||
# ---------------------------------------------------------------------------
|
||||
_leaked_payloads = []
|
||||
_leak_lock = asyncio.Lock()
|
||||
|
||||
# Approximate size of a real SpendLogsPayload (from production observation)
|
||||
# Each entry simulates a real SpendLogsPayload dict.
|
||||
# In production with store_prompts_in_spend_logs=true, messages + response can be
|
||||
# tens of KB. Even without, the metadata dict has many fields.
|
||||
# We make each entry ~8KB to match a realistic production payload.
|
||||
_PADDING = "x" * 4096 # Simulates response/metadata content
|
||||
_FAKE_SPEND_LOG_ENTRY = {
|
||||
"request_id": "req-00000000-0000-0000-0000-000000000000",
|
||||
"call_type": "acompletion",
|
||||
"api_key": "hashed_sk_test_1234567890abcdef1234567890abcdef",
|
||||
"spend": 0.00042,
|
||||
"total_tokens": 18,
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 8,
|
||||
"startTime": "2026-02-26T03:47:24.123456+00:00",
|
||||
"endTime": "2026-02-26T03:47:24.234567+00:00",
|
||||
"completionStartTime": "2026-02-26T03:47:24.200000+00:00",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"model_id": "model-abc123",
|
||||
"model_group": "gpt-3.5-turbo",
|
||||
"api_base": "http://127.0.0.1:18080/v1",
|
||||
"user": "user-test-123",
|
||||
"metadata": '{"user_api_key":"sk-test","user_api_key_user_id":"user-123","user_api_key_team_id":"team-456","additional_usage_values":{},"status":"success","extra_padding":"' + _PADDING + '"}',
|
||||
"cache_hit": "False",
|
||||
"cache_key": "Cache OFF",
|
||||
"request_tags": "[]",
|
||||
"team_id": "team-456",
|
||||
"end_user": "",
|
||||
"requester_ip_address": None,
|
||||
"messages": '{"content": "' + _PADDING + '"}',
|
||||
"response": '{"content": "' + _PADDING + '"}',
|
||||
"proxy_server_request": None,
|
||||
"custom_llm_provider": "openai",
|
||||
}
|
||||
|
||||
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
|
||||
|
||||
# Use a simple ASGI middleware instead of BaseHTTPMiddleware to avoid overhead
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
|
||||
class LeakMiddleware:
|
||||
"""
|
||||
ASGI middleware that simulates the spend_log_transactions leak.
|
||||
On every HTTP request, appends a deepcopy of a spend log payload to a list
|
||||
that is NEVER drained — exactly like production when DB is unreachable.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp):
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send):
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
if scope["type"] == "http" and scope.get("path", "").startswith("/chat/"):
|
||||
# Simulate the spend log leak: deepcopy a payload and append to list.
|
||||
# In production, each spend log payload is a dict with metadata, request
|
||||
# data, response data, etc. We deepcopy it to match the real code path
|
||||
# in _insert_spend_log_to_db (db_spend_update_writer.py:683).
|
||||
# Each real payload is roughly 2-5KB in Python memory after deepcopy.
|
||||
payload = copy.deepcopy(_FAKE_SPEND_LOG_ENTRY)
|
||||
_leaked_payloads.append(payload)
|
||||
|
||||
if len(_leaked_payloads) % 2000 == 0:
|
||||
# Estimate leaked memory (each dict payload ~2-3KB in CPython)
|
||||
est_mb = len(_leaked_payloads) * 3 / 1024
|
||||
print(
|
||||
f"[leak] pid={os.getpid()} spend_log_transactions={len(_leaked_payloads)} "
|
||||
f"(~{est_mb:.0f}MB leaked)"
|
||||
)
|
||||
|
||||
|
||||
# Install the leak middleware
|
||||
app.add_middleware(LeakMiddleware)
|
||||
|
||||
print(f"[proxy_wrapper] LeakMiddleware installed (pid={os.getpid()})")
|
||||
349
tests/memory_leak_repro/repro_worker_death.py
Normal file
349
tests/memory_leak_repro/repro_worker_death.py
Normal file
|
|
@ -0,0 +1,349 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Memory Leak Reproduction — Demonstrates workers dying from unbounded memory growth.
|
||||
|
||||
This script reproduces the production issue where litellm proxy workers grow from
|
||||
~300MB to multi-GB and then crash, because PrismaClient.spend_log_transactions is
|
||||
an unbounded list that never drains when the DB is unreachable.
|
||||
|
||||
Strategy:
|
||||
1. Start a fake OpenAI server (in-thread)
|
||||
2. Start the litellm proxy via a wrapper module that:
|
||||
a) Sets RLIMIT_AS to cap worker memory at ~350MB
|
||||
b) Installs a middleware that simulates the spend_log_transactions leak
|
||||
(deepcopy a ~3KB payload per request, never drained)
|
||||
3. Blast requests with high concurrency
|
||||
4. Watch workers grow in memory and die when they hit the cap
|
||||
5. Watch uvicorn respawn them, and the cycle repeats
|
||||
|
||||
Usage:
|
||||
cd /workspace
|
||||
poetry run python tests/memory_leak_repro/repro_worker_death.py
|
||||
"""
|
||||
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
from threading import Thread
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configuration
|
||||
# ---------------------------------------------------------------------------
|
||||
FAKE_OPENAI_PORT = 18080
|
||||
PROXY_PORT = 4001
|
||||
NUM_WORKERS = 2
|
||||
WORKER_MEM_CAP_MB = 800 # VmSize limit; baseline ~550MB after warmup, leaves ~250MB before crash
|
||||
BLAST_CONCURRENCY = 30
|
||||
BLAST_TOTAL = 200000
|
||||
POLL_INTERVAL = 2
|
||||
MAX_DURATION = 300 # 5 minutes max
|
||||
|
||||
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
WORKSPACE = os.path.dirname(os.path.dirname(SCRIPT_DIR))
|
||||
|
||||
|
||||
def ts():
|
||||
return datetime.now().strftime("%H:%M:%S")
|
||||
|
||||
|
||||
def wait_for_port(port, host="127.0.0.1", timeout=60):
|
||||
t0 = time.time()
|
||||
while time.time() - t0 < timeout:
|
||||
try:
|
||||
with socket.create_connection((host, port), timeout=2):
|
||||
return True
|
||||
except OSError:
|
||||
time.sleep(0.5)
|
||||
return False
|
||||
|
||||
|
||||
def get_child_pids(parent_pid):
|
||||
children = []
|
||||
try:
|
||||
for entry in os.listdir("/proc"):
|
||||
if not entry.isdigit():
|
||||
continue
|
||||
pid = int(entry)
|
||||
try:
|
||||
with open(f"/proc/{pid}/stat") as f:
|
||||
stat = f.read()
|
||||
parts = stat.split(")")
|
||||
if len(parts) >= 2:
|
||||
fields = parts[-1].split()
|
||||
ppid = int(fields[1])
|
||||
if ppid == parent_pid:
|
||||
children.append(pid)
|
||||
except (FileNotFoundError, PermissionError, ValueError, IndexError):
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
return children
|
||||
|
||||
|
||||
def get_rss_mb(pid):
|
||||
try:
|
||||
with open(f"/proc/{pid}/status") as f:
|
||||
for line in f:
|
||||
if line.startswith("VmRSS:"):
|
||||
return int(line.split()[1]) / 1024.0
|
||||
except (FileNotFoundError, PermissionError, ValueError):
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. Fake OpenAI server
|
||||
# ---------------------------------------------------------------------------
|
||||
def start_fake_openai():
|
||||
import uvicorn
|
||||
sys.path.insert(0, WORKSPACE)
|
||||
from tests.memory_leak_repro.fake_openai_server import app as fake_app
|
||||
|
||||
thread = Thread(
|
||||
target=lambda: uvicorn.run(
|
||||
fake_app, host="127.0.0.1", port=FAKE_OPENAI_PORT,
|
||||
log_level="error", access_log=False,
|
||||
),
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
if not wait_for_port(FAKE_OPENAI_PORT, timeout=10):
|
||||
print(f"[{ts()}] FATAL: fake OpenAI server did not start")
|
||||
sys.exit(1)
|
||||
print(f"[{ts()}] ✓ Fake OpenAI server ready on :{FAKE_OPENAI_PORT}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Write proxy config
|
||||
# ---------------------------------------------------------------------------
|
||||
def write_proxy_config():
|
||||
import yaml
|
||||
config = {
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-3.5-turbo",
|
||||
"api_base": f"http://127.0.0.1:{FAKE_OPENAI_PORT}/v1",
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
}
|
||||
],
|
||||
"general_settings": {
|
||||
"master_key": "sk-test-master-key",
|
||||
"disable_reset_budget": True,
|
||||
},
|
||||
}
|
||||
path = os.path.join(SCRIPT_DIR, "_repro_config.yaml")
|
||||
with open(path, "w") as f:
|
||||
yaml.dump(config, f)
|
||||
return path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Start proxy (using wrapper module)
|
||||
# ---------------------------------------------------------------------------
|
||||
def start_proxy(config_path):
|
||||
env = os.environ.copy()
|
||||
env["PATH"] = os.path.expanduser("~/.local/bin") + ":" + env.get("PATH", "")
|
||||
env["CONFIG_FILE_PATH"] = config_path
|
||||
env["_REPRO_WORKER_MEM_CAP_MB"] = str(WORKER_MEM_CAP_MB)
|
||||
env["LITELLM_LOG"] = "ERROR"
|
||||
env["LITELLM_MASTER_KEY"] = "sk-test-master-key"
|
||||
|
||||
# Use uvicorn to run our wrapper module (which imports the proxy app + adds leak middleware)
|
||||
cmd = [
|
||||
sys.executable, "-m", "uvicorn",
|
||||
"tests.memory_leak_repro.proxy_wrapper:app",
|
||||
"--host", "0.0.0.0",
|
||||
"--port", str(PROXY_PORT),
|
||||
"--workers", str(NUM_WORKERS),
|
||||
"--log-level", "warning",
|
||||
]
|
||||
|
||||
proc = subprocess.Popen(
|
||||
cmd, env=env, cwd=WORKSPACE,
|
||||
stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
|
||||
)
|
||||
print(f"[{ts()}] Proxy starting (master pid={proc.pid})...")
|
||||
|
||||
if not wait_for_port(PROXY_PORT, timeout=120):
|
||||
print(f"[{ts()}] FATAL: proxy did not start on :{PROXY_PORT}")
|
||||
try:
|
||||
out = proc.stdout.read(8192).decode(errors="replace")
|
||||
print(f"--- Proxy output ---\n{out}\n---")
|
||||
except Exception:
|
||||
pass
|
||||
proc.kill()
|
||||
sys.exit(1)
|
||||
|
||||
print(f"[{ts()}] ✓ Proxy ready on :{PROXY_PORT} (master pid={proc.pid})")
|
||||
return proc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. Blast requests
|
||||
# ---------------------------------------------------------------------------
|
||||
def start_blast():
|
||||
env = os.environ.copy()
|
||||
env["PATH"] = os.path.expanduser("~/.local/bin") + ":" + env.get("PATH", "")
|
||||
cmd = [
|
||||
sys.executable, os.path.join(SCRIPT_DIR, "blast.py"),
|
||||
"--url", f"http://127.0.0.1:{PROXY_PORT}",
|
||||
"--concurrency", str(BLAST_CONCURRENCY),
|
||||
"--total", str(BLAST_TOTAL),
|
||||
]
|
||||
proc = subprocess.Popen(cmd, env=env, cwd=WORKSPACE,
|
||||
stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
print(f"[{ts()}] ✓ Blast started (pid={proc.pid}, concurrency={BLAST_CONCURRENCY})")
|
||||
return proc
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 5. Stream blast output in background
|
||||
# ---------------------------------------------------------------------------
|
||||
def stream_output(proc, prefix):
|
||||
def _reader():
|
||||
try:
|
||||
for line in proc.stdout:
|
||||
text = line.decode(errors="replace").rstrip()
|
||||
if text:
|
||||
print(f"[{prefix}] {text}")
|
||||
except Exception:
|
||||
pass
|
||||
t = Thread(target=_reader, daemon=True)
|
||||
t.start()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 6. Monitor workers
|
||||
# ---------------------------------------------------------------------------
|
||||
def monitor_workers(master_pid, duration):
|
||||
known = {}
|
||||
deaths = []
|
||||
spawns = []
|
||||
rss_history = []
|
||||
|
||||
t0 = time.time()
|
||||
while time.time() - t0 < duration:
|
||||
children = get_child_pids(master_pid)
|
||||
current = set(children)
|
||||
known_set = set(known.keys())
|
||||
|
||||
for pid in known_set - current:
|
||||
last_rss = known.pop(pid)
|
||||
deaths.append((ts(), pid, last_rss))
|
||||
print(f"[{ts()}] *** WORKER DIED pid={pid} last_rss={last_rss:.1f}MB ***")
|
||||
|
||||
for pid in current - known_set:
|
||||
rss = get_rss_mb(pid) or 0
|
||||
known[pid] = rss
|
||||
spawns.append((ts(), pid))
|
||||
if time.time() - t0 > 3: # Skip initial spawn noise
|
||||
print(f"[{ts()}] +++ WORKER SPAWN pid={pid} rss={rss:.1f}MB +++")
|
||||
|
||||
parts = []
|
||||
for pid in sorted(current):
|
||||
rss = get_rss_mb(pid)
|
||||
if rss is not None:
|
||||
known[pid] = rss
|
||||
parts.append(f"pid={pid}:{rss:.0f}MB")
|
||||
rss_history.append((time.time() - t0, pid, rss))
|
||||
|
||||
if parts:
|
||||
print(f"[{ts()}] RSS: {' | '.join(parts)}")
|
||||
|
||||
if len(deaths) >= 2:
|
||||
print(f"\n[{ts()}] === CONFIRMED: {len(deaths)} worker deaths ===")
|
||||
break
|
||||
|
||||
if not os.path.exists(f"/proc/{master_pid}"):
|
||||
print(f"[{ts()}] Master gone")
|
||||
break
|
||||
|
||||
time.sleep(POLL_INTERVAL)
|
||||
|
||||
return deaths, spawns, rss_history
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Main
|
||||
# ---------------------------------------------------------------------------
|
||||
def main():
|
||||
print("=" * 70)
|
||||
print("LiteLLM Proxy Memory Leak Reproduction")
|
||||
print("=" * 70)
|
||||
print(f"Workers: {NUM_WORKERS}, Mem cap: {WORKER_MEM_CAP_MB}MB, "
|
||||
f"Concurrency: {BLAST_CONCURRENCY}, Requests: {BLAST_TOTAL}")
|
||||
print()
|
||||
|
||||
start_fake_openai()
|
||||
config_path = write_proxy_config()
|
||||
proxy_proc = start_proxy(config_path)
|
||||
|
||||
# Give proxy a moment to fully start
|
||||
time.sleep(3)
|
||||
|
||||
blast_proc = start_blast()
|
||||
stream_output(blast_proc, "blast")
|
||||
|
||||
# Give blast a moment to ramp up
|
||||
time.sleep(2)
|
||||
|
||||
try:
|
||||
deaths, spawns, rss_history = monitor_workers(proxy_proc.pid, MAX_DURATION)
|
||||
except KeyboardInterrupt:
|
||||
deaths, spawns, rss_history = [], [], []
|
||||
|
||||
# Summary
|
||||
print()
|
||||
print("=" * 70)
|
||||
print("REPRODUCTION SUMMARY")
|
||||
print("=" * 70)
|
||||
print(f"Worker deaths: {len(deaths)}")
|
||||
for t, pid, rss in deaths:
|
||||
print(f" [{t}] pid={pid} last_rss={rss:.1f}MB")
|
||||
print(f"Worker spawns: {len(spawns)}")
|
||||
|
||||
if rss_history:
|
||||
by_pid = {}
|
||||
for elapsed, pid, rss in rss_history:
|
||||
by_pid.setdefault(pid, []).append((elapsed, rss))
|
||||
print("\nRSS growth per worker:")
|
||||
for pid, points in sorted(by_pid.items()):
|
||||
if len(points) >= 2:
|
||||
t0_p, rss0 = points[0]
|
||||
t1_p, rss1 = points[-1]
|
||||
dt = t1_p - t0_p
|
||||
if dt > 0:
|
||||
rate = (rss1 - rss0) / dt
|
||||
print(f" pid={pid}: {rss0:.0f}MB → {rss1:.0f}MB over {dt:.0f}s ({rate:+.1f} MB/s)")
|
||||
|
||||
if len(deaths) >= 1:
|
||||
print("\n✓ REPRODUCTION SUCCESSFUL: Workers died from memory growth")
|
||||
print(" Root cause: unbounded spend_log_transactions list (simulated via middleware)")
|
||||
else:
|
||||
print("\n⚠ Workers did not die during test window")
|
||||
if rss_history:
|
||||
max_rss = max(r for _, _, r in rss_history)
|
||||
print(f" Max RSS observed: {max_rss:.0f}MB (cap={WORKER_MEM_CAP_MB}MB)")
|
||||
|
||||
# Cleanup
|
||||
print(f"\n[{ts()}] Cleaning up...")
|
||||
for p in [blast_proc, proxy_proc]:
|
||||
try:
|
||||
p.send_signal(signal.SIGTERM)
|
||||
p.wait(timeout=3)
|
||||
except Exception:
|
||||
try:
|
||||
p.kill()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
159
tests/memory_leak_repro/watch_workers.py
Normal file
159
tests/memory_leak_repro/watch_workers.py
Normal file
|
|
@ -0,0 +1,159 @@
|
|||
"""
|
||||
Worker RSS monitor — watches litellm/uvicorn worker processes for memory growth and deaths.
|
||||
|
||||
Scans /proc every 2 seconds, prints RSS for each worker, detects deaths and spawns.
|
||||
Run: python watch_workers.py [--parent-pid PID]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
def get_children_pids(parent_pid):
|
||||
"""Get all child PIDs of a given parent PID by scanning /proc."""
|
||||
children = []
|
||||
try:
|
||||
for entry in os.listdir("/proc"):
|
||||
if not entry.isdigit():
|
||||
continue
|
||||
pid = int(entry)
|
||||
try:
|
||||
with open(f"/proc/{pid}/stat", "r") as f:
|
||||
stat = f.read()
|
||||
# Field 4 (0-indexed 3) is the parent PID
|
||||
parts = stat.split(")") # Split after comm field which may contain spaces
|
||||
if len(parts) >= 2:
|
||||
fields = parts[-1].split()
|
||||
ppid = int(fields[1]) # Field index 1 after closing paren = ppid
|
||||
if ppid == parent_pid:
|
||||
children.append(pid)
|
||||
except (FileNotFoundError, PermissionError, ValueError, IndexError):
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
return children
|
||||
|
||||
|
||||
def get_rss_mb(pid):
|
||||
"""Get RSS in MB for a given PID from /proc/[pid]/status."""
|
||||
try:
|
||||
with open(f"/proc/{pid}/status", "r") as f:
|
||||
for line in f:
|
||||
if line.startswith("VmRSS:"):
|
||||
# VmRSS: 12345 kB
|
||||
parts = line.split()
|
||||
return int(parts[1]) / 1024.0 # kB → MB
|
||||
except (FileNotFoundError, PermissionError, ValueError):
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def get_cmdline(pid):
|
||||
"""Get command line for a PID."""
|
||||
try:
|
||||
with open(f"/proc/{pid}/cmdline", "r") as f:
|
||||
return f.read().replace("\0", " ").strip()[:80]
|
||||
except (FileNotFoundError, PermissionError):
|
||||
return ""
|
||||
|
||||
|
||||
def find_proxy_master_pid():
|
||||
"""Find the litellm/uvicorn master process PID."""
|
||||
for entry in os.listdir("/proc"):
|
||||
if not entry.isdigit():
|
||||
continue
|
||||
pid = int(entry)
|
||||
cmdline = get_cmdline(pid)
|
||||
if "litellm" in cmdline and ("--port" in cmdline or "--config" in cmdline):
|
||||
return pid
|
||||
if "uvicorn" in cmdline and "litellm" in cmdline:
|
||||
return pid
|
||||
return None
|
||||
|
||||
|
||||
def now_str():
|
||||
return datetime.now().strftime("%H:%M:%S")
|
||||
|
||||
|
||||
def main(parent_pid=None):
|
||||
if parent_pid is None:
|
||||
print("[watch] Searching for litellm/uvicorn master process...")
|
||||
for _ in range(60):
|
||||
parent_pid = find_proxy_master_pid()
|
||||
if parent_pid:
|
||||
break
|
||||
time.sleep(1)
|
||||
if parent_pid is None:
|
||||
print("[watch] ERROR: Could not find litellm/uvicorn master process")
|
||||
sys.exit(1)
|
||||
|
||||
print(f"[watch] Monitoring children of master PID {parent_pid} (RSS in MB)")
|
||||
print(f"[watch] Master cmdline: {get_cmdline(parent_pid)}")
|
||||
print()
|
||||
|
||||
known_pids = {} # pid -> last_rss
|
||||
deaths = []
|
||||
spawns = []
|
||||
|
||||
try:
|
||||
while True:
|
||||
children = get_children_pids(parent_pid)
|
||||
current_set = set(children)
|
||||
known_set = set(known_pids.keys())
|
||||
|
||||
# Detect deaths
|
||||
for pid in known_set - current_set:
|
||||
last_rss = known_pids.pop(pid)
|
||||
msg = f"[{now_str()}] *** WORKER DIED pid={pid} last_rss={last_rss:.1f}MB ***"
|
||||
print(msg)
|
||||
deaths.append((now_str(), pid, last_rss))
|
||||
|
||||
# Detect spawns
|
||||
for pid in current_set - known_set:
|
||||
rss = get_rss_mb(pid) or 0.0
|
||||
known_pids[pid] = rss
|
||||
msg = f"[{now_str()}] +++ WORKER SPAWNED pid={pid} rss={rss:.1f}MB +++"
|
||||
print(msg)
|
||||
spawns.append((now_str(), pid))
|
||||
|
||||
# Report RSS for all workers
|
||||
line_parts = []
|
||||
for pid in sorted(current_set):
|
||||
rss = get_rss_mb(pid)
|
||||
if rss is not None:
|
||||
known_pids[pid] = rss
|
||||
line_parts.append(f"pid={pid} rss={rss:.1f}MB")
|
||||
else:
|
||||
line_parts.append(f"pid={pid} rss=???")
|
||||
|
||||
if line_parts:
|
||||
print(f"[{now_str()}] {' | '.join(line_parts)}")
|
||||
|
||||
# Check if master is still alive
|
||||
if not os.path.exists(f"/proc/{parent_pid}"):
|
||||
print(f"[{now_str()}] Master PID {parent_pid} is gone. Exiting.")
|
||||
break
|
||||
|
||||
time.sleep(2)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
|
||||
print(f"\n=== SUMMARY ===")
|
||||
print(f"Worker deaths: {len(deaths)}")
|
||||
for ts, pid, rss in deaths:
|
||||
print(f" [{ts}] pid={pid} last_rss={rss:.1f}MB")
|
||||
print(f"Worker spawns: {len(spawns)}")
|
||||
for ts, pid in spawns:
|
||||
print(f" [{ts}] pid={pid}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--parent-pid", type=int, default=None)
|
||||
args = parser.parse_args()
|
||||
main(args.parent_pid)
|
||||
Loading…
Add table
Reference in a new issue