# SPDX-License-Identifier: Apache-2.0 """ Tests for interleaved chunked prefill + decode (SchedulerConfig.chunked_prefill). Strategy: keep tests fast by mocking MLX model calls and cache operations. _begin_prefill() and _step_prefill_chunk() are tested by patching make_prompt_cache and mx.eval; the scheduler-level flow is tested by patching _step_prefill_chunk directly. """ from collections import deque from types import SimpleNamespace from unittest.mock import MagicMock, patch from omlx.exceptions import PrefillMemoryExceededError from omlx.request import Request, RequestStatus, SamplingParams from omlx.scheduler import ( PrefillEvictionRequest, Scheduler, SchedulerConfig, _PrefillAbortedError, _PrefillEvictionNeeded, _PrefillState, ) # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_scheduler(chunked_prefill: bool = True, step_size: int = 4) -> Scheduler: """Return a Scheduler with a mock model/tokenizer and chunked_prefill config.""" model = MagicMock() model.layers = [] # No attention layers — keeps _build_state_machine simple tokenizer = MagicMock() tokenizer.eos_token_id = 2 config = SchedulerConfig( max_num_seqs=8, prefill_step_size=step_size, chunked_prefill=chunked_prefill, paged_cache_block_size=0, # Disable boundary snapshots ) scheduler = Scheduler(model=model, tokenizer=tokenizer, config=config) # Replace the real batch_generator factory so insert() returns a uid. mock_bg = MagicMock() mock_bg.insert.return_value = [42] mock_bg.next_generated.return_value = iter([]) scheduler.batch_generator = mock_bg scheduler._current_sampler_params = () return scheduler def _make_request(request_id: str = "req-1", n_tokens: int = 10) -> Request: """Return a pre-tokenized request with *n_tokens* prompt tokens.""" req = Request( request_id=request_id, prompt=list(range(n_tokens)), sampling_params=SamplingParams(max_tokens=32), ) req.prompt_token_ids = list(range(n_tokens)) req.num_prompt_tokens = n_tokens req.remaining_tokens = list(range(n_tokens)) return req def _make_prefill_state( scheduler: Scheduler, request: Request, n_remaining: int = 20 ) -> _PrefillState: """Build a minimal _PrefillState for direct testing.""" import mlx.core as mx tokens_remaining = mx.zeros((1, n_remaining), dtype=mx.int32) state = _PrefillState( request=request, cache=[], tokens_remaining=tokens_remaining, last_token=[99], tokens_processed=0, base_size=0, emitted_boundaries={}, boundary_enabled=False, block_size=0, total_length=n_remaining + 1, sampler=MagicMock(), sm=MagicMock(), per_row_lps=[], ) return state class _RecordingModel: def __init__(self, model_type: str): self.model_type = model_type self.layers = [] self.chunk_lengths: list[int] = [] def __call__(self, tokens, cache=None): self.chunk_lengths.append(int(tokens.shape[1])) def _make_recording_scheduler( model_type: str, *, uses_minimax_m3_positions: bool = False, nested_vlm_model_type: str | None = None, model_name: str = "", ) -> tuple[Scheduler, _RecordingModel]: model = _RecordingModel(model_type) if uses_minimax_m3_positions: model._uses_minimax_m3_positions = True if nested_vlm_model_type is not None: model._vlm_model = SimpleNamespace( config=SimpleNamespace(model_type=nested_vlm_model_type) ) tokenizer = MagicMock() tokenizer.eos_token_id = 2 scheduler = Scheduler( model=model, tokenizer=tokenizer, config=SchedulerConfig( prefill_step_size=2048, chunked_prefill=True, paged_cache_block_size=0, model_name=model_name, ), ) return scheduler, model # --------------------------------------------------------------------------- # SchedulerConfig # --------------------------------------------------------------------------- class TestSchedulerConfigChunkedPrefill: def test_default_is_false(self): config = SchedulerConfig() assert config.chunked_prefill is False def test_can_be_enabled(self): config = SchedulerConfig(chunked_prefill=True) assert config.chunked_prefill is True # --------------------------------------------------------------------------- # _PrefillState # --------------------------------------------------------------------------- class TestPrefillState: def test_fields_accessible(self): import mlx.core as mx state = _PrefillState( request=MagicMock(), cache=[], tokens_remaining=mx.zeros((1, 5), dtype=mx.int32), last_token=[7], tokens_processed=0, base_size=0, emitted_boundaries={}, boundary_enabled=False, block_size=256, total_length=6, ) assert state.tokens_processed == 0 assert state.sampler is None assert state.per_row_lps is None def test_insert_params_settable(self): import mlx.core as mx state = _PrefillState( request=MagicMock(), cache=[], tokens_remaining=mx.zeros((1, 3), dtype=mx.int32), last_token=[1], tokens_processed=0, base_size=0, emitted_boundaries={}, boundary_enabled=False, block_size=256, total_length=4, ) state.sampler = "s" state.sm = "sm" state.per_row_lps = [] assert state.sampler == "s" # --------------------------------------------------------------------------- # Scheduler queues initialised # --------------------------------------------------------------------------- class TestSchedulerQueues: def test_prefilling_queue_exists(self): sched = _make_scheduler() assert hasattr(sched, "prefilling") assert isinstance(sched.prefilling, deque) assert len(sched.prefilling) == 0 def test_prefill_states_dict_exists(self): sched = _make_scheduler() assert hasattr(sched, "_prefill_states") assert isinstance(sched._prefill_states, dict) # --------------------------------------------------------------------------- # has_requests includes prefilling # --------------------------------------------------------------------------- class TestHasRequests: def test_false_when_all_empty(self): sched = _make_scheduler() assert not sched.has_requests() def test_true_when_prefilling(self): sched = _make_scheduler() req = _make_request() sched.prefilling.append(req) assert sched.has_requests() def test_still_true_with_waiting_only(self): sched = _make_scheduler() req = _make_request() sched.waiting.append(req) assert sched.has_requests() # --------------------------------------------------------------------------- # get_stats includes num_prefilling # --------------------------------------------------------------------------- class TestGetStats: def test_num_prefilling_in_stats(self): sched = _make_scheduler() stats = sched.get_stats() assert "num_prefilling" in stats assert stats["num_prefilling"] == 0 def test_num_prefilling_counts_correctly(self): sched = _make_scheduler() sched.prefilling.append(_make_request("r1")) sched.prefilling.append(_make_request("r2")) assert sched.get_stats()["num_prefilling"] == 2 # --------------------------------------------------------------------------- # GLM adaptive chunked prefill # --------------------------------------------------------------------------- class TestGLMAdaptiveChunkedPrefill: def test_glm_uses_adaptive_prefill_chunk_size(self, monkeypatch): monkeypatch.delenv("MLX_LM_GLM_DSA_ADAPTIVE_PREFILL_STEP", raising=False) monkeypatch.delenv("MLX_LM_GLM_DSA_ADAPTIVE_PREFILL_STEP_SIZE", raising=False) monkeypatch.delenv("MLX_LM_GLM_DSA_ADAPTIVE_PREFILL_AFTER", raising=False) monkeypatch.delenv( "MLX_LM_GLM_DSA_ADAPTIVE_PREFILL_MIN_REMAINING", raising=False ) sched, model = _make_recording_scheduler("glm_moe_dsa") req = _make_request("glm", n_tokens=8194) state = _make_prefill_state(sched, req, n_remaining=8193) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [8192] assert state.tokens_processed == 8192 def test_non_glm_keeps_configured_prefill_chunk_size(self, monkeypatch): monkeypatch.delenv("MLX_LM_GLM_DSA_ADAPTIVE_PREFILL_STEP", raising=False) sched, model = _make_recording_scheduler("deepseek_v32") req = _make_request("deepseek", n_tokens=8193) state = _make_prefill_state(sched, req, n_remaining=8192) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [2048] assert state.tokens_processed == 2048 # --------------------------------------------------------------------------- # MiniMax M3 adaptive chunked prefill # --------------------------------------------------------------------------- class TestMiniMaxM3AdaptiveChunkedPrefill: def test_minimax_m3_uses_4096_for_long_prefill(self, monkeypatch): monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", raising=False) monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP_SIZE", raising=False) monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_AFTER", raising=False) monkeypatch.delenv( "MLX_MINIMAX_M3_ADAPTIVE_PREFILL_MIN_REMAINING", raising=False ) sched, model = _make_recording_scheduler("minimax_m3") req = _make_request("minimax", n_tokens=4098) state = _make_prefill_state(sched, req, n_remaining=4097) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [4096] assert state.tokens_processed == 4096 def test_minimax_m3_keeps_2048_for_short_prefill(self, monkeypatch): monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", raising=False) sched, model = _make_recording_scheduler("minimax_m3_vl") req = _make_request("minimax-short", n_tokens=4096) state = _make_prefill_state(sched, req, n_remaining=4095) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [2048] assert state.tokens_processed == 2048 def test_minimax_m3_env_can_disable_adaptive_prefill(self, monkeypatch): monkeypatch.setenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", "0") sched, model = _make_recording_scheduler("minimax_m3") req = _make_request("minimax-disabled", n_tokens=4098) state = _make_prefill_state(sched, req, n_remaining=4097) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [2048] assert state.tokens_processed == 2048 def test_minimax_m3_vlm_adapter_flag_enables_adaptive_prefill(self, monkeypatch): monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", raising=False) sched, model = _make_recording_scheduler( "vlm", uses_minimax_m3_positions=True, ) req = _make_request("minimax-adapter", n_tokens=4098) state = _make_prefill_state(sched, req, n_remaining=4097) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [4096] assert state.tokens_processed == 4096 def test_minimax_m3_nested_vlm_model_enables_adaptive_prefill(self, monkeypatch): monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", raising=False) sched, model = _make_recording_scheduler( "vlm", nested_vlm_model_type="minimax_m3_vl", ) req = _make_request("minimax-nested-vlm", n_tokens=4098) state = _make_prefill_state(sched, req, n_remaining=4097) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [4096] assert state.tokens_processed == 4096 def test_minimax_m3_model_path_enables_adaptive_prefill( self, tmp_path, monkeypatch ): monkeypatch.delenv("MLX_MINIMAX_M3_ADAPTIVE_PREFILL_STEP", raising=False) (tmp_path / "config.json").write_text( '{"model_type": "minimax_m3_vl"}', encoding="utf-8", ) sched, model = _make_recording_scheduler( "vlm", model_name=str(tmp_path), ) req = _make_request("minimax-model-path", n_tokens=4098) state = _make_prefill_state(sched, req, n_remaining=4097) with patch("omlx.scheduler._sync_and_clear_cache"): done = sched._step_prefill_chunk(state) assert not done assert model.chunk_lengths == [4096] assert state.tokens_processed == 4096 # --------------------------------------------------------------------------- # reset() clears prefilling # --------------------------------------------------------------------------- class TestReset: def test_reset_clears_prefilling(self): sched = _make_scheduler() req = _make_request() sched.prefilling.append(req) sched._prefill_states[req.request_id] = MagicMock() sched.requests[req.request_id] = req sched.reset() assert len(sched.prefilling) == 0 assert len(sched._prefill_states) == 0 # --------------------------------------------------------------------------- # fail_all_requests() includes prefilling # --------------------------------------------------------------------------- class TestFailAllRequests: def test_fail_all_includes_prefilling(self): sched = _make_scheduler() req = _make_request("pf-req") sched.prefilling.append(req) sched._prefill_states[req.request_id] = MagicMock() sched.requests[req.request_id] = req failed = sched.fail_all_requests() assert "pf-req" in failed assert len(sched.prefilling) == 0 assert len(sched._prefill_states) == 0 # --------------------------------------------------------------------------- # _do_abort_request() cleans up prefilling # --------------------------------------------------------------------------- class TestAbortPrefilling: def test_abort_removes_from_prefilling(self): sched = _make_scheduler() req = _make_request("abort-me") req.status = RequestStatus.WAITING sched.prefilling.append(req) sched._prefill_states[req.request_id] = MagicMock() sched.requests[req.request_id] = req sched._do_abort_request(req.request_id) assert req.request_id not in sched._prefill_states assert all(r.request_id != req.request_id for r in sched.prefilling) # --------------------------------------------------------------------------- # _advance_chunked_prefills(): core logic # --------------------------------------------------------------------------- class TestAdvanceChunkedPrefills: def test_no_op_when_queue_empty(self): sched = _make_scheduler() scheduled = [] rejected = [] # Should not raise sched._advance_chunked_prefills(scheduled, rejected) assert scheduled == [] assert rejected == [] def test_advances_chunk_when_not_done(self): """Requests that still have tokens stay in prefilling queue.""" sched = _make_scheduler() req = _make_request("r1") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req, n_remaining=20) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object( sched, "_step_prefill_chunk", return_value=False ) as mock_step: scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) mock_step.assert_called_once_with(state) # Not done → stays in prefilling, not moved to running assert req in sched.prefilling assert scheduled == [] assert rejected == [] assert req.request_id not in sched.running def test_inserts_when_done(self): """Completed prefill is inserted into BatchGenerator and moved to running.""" sched = _make_scheduler() req = _make_request("r1") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req, n_remaining=1) state.sampler = MagicMock() state.sm = MagicMock() state.per_row_lps = [] sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object(sched, "_step_prefill_chunk", return_value=True): with patch.object(sched, "_emit_final_boundary_if_needed"): scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) # Moved to running, removed from prefilling assert req not in sched.prefilling assert req.request_id not in sched._prefill_states assert req.request_id in sched.running assert req in scheduled assert rejected == [] assert req.status == RequestStatus.RUNNING def test_skips_aborted_request(self): """Request whose state was cleared by abort is silently skipped.""" sched = _make_scheduler() req = _make_request("gone") # State NOT added to _prefill_states (simulates post-abort cleanup) sched.prefilling.append(req) scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) # Must not raise assert scheduled == [] assert rejected == [] assert len(sched.prefilling) == 0 def test_abort_during_chunk_discards_state(self): """_PrefillAbortedError from _step_prefill_chunk is swallowed cleanly.""" sched = _make_scheduler() req = _make_request("r1") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object( sched, "_step_prefill_chunk", side_effect=_PrefillAbortedError([], 4) ): scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) # Must not raise assert req.request_id not in sched._prefill_states assert req not in sched.prefilling assert scheduled == [] assert rejected == [] def test_runtime_error_surfaces_as_request_error(self): """A non-memory RuntimeError mid-chunk yields a finish_reason="error" RequestOutput immediately (only memory-pressure errors are requeued).""" sched = _make_scheduler() req = _make_request("oom") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object( sched, "_step_prefill_chunk", side_effect=RuntimeError("kernel panic") ): scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) assert req.request_id not in sched._prefill_states assert req not in sched.prefilling assert req.request_id not in sched.requests assert scheduled == [] assert len(rejected) == 1 out = rejected[0] assert out.request_id == "oom" assert out.finished is True assert out.finish_reason == "error" assert "kernel panic" in out.error def test_memory_error_requeues_instead_of_surfacing(self): """A memory-pressure RuntimeError mid-chunk requeues the request for a fresh attempt instead of immediately surfacing an error to the client.""" sched = _make_scheduler() req = _make_request("oom-mem") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object( sched, "_step_prefill_chunk", side_effect=RuntimeError("Memory limit exceeded during chunked prefill"), ): scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) # No client-facing error; the request is reset and back on the queue. assert rejected == [] assert req.request_id not in sched._prefill_states assert req not in sched.prefilling assert sched.requests.get(req.request_id) is req assert req in sched.waiting assert req.prefill_oom_retries == 1 def test_capacity_error_surfaces_as_typed_request_error(self): """A deterministic capacity rejection is not retried as transient OOM.""" sched = _make_scheduler() req = _make_request("capacity") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state err = PrefillMemoryExceededError( message="Prefill context too large for available memory", request_id=req.request_id, estimated_bytes=123, limit_bytes=100, ) with patch.object(sched, "_step_prefill_chunk", side_effect=err): scheduled = [] rejected = [] sched._advance_chunked_prefills(scheduled, rejected) assert scheduled == [] assert len(rejected) == 1 out = rejected[0] assert out.error == str(err) assert out.error_code == "prefill_memory_exceeded" assert out.error_metadata == { "request_id": req.request_id, "estimated_bytes": 123, "limit_bytes": 100, } assert req.prefill_oom_retries == 0 def test_multiple_requests_all_advanced(self): """All requests in prefilling get one chunk advanced per call.""" sched = _make_scheduler() reqs = [_make_request(f"r{i}") for i in range(3)] for req in reqs: sched.requests[req.request_id] = req state = _make_prefill_state(sched, req, n_remaining=20) state.sampler = MagicMock() state.sm = MagicMock() state.per_row_lps = [] sched.prefilling.append(req) sched._prefill_states[req.request_id] = state call_count = 0 def fake_step(state): nonlocal call_count call_count += 1 return False # All still in-progress with patch.object(sched, "_step_prefill_chunk", side_effect=fake_step): sched._advance_chunked_prefills([], []) assert call_count == 3 # One chunk per request # --------------------------------------------------------------------------- # _schedule_waiting(): chunked fork is taken for long prompts # --------------------------------------------------------------------------- class TestScheduleWaitingChunkedFork: def _setup(self, n_tokens: int, chunked: bool = True, step_size: int = 4): sched = _make_scheduler(chunked_prefill=chunked, step_size=step_size) req = _make_request("r1", n_tokens=n_tokens) sched.add_request(req) return sched, req def test_short_prompt_stays_on_normal_path(self): """Prompts that fit in one chunk use the normal prefill path.""" # step_size=4, prompt=3 tokens → not long enough to trigger chunked fork sched, req = self._setup(n_tokens=3, step_size=4) with patch.object( sched, "_do_external_prefill", return_value=([], [0]) ) as mock_ep: with patch.object(sched, "_begin_prefill") as mock_bp: sched._schedule_waiting() mock_ep.assert_called_once() mock_bp.assert_not_called() def test_long_prompt_enters_prefilling_queue(self): """Prompts longer than step_size+1 enter the chunked prefill queue.""" # step_size=4, 10 tokens → triggers chunked path sched, req = self._setup(n_tokens=10, step_size=4) with patch.object( sched, "_begin_prefill", return_value=_make_prefill_state(sched, req) ) as mock_bp: with patch.object(sched, "_step_prefill_chunk", return_value=False): sched._schedule_waiting() mock_bp.assert_called_once() assert req.request_id in sched._prefill_states assert req in sched.prefilling assert req.request_id not in sched.running def test_prefilling_request_counts_against_concurrency_cap(self): """A chunked prefill already in flight consumes a scheduler slot.""" sched = _make_scheduler(chunked_prefill=True, step_size=4) sched.config.max_num_seqs = 1 inflight = _make_request("inflight", n_tokens=10) sched.requests[inflight.request_id] = inflight sched.prefilling.append(inflight) sched._prefill_states[inflight.request_id] = _make_prefill_state( sched, inflight, ) queued = _make_request("queued", n_tokens=10) sched.add_request(queued) with patch.object(sched, "_begin_prefill") as mock_begin: scheduled, rejected = sched._schedule_waiting() mock_begin.assert_not_called() assert scheduled == [] assert rejected == [] assert queued in sched.waiting assert inflight in sched.prefilling def test_long_prompt_completes_in_first_chunk_goes_to_running(self): """If the first chunk happens to finish the prefill, request goes to running.""" sched, req = self._setup(n_tokens=10, step_size=4) fake_state = _make_prefill_state(sched, req, n_remaining=1) with patch.object(sched, "_begin_prefill", return_value=fake_state): with patch.object(sched, "_step_prefill_chunk", return_value=True): with patch.object(sched, "_emit_final_boundary_if_needed"): with patch("omlx.scheduler._sync_and_clear_cache"): sched._schedule_waiting() assert req.request_id not in sched._prefill_states assert req not in sched.prefilling assert req.request_id in sched.running def test_chunked_disabled_uses_normal_path(self): """chunked_prefill=False always uses the full-prefill path.""" sched, req = self._setup(n_tokens=100, chunked=False, step_size=4) with patch.object( sched, "_do_external_prefill", return_value=([], [0]) ) as mock_ep: with patch.object(sched, "_begin_prefill") as mock_bp: sched._schedule_waiting() mock_ep.assert_called_once() mock_bp.assert_not_called() def test_non_chunked_path_runtime_error_cleans_up_and_rejects(self): """RuntimeError from _do_external_prefill in the non-chunked path must pop self.requests, drop the temp uid mappings, remove the PrefillProgressTracker entry, and emit a finish_reason=\"error\" RequestOutput so the client sees the failure (#1405).""" from omlx.prefill_progress import get_prefill_tracker sched, req = self._setup(n_tokens=3, step_size=4) rid = req.request_id tracker = get_prefill_tracker() tracker.clear() tracker.update(rid, processed=1, total=3, model_id="test") assert tracker.get_model_progress("test"), "tracker entry not set up" try: with patch.object( sched, "_do_external_prefill", side_effect=RuntimeError("Memory limit exceeded during prefill"), ): scheduled, rejected = sched._schedule_waiting() assert rid not in sched.requests assert rid not in sched.request_id_to_uid assert not any(v == rid for v in sched.uid_to_request_id.values()) assert tracker.get_model_progress("test") == [] assert scheduled == [] assert len(rejected) == 1 out = rejected[0] assert out.request_id == rid assert out.finished is True assert out.finish_reason == "error" assert "Memory limit" in out.error finally: tracker.clear() def _setup_throttle(self, max_bytes_gb=10, hard_cap_gb=12): """Build a scheduler with watermark fields set for throttle tests.""" sched = _make_scheduler() sched._memory_limit_bytes = max_bytes_gb * 1024**3 sched._memory_hard_limit_bytes = hard_cap_gb * 1024**3 sched._prefill_safe_zone_ratio = 0.80 sched._prefill_min_chunk_tokens = 32 return sched def _mock_current(self, sched, current_gb): """Context manager-ish — patch both memory probes to current_gb.""" target = int(current_gb * 1024**3) return patch("omlx.scheduler.mx.get_active_memory", return_value=target), patch( "omlx.scheduler.get_phys_footprint", return_value=target ) def test_adaptive_throttle_below_soft_watermark_passthrough(self): """current < soft watermark → no throttle, full chunk.""" sched = self._setup_throttle(max_bytes_gb=10, hard_cap_gb=12) # soft_watermark = 10 * 0.80 = 8 GB; current 5 GB is below a, b = self._mock_current(sched, 5) with a, b: result = sched._adaptive_chunk_size( 2048, request_id="r1", loop_label="external" ) assert result == 2048 def test_adaptive_throttle_tier_1024(self): """First quarter of the soft-to-hard band → 1024.""" sched = self._setup_throttle(max_bytes_gb=10, hard_cap_gb=12) # soft_wm = 8 GB, band = 12 - 8 = 4 GB. 10% into band = 8.4 GB. a, b = self._mock_current(sched, 8.4) with a, b: result = sched._adaptive_chunk_size( 2048, request_id="r1", loop_label="external" ) assert result == 1024 def test_adaptive_throttle_tier_512(self): """50%+ of band → 512.""" sched = self._setup_throttle(max_bytes_gb=10, hard_cap_gb=12) # 60% of band: 8 + 4*0.60 = 10.4 GB a, b = self._mock_current(sched, 10.4) with a, b: result = sched._adaptive_chunk_size( 2048, request_id="r1", loop_label="external" ) assert result == 512 def test_adaptive_throttle_requested_smaller_than_tier(self): """Requested chunk already smaller than the tier target → pass through.""" sched = self._setup_throttle(max_bytes_gb=10, hard_cap_gb=12) # 60% of band → tier 512. But requested=256 < 512. a, b = self._mock_current(sched, 10.4) with a, b: result = sched._adaptive_chunk_size( 256, request_id="r1", loop_label="external" ) assert result == 256 def test_adaptive_throttle_no_cap_passthrough(self): """When hard limit or soft base is unset (=0), no throttle.""" sched = self._setup_throttle() sched._memory_hard_limit_bytes = 0 result = sched._adaptive_chunk_size( 2048, request_id="r1", loop_label="external" ) assert result == 2048 sched._memory_hard_limit_bytes = 10 * 1024**3 sched._memory_limit_bytes = 0 result = sched._adaptive_chunk_size( 2048, request_id="r1", loop_label="external" ) assert result == 2048 def test_chunked_first_chunk_runtime_error_cleans_up_and_rejects(self): """RuntimeError on the chunked first chunk must pop self.requests, remove the PrefillProgressTracker entry, and emit an error RequestOutput. _step_prefill_chunk updates the tracker before the hard-limit check, so without this catch the entry would leak (#1405).""" from omlx.prefill_progress import get_prefill_tracker sched, req = self._setup(n_tokens=10, step_size=4) rid = req.request_id tracker = get_prefill_tracker() tracker.clear() tracker.update(rid, processed=2, total=10, model_id="test") assert tracker.get_model_progress("test"), "tracker entry not set up" try: with patch.object( sched, "_begin_prefill", return_value=_make_prefill_state(sched, req), ): with patch.object( sched, "_step_prefill_chunk", side_effect=RuntimeError( "Memory limit exceeded during chunked prefill" ), ): scheduled, rejected = sched._schedule_waiting() assert rid not in sched.requests assert rid not in sched._prefill_states assert req not in sched.prefilling assert tracker.get_model_progress("test") == [] assert scheduled == [] assert len(rejected) == 1 out = rejected[0] assert out.request_id == rid assert out.finished is True assert out.finish_reason == "error" assert "Memory limit" in out.error finally: tracker.clear() # --------------------------------------------------------------------------- # Prefill-rejection paged-cache cleanup # --------------------------------------------------------------------------- class TestPrefillRejectionReleasesPagedCache: """Rejection paths must release block_aware_cache refs / paged_cache block_table entries that ``add_request`` populated via ``fetch_cache``. Without this, every rejected request leaks an entry in ``BlockAwarePrefixCache._request_tables`` plus the ref counts on its prefix-matched blocks — pinning the paged cache and compounding the very memory pressure that triggered the rejection. The existing ``self.requests.pop(...)`` and ``get_prefill_tracker().remove(...)`` cleanups handle scheduler-side state but never reach into the paged-cache layer. """ def test_helper_calls_block_aware_cache_release(self): """The helper delegates to block_aware_cache.release_cache when one is attached — the normal production wiring.""" sched = _make_scheduler() sched.block_aware_cache = MagicMock() sched.paged_cache_manager = MagicMock() sched._release_paged_cache_for_request("rid-1") sched.block_aware_cache.release_cache.assert_called_once_with("rid-1") # release_cache delegates to delete_block_table internally; the # helper must NOT also call it directly (double-delete). sched.paged_cache_manager.delete_block_table.assert_not_called() def test_helper_falls_back_to_paged_cache_manager(self): """Without a BlockAwarePrefixCache, fall back to deleting the block table directly on the paged cache manager.""" sched = _make_scheduler() sched.block_aware_cache = None sched.paged_cache_manager = MagicMock() sched._release_paged_cache_for_request("rid-2") sched.paged_cache_manager.delete_block_table.assert_called_once_with("rid-2") def test_helper_is_noop_without_any_paged_cache(self): """No paged-cache layer attached → silent no-op.""" sched = _make_scheduler() sched.block_aware_cache = None sched.paged_cache_manager = None # Should not raise. sched._release_paged_cache_for_request("rid-3") def test_advance_chunked_prefills_releases_on_runtime_error(self): """_advance_chunked_prefills' RuntimeError handler must call release_cache so the paged-cache block refs from the request's prefix-cache lookup don't leak.""" sched = _make_scheduler() sched.block_aware_cache = MagicMock() req = _make_request("oom-chunked") sched.requests[req.request_id] = req state = _make_prefill_state(sched, req) sched.prefilling.append(req) sched._prefill_states[req.request_id] = state with patch.object( sched, "_step_prefill_chunk", side_effect=RuntimeError("Memory limit exceeded"), ): sched._advance_chunked_prefills([], []) sched.block_aware_cache.release_cache.assert_called_once_with("oom-chunked") def test_schedule_waiting_non_chunked_releases_on_runtime_error(self): """The non-chunked _do_external_prefill rejection path must release the paged-cache footprint before popping self.requests.""" sched = _make_scheduler(step_size=4) sched.block_aware_cache = MagicMock() # No prefix-cache hit: fetch_cache returns (None, prompt_tokens) so # add_request falls through to the waiting queue without trying to # preload/reconstruct. sched.block_aware_cache.fetch_cache.return_value = (None, [0, 1, 2]) req = _make_request("oom-direct", n_tokens=3) sched.add_request(req) sched.block_aware_cache.reset_mock() with patch.object( sched, "_do_external_prefill", side_effect=RuntimeError("kernel panic"), ): sched._schedule_waiting() sched.block_aware_cache.release_cache.assert_called_once_with("oom-direct") def test_schedule_waiting_chunked_first_chunk_releases_on_runtime_error(self): """The chunked first-chunk rejection path must release the paged-cache footprint before popping self.requests.""" sched = _make_scheduler(step_size=4) sched.block_aware_cache = MagicMock() sched.block_aware_cache.fetch_cache.return_value = (None, list(range(10))) req = _make_request("oom-first-chunk", n_tokens=10) sched.add_request(req) sched.block_aware_cache.reset_mock() with patch.object( sched, "_begin_prefill", return_value=_make_prefill_state(sched, req), ): with patch.object( sched, "_step_prefill_chunk", side_effect=RuntimeError("kernel panic"), ): sched._schedule_waiting() sched.block_aware_cache.release_cache.assert_called_once_with("oom-first-chunk") def test_schedule_waiting_preflight_rejection_releases(self): """_preflight_memory_check rejection (the non-RuntimeError path inside _schedule_waiting) must also release the paged-cache footprint. Same leak shape as the RuntimeError rejections — the request reached this point via add_request → fetch_cache so _request_tables is populated and prefix block refs are held.""" sched = _make_scheduler(step_size=4) sched.block_aware_cache = MagicMock() sched.block_aware_cache.fetch_cache.return_value = (None, list(range(5))) req = _make_request("oom-preflight", n_tokens=5) sched.add_request(req) sched.block_aware_cache.reset_mock() from omlx.scheduler import _PreflightRejection with patch.object( sched, "_preflight_memory_check", return_value=_PreflightRejection( message="Memory limit exceeded by preflight estimate", estimated_bytes=1, limit_bytes=1, ), ): scheduled, rejected = sched._schedule_waiting() assert scheduled == [] assert len(rejected) == 1 assert rejected[0].request_id == "oom-preflight" assert rejected[0].finish_reason == "error" sched.block_aware_cache.release_cache.assert_called_once_with("oom-preflight") # --------------------------------------------------------------------------- # First-chunk eviction pause must preserve a reconstructed prefix (#2180) # --------------------------------------------------------------------------- class TestFirstChunkEvictionPreservesPrefix: def test_first_chunk_eviction_pause_keeps_reconstructed_prefix(self): """_PrefillEvictionNeeded raised before the first chunk's forward pass must not discard a reconstructed SSD prefix. The eviction pause keeps prompt_cache / block_table / cached_tokens / remaining_tokens attached, so when no idle model can be evicted the retry prefills only the uncached suffix instead of recomputing the whole prompt cold (#2180).""" sched = _make_scheduler(step_size=4) sched.block_aware_cache = MagicMock() sched.block_aware_cache.fetch_cache.return_value = (None, list(range(100))) req = _make_request("evict-first-chunk", n_tokens=100) sched.add_request(req) sched.block_aware_cache.reset_mock() # Simulate the state _prepare_prefix_cache_for_request leaves after a # successful paged/SSD cache hit + reconstruction: 90 cached tokens, # a 10-token uncached suffix, and a live block table. prompt_cache = [MagicMock()] block_table = MagicMock() sched._prefix_cache_prepared.add(req.request_id) req.prompt_cache = prompt_cache req.cached_tokens = 90 req.remaining_tokens = req.prompt_token_ids[90:] req.block_table = block_table req.shared_prefix_blocks = 3 eviction = PrefillEvictionRequest( request_id=req.request_id, model_id="test", current_bytes=1, target_cap_bytes=1, predicted_transient_bytes=1, requested_tokens=4, reason="adaptive_prefill_throttle", ) with patch.object( sched, "_begin_prefill", return_value=_make_prefill_state(sched, req), ): with patch.object( sched, "_step_prefill_chunk", side_effect=_PrefillEvictionNeeded(eviction), ): scheduled, rejected = sched._schedule_waiting() assert scheduled == [] assert rejected == [] # Paused back into the waiting queue with the eviction request pending. assert req in sched.waiting assert sched._pending_prefill_eviction_request is eviction # The reconstructed prefix must survive the pause untouched. assert req.prompt_cache is prompt_cache assert req.cached_tokens == 90 assert req.remaining_tokens == req.prompt_token_ids[90:] assert req.block_table is block_table assert req.shared_prefix_blocks == 3 sched.block_aware_cache.release_cache.assert_not_called()