diff --git a/tests/memory_leak_repro/blast.py b/tests/memory_leak_repro/blast.py new file mode 100644 index 00000000000..aaa21e09960 --- /dev/null +++ b/tests/memory_leak_repro/blast.py @@ -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)) diff --git a/tests/memory_leak_repro/fake_openai_server.py b/tests/memory_leak_repro/fake_openai_server.py new file mode 100644 index 00000000000..689e1c9edb1 --- /dev/null +++ b/tests/memory_leak_repro/fake_openai_server.py @@ -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 + ) diff --git a/tests/memory_leak_repro/leak_inducer.py b/tests/memory_leak_repro/leak_inducer.py new file mode 100644 index 00000000000..785ea411924 --- /dev/null +++ b/tests/memory_leak_repro/leak_inducer.py @@ -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() diff --git a/tests/memory_leak_repro/proxy_config.yaml b/tests/memory_leak_repro/proxy_config.yaml new file mode 100644 index 00000000000..80643c4cbf8 --- /dev/null +++ b/tests/memory_leak_repro/proxy_config.yaml @@ -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 diff --git a/tests/memory_leak_repro/proxy_wrapper.py b/tests/memory_leak_repro/proxy_wrapper.py new file mode 100644 index 00000000000..ef817010704 --- /dev/null +++ b/tests/memory_leak_repro/proxy_wrapper.py @@ -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()})") diff --git a/tests/memory_leak_repro/repro_worker_death.py b/tests/memory_leak_repro/repro_worker_death.py new file mode 100644 index 00000000000..4cabbe673b0 --- /dev/null +++ b/tests/memory_leak_repro/repro_worker_death.py @@ -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() diff --git a/tests/memory_leak_repro/watch_workers.py b/tests/memory_leak_repro/watch_workers.py new file mode 100644 index 00000000000..d5fa3ea1154 --- /dev/null +++ b/tests/memory_leak_repro/watch_workers.py @@ -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)