diff --git a/lmdeploy/pytorch/engine/inputs_maker.py b/lmdeploy/pytorch/engine/inputs_maker.py index b35fb3836f..6f93285f9a 100644 --- a/lmdeploy/pytorch/engine/inputs_maker.py +++ b/lmdeploy/pytorch/engine/inputs_maker.py @@ -1229,16 +1229,17 @@ def create_model_inputs_delta_valid_only(self): def update_running_seqs(self, running: 'SeqList', inputs: 'ModelInputs|None'): """Update running seqs.""" - if self.config.role == EngineRole.Prefill: - # p node will not update running seqs - return - is_decoding = inputs is None if self.long_context_chunker.enabled() and not is_decoding and inputs.is_chunk: # long context chunk does not need to update running seqs self.long_context_chunker.update_step(inputs) return + # Prefill workers must advance chunk history above, but must not add + # completed prefills to the decode-running list. + if self.config.role == EngineRole.Prefill: + return + if is_decoding: self.running_seqs = running else: diff --git a/lmdeploy/serve/openai/chat_completions/serving.py b/lmdeploy/serve/openai/chat_completions/serving.py index ea57c43221..75652b69ba 100644 --- a/lmdeploy/serve/openai/chat_completions/serving.py +++ b/lmdeploy/serve/openai/chat_completions/serving.py @@ -303,6 +303,7 @@ 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: + 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 +384,8 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: ).model_dump() if with_cache: + # Keep the OpenAI response ID independent of the engine's cache owner. + 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..91e7d94841 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,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: finish_reason=res.finish_reason, logprobs=logprobs) if res.cache_block_ids is not None: + 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 +334,8 @@ async def _inner_call(i, generator, session): ).model_dump() if with_cache: + # Cache ownership uses the engine's ID, not the public response ID. + 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..32888fcb3f 100644 --- a/lmdeploy/serve/proxy/proxy.py +++ b/lmdeploy/serve/proxy/proxy.py @@ -710,7 +710,12 @@ 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 + # The public response ID is not the engine's cache owner. + remote_session_id = 0 if node_manager.dummy_prefill else prefill_info.get('cache_session_id') + if type(remote_session_id) is not int or remote_session_id < 0: + return create_error_response(HTTPStatus.BAD_GATEWAY, + 'Prefill response is missing a valid cache_session_id; ' + 'upgrade the Prefill server together with the proxy.') 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 +729,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 +740,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 +860,13 @@ 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 + # The public response ID is not the engine's cache owner. + remote_session_id = 0 if node_manager.dummy_prefill else prefill_info.get('cache_session_id') + if type(remote_session_id) is not int or remote_session_id < 0: + return create_error_response(HTTPStatus.BAD_GATEWAY, + 'Prefill response is missing a valid cache_session_id; ' + 'upgrade the Prefill server together with the proxy.') 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 +879,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 +889,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/pytorch/engine/test_inputs_maker.py b/tests/pytorch/engine/test_inputs_maker.py index 654cee8e13..c5ba3b2a4c 100644 --- a/tests/pytorch/engine/test_inputs_maker.py +++ b/tests/pytorch/engine/test_inputs_maker.py @@ -1388,6 +1388,38 @@ def has_waiting_long_prefill(self): assert calls == [(True, True), (False, False)] +@pytest.mark.parametrize('role', list(EngineRole)) +@pytest.mark.parametrize('length', [511, 512, 513, 1106]) +def test_chunk_history_advances_for_all_roles_without_decoding_on_prefill(role, length): + from lmdeploy.pytorch.config import CacheConfig, SchedulerConfig + from lmdeploy.pytorch.messages import SequenceMeta + from lmdeploy.pytorch.paging import Scheduler + from lmdeploy.pytorch.strategies.ar.sequence import ARSequenceStrategy + + scheduler = Scheduler( + SchedulerConfig(max_batches=4, max_session_len=4096, max_request_output_len=2048), + CacheConfig(max_batches=4, block_size=16, num_cpu_blocks=0, num_gpu_blocks=256), + SequenceMeta(16, strategy=ARSequenceStrategy())) + seq = scheduler.add_session(0).add_sequence(list(range(length)), preserve_cache=True) + maker = InputsMakerAsync.__new__(InputsMakerAsync) + maker.config = SimpleNamespace(role=role) + maker.running_seqs = [] + maker.long_context_chunker = LongContextChunker(512) + chunker = maker.long_context_chunker + chunker.set_seq(seq) + for expected in range(512, length, 512): + assert not chunker.is_last_chunk() + size, _ = chunker.next_chunk_size() + maker.update_running_seqs([seq], SimpleNamespace(is_chunk=True, max_q_seqlen=size)) + assert seq.num_history_ids == expected + assert maker.running_seqs == [] + assert chunker.is_last_chunk() + # The input builder clears the chunker before submitting the final chunk. + chunker.clear() + maker.update_running_seqs([seq], SimpleNamespace(is_chunk=True)) + assert maker.running_seqs == ([] if role == EngineRole.Prefill else [seq]) + + def test_normal_prefill_can_update_running_while_long_chunker_is_active(): long_seq = _DummySeq(history_ids=0, token_ids=1024, all_multimodals={}, input_multimodals={}) short_seq = _DummySeq(history_ids=0, token_ids=16, all_multimodals={}, input_multimodals={}) diff --git a/tests/test_lmdeploy/serve/openai/test_distserve.py b/tests/test_lmdeploy/serve/openai/test_distserve.py new file mode 100644 index 0000000000..cc802b043f --- /dev/null +++ b/tests/test_lmdeploy/serve/openai/test_distserve.py @@ -0,0 +1,224 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Cache-owner IDs travel from the real API to migration ACK and scheduler GC. + +Only generation/transport are simulated; no model, CUDA, or RDMA is needed. +""" +import asyncio +import importlib +import json +from types import SimpleNamespace +from unittest.mock import Mock + +import httpx +import pytest +from fastapi import APIRouter, FastAPI + +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.engine.engine import Engine +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.managers import SessionManager +from lmdeploy.serve.openai.api_server import ServerContext +from lmdeploy.serve.openai.chat_completions import register as register_chat +from lmdeploy.serve.openai.endpoints.completions import register as register_completion +from lmdeploy.serve.openai.protocol import ChatCompletionRequest, CompletionRequest, DeltaMessage + + +class _Parser: + tool_parser_cls = None + + def __init__(self, request): + self.request = request + self.tool_parser = None + self.reasoning_tokens = 0 + + def stream_chunk(self, text, token_ids, **kwargs): + return [(DeltaMessage(content=text), False)] + + def parse_complete(self, text, token_ids=None): + return text, None, None + + def validate_complete(self, **kwargs): + return True + + +@pytest.fixture +def prefill(): + scheduler = 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())) + engine = object.__new__(Engine) + engine.scheduler = scheduler + + class Generation: + model_name = 'fake-model' + epoch = 0 + backend_config = SimpleNamespace(role=EngineRole.Prefill, logprobs_mode=None, adapters=[]) + tokenizer = SimpleNamespace(model=SimpleNamespace(model=SimpleNamespace(model=None))) + + def __init__(self): + self.session_mgr = SessionManager() + self.ids = [] + + async def preprocess(self, prompt, session, **kwargs): + return SimpleNamespace(session=session, gen_config=kwargs['gen_config']) + + async def generate(self, prepared, **kwargs): + sid = prepared.session.session_id + self.ids.append(sid) + seq = scheduler.add_session(sid).add_sequence([1, 2], preserve_cache=prepared.gen_config.preserve_cache) + seq.state.to_state(ToBeMigratedState) + engine._on_end_session([SimpleNamespace(data={'session_id': sid, 'response': False})]) + 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, + routed_experts=None) + + context = ServerContext() + context.async_engine = Generation() + context.response_parser_cls = _Parser + router = APIRouter() + register_chat(router, context) + register_completion(router, context) + app = FastAPI() + app.include_router(router) + return context, engine, app + + +def _request(chat, **kwargs): + if chat: + return ChatCompletionRequest(model='fake-model', messages=[dict(role='user', content='hi')], **kwargs) + return CompletionRequest(model='fake-model', prompt='hi', **kwargs) + + +@pytest.mark.parametrize('chat', [False, True]) +@pytest.mark.parametrize('user_id', [-1, 42]) +@pytest.mark.parametrize('stream', [False, True]) +def test_prefill_cache_handle_is_not_the_public_response_id(prefill, chat, user_id, stream): + context, engine, app = prefill + + async def run(): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://prefill') as client: + for _ in range(3): + payload = _request(chat, session_id=user_id, stream=stream).model_dump() + payload.update(with_cache=True, preserve_cache=True) + response = await client.post('/v1/chat/completions' if chat else '/v1/completions', json=payload) + assert response.status_code == 200, response.text + if stream: + chunks = [json.loads(line[6:]) for line in response.text.splitlines() + if line.startswith('data: ') and line != 'data: [DONE]'] + result = next(chunk for chunk in chunks if 'cache_session_id' in chunk) + else: + result = response.json() + owner = result['cache_session_id'] + assert owner == context.async_engine.ids[-1] + assert result['id'].startswith('chatcmpl-') if chat else result['id'] == str(user_id) + assert len(engine.scheduler.sessions) == 1 + assert engine.end_session(owner) + assert not engine.scheduler.sessions + assert not engine.scheduler.seq_manager._seq_map + assert not context.session_manager.sessions + + asyncio.run(run()) + + +@pytest.mark.parametrize('chat', [False, True]) +@pytest.mark.parametrize('user_id', [-1, 42]) +@pytest.mark.parametrize('stream', [False, True]) +def test_proxy_passes_internal_cache_owner_to_decode(prefill, monkeypatch, chat, user_id, stream): + proxy = importlib.import_module('lmdeploy.serve.proxy.proxy') + context, engine, app = prefill + 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) + owners = [] + + async def check_model(model): + return None + + def acknowledge(payload): + owner = payload['migration_request']['remote_session_id'] + assert owner == context.async_engine.ids[-1] + owners.append(owner) + assert engine.end_session(owner) + + async def generate(payload, url, endpoint): + if url == 'd': + acknowledge(payload) + return json.dumps(dict(choices=[dict(finish_reason='stop')])) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://prefill') as client: + response = await client.post(endpoint, json=payload) + assert response.status_code == 200 + return response.text + + async def stream_generate(payload, url, endpoint): + acknowledge(payload) + yield b'data: {"choices": [{"finish_reason": "stop"}]}\n\n' + yield b'data: [DONE]\n\n' + + manager = 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) + monkeypatch.setattr(proxy, 'node_manager', manager) + + async def run(): + request = _request(chat, session_id=user_id, stream=stream, max_tokens=8) + handler = proxy.chat_completions_v1 if chat else proxy.completions_v1 + response = await handler(request) + assert response.status_code == 200 + if stream: + async for _ in response.body_iterator: + pass + assert owners == [context.async_engine.ids[-1]] + pool.shelf_prefill_session.assert_called_once_with(('p', 'd'), owners[0]) + pool.unshelf_prefill_session.assert_called_once_with(('p', 'd'), owners[0]) + assert not engine.scheduler.sessions + assert not engine.scheduler.seq_manager._seq_map + + asyncio.run(run()) + + +@pytest.mark.parametrize('chat', [False, True]) +@pytest.mark.parametrize('owner', [None, -1, '42', True]) +def test_proxy_rejects_invalid_cache_owner_without_public_id_fallback(monkeypatch, chat, owner): + proxy = importlib.import_module('lmdeploy.serve.proxy.proxy') + pool = PDConnectionPool() + pool.is_connected = lambda *args: True + calls = [] + + async def check_model(model): + return None + + async def generate(payload, url, endpoint): + calls.append(url) + result = dict(id='42', cache_block_ids=[0], remote_token_ids=[1]) + if owner is not None: + result['cache_session_id'] = owner + return json.dumps(result) + + manager = SimpleNamespace( + serving_strategy=ServingStrategy.DistServe, check_request_model=check_model, + dummy_prefill=False, pd_connection_pool=pool, generate=generate, + get_node_url=lambda model, role: 'p' if role == EngineRole.Prefill else 'd', + pre_call=lambda url: 0, post_call=lambda *args: None) + monkeypatch.setattr(proxy, 'node_manager', manager) + + async def run(): + handler = proxy.chat_completions_v1 if chat else proxy.completions_v1 + response = await handler(_request(chat)) + assert response.status_code == 502 + assert 'cache_session_id' in response.body.decode() + assert calls == ['p'] + assert not pool.migration_session_shelf + + asyncio.run(run())