From 7796634492adb72e33a107261ae2f2534a720903 Mon Sep 17 00:00:00 2001 From: GokayAI <60583610+gokay-ai@users.noreply.github.com> Date: Fri, 18 Sep 2026 10:40:53 +0000 Subject: [PATCH] fix(distserve): free Prefill cache with scheduler session id The DistServe proxy used the OpenAI response id as remote_session_id. That id is the client session_id (default -1) or a chatcmpl-* string, not the Prefill scheduler key, so cache-free logged invalid free and leaked preserved Prefill sessions. Return cache_session_id from Prefill and use it for migration/free. Fixes #4967 Signed-off-by: GokayAI <60583610+gokay-ai@users.noreply.github.com> --- .../serve/openai/chat_completions/serving.py | 6 + .../serve/openai/endpoints/completions.py | 8 +- lmdeploy/serve/proxy/proxy.py | 33 ++- .../serve/test_distserve_cache_session_id.py | 227 ++++++++++++++++++ 4 files changed, 266 insertions(+), 8 deletions(-) create mode 100644 tests/test_lmdeploy/serve/test_distserve_cache_session_id.py diff --git a/lmdeploy/serve/openai/chat_completions/serving.py b/lmdeploy/serve/openai/chat_completions/serving.py index ea57c43221..4a14c6f30c 100644 --- a/lmdeploy/serve/openai/chat_completions/serving.py +++ b/lmdeploy/serve/openai/chat_completions/serving.py @@ -303,6 +303,9 @@ async def _completion_stream_generator() -> AsyncGenerator[str, None]: routed_experts=chunk.routed_experts, output_ids=stream_output_ids) if chunk.cache_block_ids is not None and chunk.is_last_delta: + # DistServe cache-free uses the engine session id, + # not the public chatcmpl-* response id. + response_json['cache_session_id'] = chat_runner.session.session_id response_json['cache_block_ids'] = chunk.cache_block_ids response_json['remote_token_ids'] = chunk.token_ids yield f'data: {json.dumps(response_json)}\n\n' @@ -383,6 +386,9 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: ).model_dump() if with_cache: + # DistServe cache-free uses the engine session id, not the public + # chatcmpl-* response id. + response['cache_session_id'] = chat_runner.session.session_id response['cache_block_ids'] = cache_block_ids response['remote_token_ids'] = remote_token_ids diff --git a/lmdeploy/serve/openai/endpoints/completions.py b/lmdeploy/serve/openai/endpoints/completions.py index 6c704f47db..7b9789c429 100644 --- a/lmdeploy/serve/openai/endpoints/completions.py +++ b/lmdeploy/serve/openai/endpoints/completions.py @@ -223,7 +223,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: # First chunk with role final_usage = UsageInfo() if include_usage else None try: - for generator in generators: + for generator, session in zip(generators, sessions): async for res in generator: logprobs = None if request.logprobs and res.logprobs: @@ -243,6 +243,9 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: finish_reason=res.finish_reason, logprobs=logprobs) if res.cache_block_ids is not None: + # DistServe cache-free uses the engine session id, + # not the public OpenAI response id. + response_json['cache_session_id'] = session.session_id response_json['cache_block_ids'] = res.cache_block_ids response_json['remote_token_ids'] = res.token_ids yield f'data: {json.dumps(response_json)}\n\n' @@ -333,6 +336,9 @@ async def _inner_call(i, generator, session): ).model_dump() if with_cache: + # DistServe cache-free uses the engine session id, not the public + # OpenAI response id (client session_id, default -1). + response['cache_session_id'] = sessions[0].session_id response['cache_block_ids'] = cache_block_ids response['remote_token_ids'] = remote_token_ids diff --git a/lmdeploy/serve/proxy/proxy.py b/lmdeploy/serve/proxy/proxy.py index 7b39790ae0..8d020c8507 100644 --- a/lmdeploy/serve/proxy/proxy.py +++ b/lmdeploy/serve/proxy/proxy.py @@ -564,6 +564,22 @@ async def connection_warmup(): return JSONResponse({'SUCCESS': True}) +def _distserve_remote_session_id(prefill_info: dict, dummy_prefill: bool): + """Return the Prefill scheduler session id used for DistServe cache-free. + + The OpenAI response ``id`` is the client ``session_id`` (default ``-1``) or + a ``chatcmpl-*`` string. Cache-free looks up ``scheduler.sessions[id]``, so + those public ids miss the real Prefill session and leak preserved metadata. + """ + if dummy_prefill: + return 0 + cache_session_id = prefill_info.get('cache_session_id') + if type(cache_session_id) is not int or cache_session_id < 0: + return create_error_response(HTTPStatus.BAD_GATEWAY, + 'Prefill response is missing a valid cache_session_id.') + return cache_session_id + + @app.post('/distserve/gc', dependencies=[Depends(validate_json_request)]) async def cache_block_gc_to_be_migrated(): # TODO (JimyMa): add garbage collection of to be migrated request @@ -710,7 +726,9 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque rdma_config=node_manager.rdma_config, )) - remote_session_id = int(prefill_info.get('id')) if prefill_info.get('id') else 0 + remote_session_id = _distserve_remote_session_id(prefill_info, node_manager.dummy_prefill) + if not isinstance(remote_session_id, int): + return remote_session_id remote_block_ids = prefill_info.get('cache_block_ids') or [] remote_token_id = prefill_info.get('remote_token_ids')[-1] if prefill_info.get('remote_token_ids') else 0 @@ -724,7 +742,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque start = node_manager.pre_call(d_url) if not node_manager.dummy_prefill: - node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id']) + node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), remote_session_id) if request.stream is True: response = node_manager.stream_generate(request_dict, d_url, '/v1/chat/completions') background_task = node_manager.create_background_tasks(d_url, start) @@ -735,7 +753,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque resp = JSONResponse(json.loads(response)) if not node_manager.dummy_prefill: - node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), prefill_info['id']) + node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), remote_session_id) return resp @@ -855,9 +873,10 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None except Exception as e: logger.error(f'error Msg: {str(e)}') return {'status': f'Connection error, cannot establish connection {(p_url, d_url)}'} - node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id']) - remote_session_id = int(prefill_info.get('id')) if prefill_info.get('id') else 0 + remote_session_id = _distserve_remote_session_id(prefill_info, node_manager.dummy_prefill) + if not isinstance(remote_session_id, int): + return remote_session_id remote_block_ids = prefill_info.get('cache_block_ids') or [] remote_token_id = prefill_info.get('remote_token_ids')[-1] if prefill_info.get('remote_token_ids') else 0 request_dict['migration_request'] = MigrationRequest( @@ -870,7 +889,7 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None start = node_manager.pre_call(d_url) if not node_manager.dummy_prefill: - node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id']) + node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), remote_session_id) if request.stream is True: response = node_manager.stream_generate(request_dict, d_url, '/v1/completions') background_task = node_manager.create_background_tasks(d_url, start) @@ -880,7 +899,7 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None node_manager.post_call(d_url, start) resp = JSONResponse(json.loads(response)) if not node_manager.dummy_prefill: - node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), prefill_info.get('id')) + node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), remote_session_id) return resp else: raise ValueError(f'No serving strategy named {node_manager.serving_strategy}') diff --git a/tests/test_lmdeploy/serve/test_distserve_cache_session_id.py b/tests/test_lmdeploy/serve/test_distserve_cache_session_id.py new file mode 100644 index 0000000000..5330104fc5 --- /dev/null +++ b/tests/test_lmdeploy/serve/test_distserve_cache_session_id.py @@ -0,0 +1,227 @@ +"""Regression tests for DistServe Prefill cache-free session identity. + +The OpenAI response ``id`` is the client session_id (default -1). DistServe +cache-free looks up ``scheduler.sessions[remote_session_id]``, so using that +public id misses the real Prefill session and leaks preserved metadata. +""" +import asyncio +import json +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +from fastapi import APIRouter + +from lmdeploy.pytorch.config import CacheConfig, SchedulerConfig +from lmdeploy.pytorch.disagg.config import EngineRole, ServingStrategy +from lmdeploy.pytorch.disagg.conn.protocol import MigrationProtocol +from lmdeploy.pytorch.disagg.conn.proxy_conn import PDConnectionPool +from lmdeploy.pytorch.messages import SequenceMeta +from lmdeploy.pytorch.paging.scheduler import Scheduler +from lmdeploy.pytorch.paging.seq_states.states import ToBeMigratedState +from lmdeploy.pytorch.strategies.ar.sequence import ARSequenceStrategy +from lmdeploy.serve.openai.endpoints.completions import register +from lmdeploy.serve.openai.protocol import CompletionRequest +from lmdeploy.serve.proxy import proxy as proxy_mod + + +class _Session: + + def __init__(self, session_id): + self.session_id = session_id + + async def async_abort(self): + pass + + +class _SessionManager: + + def __init__(self): + self.removed = [] + + def has(self, session_id): + return False + + def remove(self, session): + self.removed.append(session) + + +class _RawRequest: + + def __init__(self, payload=None): + self._payload = payload or {} + + async def json(self): + return dict(self._payload) + + async def is_disconnected(self): + return False + + +def _scheduler(): + return Scheduler( + SchedulerConfig(max_batches=4, max_session_len=128, max_request_output_len=64), + CacheConfig(max_batches=4, block_size=16, num_cpu_blocks=0, num_gpu_blocks=16), + SequenceMeta(16, strategy=ARSequenceStrategy()), + ) + + +def _cache_free(scheduler, session_id): + """Mirror EngineP2PConnection.handle_zmq_recv cache-free lookup.""" + if session_id in scheduler.sessions: + scheduler.end_session(session_id) + return True + return False + + +def test_public_completion_id_misses_preserved_prefill_session(): + scheduler = _scheduler() + session = scheduler.add_session(7) + seq = session.add_sequence([1, 2, 3], preserve_cache=True) + # Prefill finish with preserve_cache keeps the session for migration. + seq.state.to_state(ToBeMigratedState) + + public_id = -1 + assert not _cache_free(scheduler, public_id) + assert 7 in scheduler.sessions + assert seq.seq_id in scheduler.seq_manager._seq_map + + assert _cache_free(scheduler, 7) + assert 7 not in scheduler.sessions + assert seq.seq_id not in scheduler.seq_manager._seq_map + + +def test_completions_with_cache_returns_engine_session_id_not_public_id(): + engine_session_id = 7 + + class _AsyncEngine: + model_name = 'fake-model' + + async def preprocess(self, prompt, session, **kwargs): + return SimpleNamespace(session=session, gen_config=kwargs['gen_config']) + + async def generate(self, prepared, **kwargs): + yield SimpleNamespace(response='ok', + token_ids=[1], + input_token_len=2, + generate_token_len=1, + finish_reason='stop', + logprobs=None, + cache_block_ids=[0], + cached_tokens=0) + + class _ServerContext: + + def __init__(self): + self.async_engine = _AsyncEngine() + self.engine_config = SimpleNamespace(logprobs_mode=None, adapters=[]) + self.session_manager = _SessionManager() + self.default_gen_config = {} + + def create_session(self, session_id=None): + return _Session(engine_session_id) + + async def _run(): + context = _ServerContext() + router = APIRouter() + register(router, context) + endpoint = router.routes[0].endpoint + return await endpoint( + CompletionRequest(model='fake-model', prompt='hi', session_id=-1, max_tokens=1), + _RawRequest({'with_cache': True, 'preserve_cache': True}), + ) + + response = asyncio.run(_run()) + assert response['id'] == '-1' + assert response['cache_session_id'] == engine_session_id + + +def _proxy_manager(generate, stream_generate=None, pool=None): + pool = pool or PDConnectionPool() + pool.is_connected = lambda *args: True + + async def check_model(model): + return None + + return SimpleNamespace( + serving_strategy=ServingStrategy.DistServe, + check_request_model=check_model, + dummy_prefill=False, + migration_protocol=MigrationProtocol.RDMA, + rdma_config=None, + pd_connection_pool=pool, + generate=generate, + stream_generate=stream_generate, + get_node_url=lambda model, role: 'p' if role == EngineRole.Prefill else 'd', + pre_call=lambda url: 0, + post_call=lambda *args: None, + create_background_tasks=lambda *args: None, + ) + + +@pytest.mark.parametrize('stream', [False, True]) +def test_proxy_completions_uses_cache_session_id_for_free(monkeypatch, stream): + owners = [] + pool = PDConnectionPool() + pool.is_connected = lambda *args: True + pool.shelf_prefill_session = Mock(wraps=pool.shelf_prefill_session) + pool.unshelf_prefill_session = Mock(wraps=pool.unshelf_prefill_session) + + async def generate(payload, url, endpoint): + if url == 'd': + owners.append(payload['migration_request']['remote_session_id']) + return json.dumps({'choices': [{'finish_reason': 'stop'}]}) + return json.dumps({ + 'id': '-1', + 'cache_session_id': 7, + 'cache_block_ids': [0], + 'remote_token_ids': [1], + }) + + async def stream_generate(payload, url, endpoint): + owners.append(payload['migration_request']['remote_session_id']) + yield b'data: {"choices": [{"finish_reason": "stop"}]}\n\n' + yield b'data: [DONE]\n\n' + + monkeypatch.setattr( + proxy_mod, 'node_manager', + _proxy_manager(generate, stream_generate=stream_generate, pool=pool)) + + async def _run(): + response = await proxy_mod.completions_v1( + CompletionRequest(model='fake-model', prompt='hi', max_tokens=1, stream=stream)) + if stream: + async for _ in response.body_iterator: + pass + return response + + response = asyncio.run(_run()) + assert response.status_code == 200 + assert owners == [7] + pool.shelf_prefill_session.assert_called_once_with(('p', 'd'), 7) + pool.unshelf_prefill_session.assert_called_once_with(('p', 'd'), 7) + + +@pytest.mark.parametrize('owner', [None, -1, '-1', True]) +def test_proxy_completions_rejects_public_id_fallback(monkeypatch, owner): + decode_calls = [] + pool = PDConnectionPool() + pool.is_connected = lambda *args: True + + async def generate(payload, url, endpoint): + decode_calls.append(url) + result = {'id': '-1', 'cache_block_ids': [0], 'remote_token_ids': [1]} + if owner is not None: + result['cache_session_id'] = owner + return json.dumps(result) + + monkeypatch.setattr(proxy_mod, 'node_manager', _proxy_manager(generate, pool=pool)) + + async def _run(): + return await proxy_mod.completions_v1(CompletionRequest(model='fake-model', prompt='hi', max_tokens=1)) + + response = asyncio.run(_run()) + assert response.status_code == 502 + assert b'cache_session_id' in response.body + assert decode_calls == ['p'] + assert not pool.migration_session_shelf