[mem_cache] Move req_pool_idx into ReqKvInfo (#37094)

This commit is contained in:
Liangsheng Yin
2026-08-30 14:46:21 -07:00
committed by GitHub
parent 8a87079dbb
commit 007ef5e23a
68 changed files with 372 additions and 357 deletions
@@ -23,7 +23,7 @@ def _drain_until_released(t: ScriptedContext, *handles: ScriptedReqHandle):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.req_pool_idx is None)
and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return
@@ -57,7 +57,7 @@ class TestAbortBasic(ScriptedTestCase):
r.kv_pages == 0
), f"abort must release KV; r.kv_pages={r.kv_pages} after abort"
assert (
r.req is None or r.req.req_pool_idx is None
r.req is None or r.req.kv.req_pool_idx is None
), f"abort must release row; r.req={r.req} after abort"
assert (
r.lock_refs == 0
@@ -78,7 +78,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
def test_abort_at_chunk_mid(self):
self.server.execute_script(self._script_abort_at_chunk_mid)
@@ -124,7 +124,7 @@ class TestAbortBasic(ScriptedTestCase):
t.abort(r)
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
def test_abort_at_admission_step(self):
@@ -137,7 +137,7 @@ class TestAbortBasic(ScriptedTestCase):
t.abort(r)
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
def test_abort_then_start_same_step_new_rid(self):
self.server.execute_script(self._script_abort_then_start_same_step_new_rid)
@@ -191,7 +191,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from _drain_until_released(t, *reqs)
for r in reqs:
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
def test_abort_unknown_rid_noop(self):
self.server.execute_script(self._script_abort_unknown_rid_noop)
@@ -264,7 +264,7 @@ class TestAbortBasic(ScriptedTestCase):
t.abort(r)
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
def test_double_abort_idempotent(self):
@@ -351,7 +351,7 @@ class TestAbortBasic(ScriptedTestCase):
f"aborted req revived and ran another chunk; "
f"chunks_done went {chunks_after_abort} -> {r.chunks_done}"
)
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
def test_abort_mid_chunk_no_extra_radix_node(self):
self.server.execute_script(self._script_abort_mid_chunk_no_extra_radix_node)
@@ -369,7 +369,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
chunks_after_release = r.chunks_done
for _ in range(4):
yield
@@ -406,7 +406,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from run_until_finished(r2)
assert r2.finished, "resubmit under same rid must complete independently"
assert r1.kv_pages == 0, "aborted r1 must release KV before resubmit"
assert r1.req is None or r1.req.req_pool_idx is None
assert r1.req is None or r1.req.kv.req_pool_idx is None
assert r1.lock_refs == 0
def test_abort_during_gap_inflight_middle_chunks_positive(self):
@@ -429,7 +429,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert not r.is_chunking, "aborted gap req must not re-enter chunking"
yield
@@ -511,7 +511,7 @@ class TestAbortBasic(ScriptedTestCase):
assert r1.kv_pages == 0, (
f"force_retract + abort same yield must release KV; got " f"{r1.kv_pages}"
)
assert r1.req is None or r1.req.req_pool_idx is None, (
assert r1.req is None or r1.req.kv.req_pool_idx is None, (
f"force_retract + abort same yield must release row; got " f"{r1.req}"
)
assert r1.lock_refs == 0, (
@@ -540,7 +540,7 @@ class TestAbortBasic(ScriptedTestCase):
yield from run_until(r2, lambda h: h.is_chunking)
assert r1.kv_pages == 0
assert r1.req is None or r1.req.req_pool_idx is None
assert r1.req is None or r1.req.kv.req_pool_idx is None
assert r1.lock_refs == 0
yield from run_until_finished(r2)
assert r2.finished, "baton handoff must let r2 complete"
@@ -570,7 +570,7 @@ class TestAbortPP(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
assert r.finished
@@ -409,13 +409,13 @@ class TestKVPressureSmallPool(ScriptedTestCase):
if (
r_chunk.kv_pages == 0
and r_chunk.lock_refs == 0
and (r_chunk.req is None or r_chunk.req.req_pool_idx is None)
and (r_chunk.req is None or r_chunk.req.kv.req_pool_idx is None)
):
break
yield
assert r_chunk.kv_pages == 0, f"kv_pages={r_chunk.kv_pages}"
assert r_chunk.lock_refs == 0, f"lock_refs={r_chunk.lock_refs}"
assert r_chunk.req is None or r_chunk.req.req_pool_idx is None
assert r_chunk.req is None or r_chunk.req.kv.req_pool_idx is None
t.abort(ballast)
for _ in range(200):
@@ -18,7 +18,7 @@ def _drain_until_released(t, *handles):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.req_pool_idx is None)
and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return
@@ -223,7 +223,7 @@ class TestLifecycleBasic(ScriptedTestCase):
r = t.start_req(prompt_len=16, max_new_tokens=2, ignore_eos=True)
yield from run_until_finished(r)
assert r.finished
assert r.req.req_pool_idx is None
assert r.req.kv.req_pool_idx is None
assert r.kv_pages == 0
assert r.lock_refs == 0
@@ -235,7 +235,7 @@ class TestLifecycleBasic(ScriptedTestCase):
r1 = t.start_req(prompt_len=16, max_new_tokens=2, ignore_eos=True)
yield from run_until_finished(r1)
yield from _drain_until_released(t, r1)
assert r1.req.req_pool_idx is None and r1.kv_pages == 0 and r1.lock_refs == 0
assert r1.req.kv.req_pool_idx is None and r1.kv_pages == 0 and r1.lock_refs == 0
r1_output_len = len(r1.req.output_ids)
r2 = t.start_req(prompt_len=16, max_new_tokens=2, ignore_eos=True)
@@ -243,7 +243,7 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from _drain_until_released(t, r2)
assert r1.finished and r2.finished
assert r1_output_len == 2 and len(r2.req.output_ids) == 2
assert r2.req.req_pool_idx is None and r2.kv_pages == 0 and r2.lock_refs == 0
assert r2.req.kv.req_pool_idx is None and r2.kv_pages == 0 and r2.lock_refs == 0
def test_five_seq_clean(self):
self.server.execute_script(self._script_five_seq_clean)
@@ -256,7 +256,7 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from run_until_finished(r)
assert r.finished
assert len(r.req.output_ids) == 2
assert r.req.req_pool_idx is None
assert r.req.kv.req_pool_idx is None
assert r.kv_pages == 0
assert r.lock_refs == 0
reqs.append(r)
@@ -300,7 +300,9 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from run_until_finished(r)
assert r.finished
assert len(r.req.output_ids) == 2
assert r.req.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
assert (
r.req.kv.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
)
if prompt == VERY_LONG_PROMPT_LEN:
assert r.chunks_done == 8
else:
@@ -319,7 +321,7 @@ class TestLifecycleBasic(ScriptedTestCase):
assert r.finished
assert len(r.req.output_ids) == 1
yield from _drain_until_released(t, r)
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.kv_pages == 0 and r.lock_refs == 0
if L > DEFAULT_CHUNK_SIZE:
assert (
@@ -341,7 +343,7 @@ class TestLifecycleBasic(ScriptedTestCase):
assert r.finished
assert len(r.req.output_ids) == 1
yield from _drain_until_released(t, r)
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.kv_pages == 0 and r.lock_refs == 0
if L > DEFAULT_CHUNK_SIZE:
assert (
@@ -360,7 +362,9 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from run_until_finished(r)
assert r.finished
assert len(r.req.output_ids) == 2
assert r.req.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
assert (
r.req.kv.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
)
for _ in range(20):
yield
@@ -377,7 +381,9 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from run_until_finished(r)
assert r.finished
assert len(r.req.output_ids) == 2
assert r.req.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
assert (
r.req.kv.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
)
if L == VERY_LONG_PROMPT_LEN:
assert r.chunks_done == 8
else:
@@ -394,7 +400,9 @@ class TestLifecycleBasic(ScriptedTestCase):
yield from run_until_finished(r)
assert r.finished
assert len(r.req.output_ids) == 2
assert r.req.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
assert (
r.req.kv.req_pool_idx is None and r.kv_pages == 0 and r.lock_refs == 0
)
for _ in range(5):
yield
t.flush_cache()
@@ -432,7 +440,7 @@ class TestLifecycleBasic(ScriptedTestCase):
)
assert r.finished or _error_message(r) is not None
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
@@ -222,7 +222,9 @@ class TestLoRAAdapterEviction(ScriptedTestCase):
t.abort(r_a)
for _ in range(12):
if r_a.kv_pages == 0 and (r_a.req is None or r_a.req.req_pool_idx is None):
if r_a.kv_pages == 0 and (
r_a.req is None or r_a.req.kv.req_pool_idx is None
):
break
yield
@@ -173,7 +173,7 @@ class TestMultiReqBasic(ScriptedTestCase):
for _ in range(5):
yield
assert r.kv_pages == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
def test_rid_reuse_after_finish(self):
self.server.execute_script(self._script_rid_reuse_after_finish)
@@ -31,7 +31,7 @@ def _expected_chunks(prompt_len: int, chunk_size: int) -> int:
def _drain_until_released(t, *handles):
for _ in range(16):
if all(
h.kv_pages == 0 and (h.req is None or h.req.req_pool_idx is None)
h.kv_pages == 0 and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return
@@ -168,13 +168,13 @@ class TestPriorityBasic(ScriptedTestCase):
if (
r.kv_pages == 0
and r.lock_refs == 0
and (r.req is None or r.req.req_pool_idx is None)
and (r.req is None or r.req.kv.req_pool_idx is None)
):
break
yield
assert r.kv_pages == 0
assert r.lock_refs == 0
assert r.req is None or r.req.req_pool_idx is None
assert r.req is None or r.req.kv.req_pool_idx is None
t.continue_generation()
yield
assert r.kv_pages == 0 and r.lock_refs == 0
@@ -19,7 +19,7 @@ def _drain_until_released(t, *handles):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.req_pool_idx is None)
and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return
@@ -41,7 +41,7 @@ class TestRegressionBasic(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert r.kv_pages == 0
assert r.req.req_pool_idx is None
assert r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
assert not r.is_chunking
assert r.req.inflight_middle_chunks == 0
@@ -58,7 +58,7 @@ class TestRegressionBasic(ScriptedTestCase):
yield
assert r.kv_pages == 0
assert r.req.req_pool_idx is None
assert r.req.kv.req_pool_idx is None
assert r.lock_refs == 0
assert not r.is_chunking
@@ -240,7 +240,7 @@ class TestRegressionBasic(ScriptedTestCase):
r = t.start_req(prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2)
yield from run_until(r, lambda h: h.is_chunking and h.chunks_done >= 1)
assert r.req.req_pool_idx is not None, "row must be held mid-chunk"
assert r.req.kv.req_pool_idx is not None, "row must be held mid-chunk"
assert r.kv_pages > 0, "committed KV must be held mid-chunk"
assert r.lock_refs >= 1, "radix lock_ref must be held mid-chunk"
@@ -248,8 +248,8 @@ class TestRegressionBasic(ScriptedTestCase):
yield from _drain_until_released(t, r)
assert (
r.req.req_pool_idx is None
), f"96d4749094: abort must release row; got row_idx={r.req.req_pool_idx!r}"
r.req.kv.req_pool_idx is None
), f"96d4749094: abort must release row; got row_idx={r.req.kv.req_pool_idx!r}"
assert (
r.kv_pages == 0
), f"96d4749094: abort must release KV; got kv_pages={r.kv_pages}"
@@ -269,14 +269,14 @@ class TestRegressionBasic(ScriptedTestCase):
def _script_pause_retract_releases_waiting_chunked_resume(t: ScriptedContext):
r = t.start_req(prompt_len=VERY_LONG_PROMPT_LEN, max_new_tokens=2)
yield from run_until(r, lambda h: h.is_chunking and h.chunks_done >= 1)
assert r.req.req_pool_idx is not None and r.kv_pages > 0 and r.lock_refs >= 1
assert r.req.kv.req_pool_idx is not None and r.kv_pages > 0 and r.lock_refs >= 1
t.pause_generation(mode="retract")
yield
assert r.req.req_pool_idx is None, (
assert r.req.kv.req_pool_idx is None, (
f"f38e69f87d: pause(retract) must release waiting "
f"chunked-resume row; got row_idx={r.req.req_pool_idx!r}"
f"chunked-resume row; got row_idx={r.req.kv.req_pool_idx!r}"
)
assert r.kv_pages == 0
assert r.lock_refs == 0
@@ -14,7 +14,7 @@ def _drain_until_released(t: ScriptedContext, *handles: ScriptedReqHandle):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.req_pool_idx is None)
and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return