Skip to content

Commit c64bd8e

Browse files
fluffy314cursoragent
authored andcommitted
fix(runtime): preserve completed sessions without forking
Guard terminal gRPC streams from late cancellation cleanup and replace ps-based footprint sampling so long-running decode workers cannot deadlock gRPC during memory checks. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent 53dbfac commit c64bd8e

5 files changed

Lines changed: 120 additions & 15 deletions

File tree

inference_engine/server/grpc_app.py

Lines changed: 26 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -240,11 +240,15 @@ def _watch_context(
240240
context,
241241
cancel_event: threading.Event,
242242
session_id: str,
243+
completed_event: threading.Event | None = None,
243244
) -> None:
244245
callback = getattr(context, "add_done_callback", None)
245246
if callback is not None:
246247
def done(done_context) -> None:
247-
if done_context.cancelled():
248+
logically_completed = (
249+
completed_event is not None and completed_event.is_set()
250+
)
251+
if done_context.cancelled() and not logically_completed:
248252
cancel_event.set()
249253
self._store.remove_session_if_present(
250254
session_id, reason="client_cancelled",
@@ -305,7 +309,13 @@ async def AppendTokens( # noqa: N802 — gRPC-generated method casing
305309
f"session {request.session_id!r} already has an active operation",
306310
)
307311
cancel_event = threading.Event()
308-
self._watch_context(context, cancel_event, request.session_id)
312+
completed_event = threading.Event()
313+
self._watch_context(
314+
context,
315+
cancel_event,
316+
request.session_id,
317+
completed_event,
318+
)
309319
try:
310320
kwargs = {
311321
"session_id": request.session_id,
@@ -316,6 +326,7 @@ async def AppendTokens( # noqa: N802 — gRPC-generated method casing
316326
new_history_length = await asyncio.to_thread(
317327
self._append.append_tokens, **kwargs,
318328
)
329+
completed_event.set()
319330
except asyncio.CancelledError:
320331
cancel_event.set()
321332
self._store.remove_session_if_present(
@@ -390,7 +401,13 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
390401
f"session {request.session_id!r} already has an active operation",
391402
)
392403
cancel_event = threading.Event()
393-
self._watch_context(context, cancel_event, request.session_id)
404+
completed_event = threading.Event()
405+
self._watch_context(
406+
context,
407+
cancel_event,
408+
request.session_id,
409+
completed_event,
410+
)
394411
kwargs = {
395412
"session_id": request.session_id,
396413
"max_tokens": request.max_tokens,
@@ -413,6 +430,7 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
413430
if error is not None:
414431
raise error
415432
if event is _STREAM_END:
433+
completed_event.set()
416434
break
417435
if context.cancelled():
418436
threaded_stream.cancel()
@@ -436,6 +454,7 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
436454
# DoneEvent — the only remaining event type per
437455
# the GenerateEvent union.
438456
assert isinstance(event, DoneEvent)
457+
completed_event.set()
439458
yield runtime_pb2.GenerateResponse(
440459
done=runtime_pb2.GenerateDone(
441460
stop_reason=_STOP_REASON_TO_PROTO[
@@ -456,9 +475,10 @@ async def Generate( # noqa: N802 — gRPC-generated method casing
456475
await context.abort(grpc.StatusCode.ABORTED, str(exc))
457476
except asyncio.CancelledError:
458477
threaded_stream.cancel()
459-
self._store.remove_session_if_present(
460-
request.session_id, reason="client_cancelled",
461-
)
478+
if not completed_event.is_set():
479+
self._store.remove_session_if_present(
480+
request.session_id, reason="client_cancelled",
481+
)
462482
raise
463483
finally:
464484
threaded_stream.cancel()

inference_engine/server/runtime_health.py

Lines changed: 4 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -5,13 +5,13 @@
55
import gc
66
import json
77
import os
8-
import subprocess
98
import threading
109
import time
1110
from dataclasses import asdict, dataclass
1211
from pathlib import Path
1312
from typing import Callable, Optional
1413

14+
import psutil
1515
from fastapi import FastAPI
1616
from fastapi.responses import JSONResponse
1717

@@ -95,15 +95,10 @@ class MemorySnapshot:
9595

9696

9797
def process_footprint_bytes(pid: Optional[int] = None) -> int:
98-
"""Return resident process bytes using the macOS/Linux ``ps`` contract."""
98+
"""Return resident process bytes without forking the threaded runtime."""
9999
try:
100-
output = subprocess.check_output(
101-
["ps", "-o", "rss=", "-p", str(pid or os.getpid())],
102-
text=True,
103-
timeout=1,
104-
)
105-
return int(output.strip()) * 1024
106-
except (OSError, ValueError, subprocess.SubprocessError):
100+
return int(psutil.Process(pid or os.getpid()).memory_info().rss)
101+
except (OSError, ValueError, psutil.Error):
107102
return 0
108103

109104

requirements.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ uvicorn>=0.32,<1.0
4444
pydantic>=2.7,<3.0
4545
httpx>=0.27,<1.0
4646
prometheus-client>=0.20,<1.0
47+
psutil>=5.9,<8.0
4748

4849
# gRPC runtime (PR-B1 of ADR 0008 Phase B; the runtime/SDK protocol).
4950
# grpcio-tools is needed only at build time (regenerating stubs from

tests/inference_engine/server/test_grpc_app.py

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -664,6 +664,50 @@ async def abort(self, code, details): # pragma: no cover
664664
store.get_session(sess.session_id)
665665

666666

667+
async def test_generate_cancellation_after_done_preserves_session():
668+
from inference_engine.session import (
669+
DoneEvent,
670+
GenerationCoordinator,
671+
STOP_REASON_MAX_TOKENS,
672+
)
673+
674+
class CompletedGeneration(GenerationCoordinator):
675+
def generate(self, session_id, *, max_tokens, **kwargs):
676+
del session_id, max_tokens, kwargs
677+
yield DoneEvent(
678+
stop_reason=STOP_REASON_MAX_TOKENS,
679+
generated_token_count=0,
680+
prefill_seconds=0.0,
681+
total_seconds=0.0,
682+
)
683+
684+
class Context:
685+
def cancelled(self):
686+
return False
687+
688+
async def abort(self, code, details):
689+
raise AssertionError((code, details))
690+
691+
store = SessionStore(capacity=1)
692+
session = store.create_session()
693+
servicer = RuntimeServiceServicer(
694+
store,
695+
generation_coordinator=CompletedGeneration(store, verifier=None),
696+
)
697+
stream = servicer.Generate(
698+
runtime_pb2.GenerateRequest(
699+
session_id=session.session_id,
700+
max_tokens=1,
701+
),
702+
Context(),
703+
)
704+
response = await anext(stream)
705+
assert response.WhichOneof("payload") == "done"
706+
with pytest.raises(asyncio.CancelledError):
707+
await stream.athrow(asyncio.CancelledError())
708+
assert store.get_session(session.session_id) is session
709+
710+
667711
async def test_wire_cancellation_releases_blocked_generate_session():
668712
"""A disconnect removes session state while model thread is still blocked."""
669713
from inference_engine.session import GenerationCoordinator
@@ -945,6 +989,26 @@ async def test_watch_context_removes_cancelled_session():
945989
assert store.active_count == 0
946990

947991

992+
async def test_watch_context_preserves_logically_completed_session():
993+
store = SessionStore(capacity=1)
994+
session = store.create_session()
995+
servicer = RuntimeServiceServicer(store)
996+
context = _DirectContext()
997+
cancel_event = threading.Event()
998+
completed_event = threading.Event()
999+
servicer._watch_context(
1000+
context,
1001+
cancel_event,
1002+
session.session_id,
1003+
completed_event,
1004+
)
1005+
completed_event.set()
1006+
assert context.callback is not None
1007+
context.callback(type("Done", (), {"cancelled": lambda self: True})())
1008+
assert not cancel_event.is_set()
1009+
assert store.active_count == 1
1010+
1011+
9481012
async def test_create_session_rejects_memory_drain():
9491013
store = SessionStore(capacity=1)
9501014
governor = type("Governor", (), {"draining": True})()

tests/inference_engine/server/test_runtime_health.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
import json
22

3+
import psutil
34
from fastapi.testclient import TestClient
45

6+
from inference_engine.server import runtime_health
57
from inference_engine.server.runtime_health import (
68
DecodeLiveness,
79
PrimaryMemoryGovernor,
@@ -18,6 +20,29 @@ def reset(self):
1820
self.resets += 1
1921

2022

23+
def test_process_footprint_uses_fork_free_psutil(monkeypatch):
24+
seen = []
25+
26+
class Process:
27+
def __init__(self, pid):
28+
seen.append(pid)
29+
30+
def memory_info(self):
31+
return type("MemoryInfo", (), {"rss": 12345})()
32+
33+
monkeypatch.setattr(runtime_health.psutil, "Process", Process)
34+
assert runtime_health.process_footprint_bytes(42) == 12345
35+
assert seen == [42]
36+
37+
38+
def test_process_footprint_returns_zero_when_process_disappears(monkeypatch):
39+
def missing(pid):
40+
raise psutil.NoSuchProcess(pid)
41+
42+
monkeypatch.setattr(runtime_health.psutil, "Process", missing)
43+
assert runtime_health.process_footprint_bytes(42) == 0
44+
45+
2146
def test_liveness_is_atomically_published(tmp_path):
2247
path = tmp_path / "live.json"
2348
live = DecodeLiveness(path, clock=lambda: 123.0, pid=42)

0 commit comments

Comments
 (0)