Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions lmdeploy/serve/openai/chat_completions/serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'
Expand Down Expand Up @@ -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

Expand Down
8 changes: 7 additions & 1 deletion lmdeploy/serve/openai/endpoints/completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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'
Expand Down Expand Up @@ -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

Expand Down
33 changes: 26 additions & 7 deletions lmdeploy/serve/proxy/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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)
Expand All @@ -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

Expand Down Expand Up @@ -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(
Expand All @@ -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)
Expand All @@ -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}')
Expand Down
227 changes: 227 additions & 0 deletions tests/test_lmdeploy/serve/test_distserve_cache_session_id.py
Original file line number Diff line number Diff line change
@@ -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