diff --git a/lmdeploy/serve/openai/api_client.py b/lmdeploy/serve/openai/api_client.py index 753b7260e9..445705e350 100644 --- a/lmdeploy/serve/openai/api_client.py +++ b/lmdeploy/serve/openai/api_client.py @@ -125,7 +125,7 @@ def chat_completions_v1( probable tokens with probabilities that add up to top_p or higher are kept for generation. n (int): How many chat completion choices to generate for each - input message. Only support one here. + input message. Accepts values from 1 to 128. stream: whether to stream the results or not. Default to false. max_completion_tokens (int | None): output token nums. Default to None. max_tokens (int | None): output token nums. Default to None. diff --git a/lmdeploy/serve/openai/chat_completions/fanout.py b/lmdeploy/serve/openai/chat_completions/fanout.py new file mode 100644 index 0000000000..b9ae570d3d --- /dev/null +++ b/lmdeploy/serve/openai/chat_completions/fanout.py @@ -0,0 +1,370 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Multi-choice collation for the chat completions endpoint.""" +from __future__ import annotations + +import asyncio +import json +import time +from collections.abc import AsyncGenerator, Awaitable, Callable +from copy import deepcopy +from dataclasses import dataclass + +import shortuuid +from fastapi import Request +from fastapi.responses import Response, StreamingResponse + +from lmdeploy.serve.openai.protocol import ChatCompletionRequest, UsageInfo + +from .streaming_response import ManagedStreamingResponse + +ChatEndpoint = Callable[[ChatCompletionRequest, Request], + Awaitable[dict | Response]] + + +@dataclass +class _FanoutResponseError(Exception): + """Carry an HTTP error response out of concurrent choice invocation.""" + + response: Response + + +class _FanoutRequest: + """Give each recursive endpoint call an isolated JSON payload.""" + + def __init__(self, request: Request, payload: dict): + self._request = request + self._payload = payload + + def __getattr__(self, name): + """Delegate request attributes not overridden by this wrapper.""" + return getattr(self._request, name) + + async def json(self) -> dict: + """Return an isolated copy of this choice's raw JSON payload.""" + return deepcopy(self._payload) + + async def is_disconnected(self) -> bool: + """Report the connection state of the original client request.""" + return await self._request.is_disconnected() + + +def _choice_request(request: ChatCompletionRequest, + index: int) -> ChatCompletionRequest: + """Create an independent single-choice request for one fan-out index.""" + choice_request = request.model_copy(deep=True) + choice_request.n = 1 + choice_request.session_id = -1 + if request.seed is not None: + choice_request.seed = (request.seed + index) % (1 << 64) + return choice_request + + +async def _cancel_tasks(tasks: list[asyncio.Task]) -> list: + """Cancel unfinished tasks and collect every terminal result.""" + for task in tasks: + if not task.done(): + task.cancel() + return await asyncio.gather(*tasks, return_exceptions=True) + + +async def _close_streaming_response(response: StreamingResponse) -> None: + """Close a child response through its explicit or iterator lifecycle.""" + close_response = getattr(response, 'close', None) + if close_response is not None: + await close_response() + return + close_iterator = getattr(response.body_iterator, 'aclose', None) + if close_iterator is not None: + await close_iterator() + + +async def _close_responses(responses) -> None: + """Close all streaming responses in an invocation result collection.""" + await asyncio.gather(*( + _close_streaming_response(response) + for response in responses + if isinstance(response, StreamingResponse) + ), return_exceptions=True) + + +async def _cleanup_invocations(tasks: list[asyncio.Task]) -> None: + """Cancel choice invocations and close responses they already created.""" + results = await _cancel_tasks(tasks) + await _close_responses(results) + + +def _consume_cleanup_result(task: asyncio.Task) -> None: + """Retrieve a detached cleanup task's result to suppress task warnings.""" + try: + task.result() + except BaseException: # cleanup is best-effort after caller cancellation + pass + + +async def _shield_cleanup(awaitable, name: str) -> None: + """Let cleanup continue if cancellation interrupts its caller.""" + cleanup_task = asyncio.create_task(awaitable, name=name) + try: + await asyncio.shield(cleanup_task) + except (asyncio.CancelledError, GeneratorExit): + cleanup_task.add_done_callback(_consume_cleanup_result) + raise + + +async def _invoke_choices( + endpoint: ChatEndpoint, + request: ChatCompletionRequest, + raw_request: Request, + payload: dict, +) -> list[dict | StreamingResponse] | Response: + """Invoke the single-choice endpoint concurrently for every choice.""" + + async def invoke(index: int): + """Invoke and validate one indexed single-choice response.""" + choice_request = _choice_request(request, index) + choice_payload = deepcopy(payload) + choice_payload.update(n=1, session_id=-1, seed=choice_request.seed) + response = await endpoint( + choice_request, + _FanoutRequest(raw_request, choice_payload), + ) + if isinstance(response, Response) and not isinstance( + response, StreamingResponse): + raise _FanoutResponseError(response) + return response + + tasks = [asyncio.create_task(invoke(index)) for index in range(request.n)] + try: + return await asyncio.gather(*tasks) + except _FanoutResponseError as error: + await _shield_cleanup( + _cleanup_invocations(tasks), 'fanout_invocation_cleanup') + return error.response + except BaseException: + await _shield_cleanup( + _cleanup_invocations(tasks), 'fanout_invocation_cleanup') + raise + + +def _cached_tokens(usage: dict) -> int: + """Read cached prompt tokens from an OpenAI-compatible usage object.""" + details = usage.get('prompt_tokens_details') or {} + return details.get('cached_tokens', 0) + + +def _aggregate_usage(usages: list[dict]) -> UsageInfo: + """Count shared prompt usage once and sum per-choice completion usage.""" + first_usage = usages[0] + completion_details = [ + usage.get('completion_tokens_details') for usage in usages + ] + reasoning_tokens = None + if all(details is not None for details in completion_details): + reasoning_tokens = sum( + details['reasoning_tokens'] for details in completion_details) + return UsageInfo.build( + prompt_tokens=first_usage.get('prompt_tokens', 0), + completion_tokens=sum( + usage.get('completion_tokens') or 0 for usage in usages), + cached_tokens=_cached_tokens(first_usage), + reasoning_tokens=reasoning_tokens, + ) + + +def _collate_responses( + responses: list[dict], + request_id: str, + created_time: int, +) -> dict: + """Combine single-choice JSON responses into one multi-choice response.""" + response = deepcopy(responses[0]) + response['id'] = request_id + response['created'] = created_time + response['choices'] = [] + usages = [] + for index, choice_response in enumerate(responses): + choices = choice_response.get('choices') or [] + if len(choices) != 1: + raise RuntimeError( + f'Expected one choice from fan-out request, got {len(choices)}' + ) + choice = deepcopy(choices[0]) + choice['index'] = index + response['choices'].append(choice) + usages.append(choice_response.get('usage') or {}) + response['usage'] = _aggregate_usage(usages).model_dump() + return response + + +async def _stream_choice( + index: int, + response: StreamingResponse, + queue: asyncio.Queue, + request_id: str, + created_time: int, + model_name: str, + stopping: asyncio.Event, +) -> None: + """Parse one child SSE stream and forward normalized events to a queue.""" + buffer = '' + try: + async for chunk in response.body_iterator: + buffer += chunk.decode() if isinstance(chunk, bytes) else chunk + while '\n\n' in buffer: + event, buffer = buffer.split('\n\n', 1) + for line in event.splitlines(): + if not line.startswith('data: '): + continue + data = line.removeprefix('data: ') + if data == '[DONE]': + continue + payload = json.loads(data) + if payload.get( + 'usage' + ) is not None and not payload.get('choices'): + await queue.put(('usage', index, payload['usage'])) + continue + choices = payload.get('choices') or [] + if len(choices) != 1: + raise RuntimeError( + 'Expected one streaming choice from fan-out request, ' + f'got {len(choices)}') + choices[0]['index'] = index + payload.update( + id=request_id, + created=created_time, + model=model_name, + ) + await queue.put(('data', payload)) + await queue.put(('done', index)) + except asyncio.CancelledError: + if not stopping.is_set(): + await queue.put(( + 'error', + RuntimeError(f'Fan-out choice {index} was cancelled.'), + )) + raise + except Exception as error: # noqa: BLE001 + await queue.put(('error', error)) + raise + finally: + await _close_streaming_response(response) + + +def _batch_stream_payloads(payloads: list[dict]) -> list[dict]: + """Combine each choice's Nth ready delta into the Nth output batch.""" + batches: list[dict] = [] + next_batch_by_index: dict[int, int] = {} + + for payload in payloads: + choice = payload['choices'][0] + index = choice['index'] + target = next_batch_by_index.get(index, 0) + if target == len(batches): + batches.append(payload) + else: + batches[target]['choices'].append(choice) + next_batch_by_index[index] = target + 1 + + for payload in batches: + payload['choices'].sort(key=lambda choice: choice['index']) + return batches + + +async def _collate_streams( + responses: list[StreamingResponse], + request: ChatCompletionRequest, + request_id: str, + created_time: int, +) -> AsyncGenerator[str, None]: + """Interleave child streams into one multi-choice SSE response.""" + queue: asyncio.Queue = asyncio.Queue(maxsize=max(1, len(responses) * 2)) + stopping = asyncio.Event() + # produce the streaming responses for each fan-out request + tasks = [ + asyncio.create_task( + _stream_choice(index, response, queue, request_id, created_time, + request.model, stopping)) + for index, response in enumerate(responses) + ] + usages: dict[int, dict] = {} + completed = 0 + include_usage = bool(request.stream_options + and request.stream_options.include_usage) + # consume the streaming responses of each fan-out request, and yield to the client + try: + while completed < len(tasks): + items = [await queue.get()] + while True: + try: + items.append(queue.get_nowait()) + except asyncio.QueueEmpty: + break + + payloads = [] + stream_error = None + for item in items: + if item[0] == 'data': + payloads.append(item[1]) + elif item[0] == 'usage': + # item[1]: index, item[2]: usage payload + usages[item[1]] = item[2] + elif item[0] == 'done': + completed += 1 + else: + stream_error = item[1] + break + + for payload in _batch_stream_payloads(payloads): + yield f'data: {json.dumps(payload)}\n\n' + if stream_error is not None: + raise stream_error + if include_usage and len(usages) == len(tasks): + ordered_usages = [usages[index] for index in range(len(tasks))] + usage_response = { + 'id': request_id, + 'object': 'chat.completion.chunk', + 'created': created_time, + 'model': request.model, + 'choices': [], + 'usage': _aggregate_usage(ordered_usages).model_dump(), + } + yield f'data: {json.dumps(usage_response)}\n\n' + yield 'data: [DONE]\n\n' + finally: + stopping.set() + await _shield_cleanup(_cancel_tasks(tasks), + 'fanout_stream_cleanup') + + +async def fanout_chat_completions( + endpoint: ChatEndpoint, + request: ChatCompletionRequest, + raw_request: Request, + payload: dict, +) -> dict | Response: + """Run the established single-choice endpoint once per requested choice.""" + request_id = f'chatcmpl-{shortuuid.random()}' + created_time = int(time.time()) + responses = await _invoke_choices(endpoint, request, raw_request, payload) + if isinstance(responses, Response): + return responses + if request.stream: + if not all( + isinstance(response, StreamingResponse) + for response in responses): + await _close_responses(responses) + raise RuntimeError( + 'Expected streaming responses from fan-out requests') + stream = _collate_streams(responses, request, request_id, created_time) + return ManagedStreamingResponse( + stream, + cleanup_callbacks=[ + lambda response=response: _close_streaming_response(response) + for response in responses + ], + media_type='text/event-stream') + if not all(isinstance(response, dict) for response in responses): + await _close_responses(responses) + raise RuntimeError('Expected JSON objects from fan-out requests') + return _collate_responses(responses, request_id, created_time) diff --git a/lmdeploy/serve/openai/chat_completions/serving.py b/lmdeploy/serve/openai/chat_completions/serving.py index b8190fc9c1..4b8b0edc7a 100644 --- a/lmdeploy/serve/openai/chat_completions/serving.py +++ b/lmdeploy/serve/openai/chat_completions/serving.py @@ -9,7 +9,6 @@ import shortuuid from fastapi import APIRouter, Depends, Request -from fastapi.responses import StreamingResponse from lmdeploy.pytorch.disagg.conn.protocol import MigrationRequest from lmdeploy.serve.core.exceptions import RequestError @@ -35,8 +34,10 @@ from lmdeploy.serve.utils.server_utils import validate_json_request from lmdeploy.utils import get_logger +from .fanout import fanout_chat_completions from .logits_processors import logit_bias_logits_processor from .logprobs import _create_chat_completion_logprobs, _create_output_token_logprobs +from .streaming_response import ManagedStreamingResponse from .validation import check_request logger = get_logger('lmdeploy') @@ -63,7 +64,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, probable tokens with probabilities that add up to top_p or higher are kept for generation. - **n** (int): How many chat completion choices to generate for each input - message. **Only support one here**. + message. Accepts values from 1 to 128. - **stream**: whether to stream the results or not. Default to false. - **stream_options**: Options for streaming response. Only set this when you set stream: true. @@ -131,10 +132,23 @@ async def chat_completions_v1(request: ChatCompletionRequest, - **presence_penalty** (replaced with repetition_penalty) - **frequency_penalty** (replaced with repetition_penalty) """ - error_check_ret = validate_request(request, server_context, - check_request) + json_request = await raw_request.json() + error_check_ret = validate_request( + request, + server_context, + check_request, + json_request=json_request, + ) if error_check_ret is not None: return error_check_ret + if request.n is not None and request.n > 1: + return await fanout_chat_completions( + chat_completions_v1, + request, + raw_request, + json_request, + ) + # Resolve input: messages has priority over input_ids/image_data messages_empty = (request.messages is None or request.messages == '' or (isinstance(request.messages, list) @@ -165,7 +179,6 @@ async def chat_completions_v1(request: ChatCompletionRequest, # input_ids only — engine requires messages=None request.messages = None - json_request = await raw_request.json() migration_request = json_request.pop('migration_request', None) with_cache = json_request.pop('with_cache', False) preserve_cache = json_request.pop('preserve_cache', False) @@ -247,9 +260,10 @@ async def chat_completions_v1(request: ChatCompletionRequest, mm_processor_kwargs=request.mm_processor_kwargs) except RequestError as error: return create_request_error_response(error) + result_generator = server_context.async_engine.generate( - preprocessed, - stream_response=True) # always use stream to enable batching + preprocessed, + stream_response=True) # always use stream to enable batching include_usage = bool(request.stream_options and request.stream_options.include_usage) @@ -383,11 +397,12 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]: # Streaming response if request.stream: - stream_generator = with_request_cleanup( - completion_stream_generator(), [result_generator], [session], - server_context.session_manager) - return StreamingResponse(stream_generator, - media_type='text/event-stream') + return ManagedStreamingResponse( + completion_stream_generator(), + result_generators=[result_generator], + sessions=[session], + session_mgr=server_context.session_manager, + media_type='text/event-stream') # Non-streaming response final_logprobs = [] diff --git a/lmdeploy/serve/openai/chat_completions/streaming_response.py b/lmdeploy/serve/openai/chat_completions/streaming_response.py new file mode 100644 index 0000000000..6919c12055 --- /dev/null +++ b/lmdeploy/serve/openai/chat_completions/streaming_response.py @@ -0,0 +1,117 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Streaming-response resource ownership for chat completions. + +Chat-completion fan-out creates multiple single-choice responses. Each child +allocates an engine session and result generator before Starlette starts its +body iterator. If another child fails, or the client disconnects before the +combined stream starts, cleanup placed only in an iterator ``finally`` block +is never activated and those resources can leak. + +``ManagedStreamingResponse`` makes that ownership explicit on the response +itself. Normal iteration still performs cleanup, while ``close()`` also lets +fan-out and the ASGI response lifecycle release resources whose iterators were +never started. +""" +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable, Iterable + +from fastapi.responses import StreamingResponse + +from lmdeploy.serve.utils.request_cleanup import cleanup_result_generators +from lmdeploy.utils import get_logger + +logger = get_logger('lmdeploy') + + +class ManagedStreamingResponse(StreamingResponse): + """A chat-completion response with explicit, cancellation-safe cleanup. + + Keeping cleanup on the response makes resource ownership independent of whether Starlette entered its body iterator. + """ + + def __init__( + self, + content, + *, + result_generators: Iterable = (), + sessions: Iterable = (), + session_mgr=None, + cleanup_callbacks: Iterable[Callable[[], Awaitable[None]]] = (), + **kwargs, + ): + self._result_generators = tuple(result_generators) + self._sessions = tuple(sessions) + self._session_mgr = session_mgr + self._cleanup_callbacks = tuple(cleanup_callbacks) + self._resource_cleanup_task: asyncio.Task | None = None + self._close_task: asyncio.Task | None = None + if self._result_generators or self._sessions: + content = self._with_resource_cleanup(content) + super().__init__(content, **kwargs) + + async def _with_resource_cleanup(self, content): + """Clean owned engine resources after normal body iteration.""" + try: + async for item in content: + yield item + finally: + await self._cleanup_resources() + + async def _cleanup_resources(self) -> None: + """Close result generators and sessions once, despite cancellation.""" + if not self._result_generators and not self._sessions: + return + if self._resource_cleanup_task is None: + self._resource_cleanup_task = asyncio.create_task( + cleanup_result_generators( + self._result_generators, + self._sessions, + self._session_mgr, + ), + name='streaming_response_resource_cleanup') + await asyncio.shield(self._resource_cleanup_task) + + async def _close(self) -> None: + """Close the response body, owned resources, and child callbacks.""" + body_iterator = self.body_iterator + close_iterator = getattr(body_iterator, 'aclose', None) + if close_iterator is not None: + try: + await close_iterator() + except (asyncio.CancelledError, GeneratorExit): + pass + except Exception: + logger.exception('Close response body iterator failed.') + + await self._cleanup_resources() + for callback in self._cleanup_callbacks: + try: + await callback() + except (asyncio.CancelledError, GeneratorExit): + pass + except Exception: + logger.exception('Streaming response cleanup callback failed.') + + async def close(self) -> None: + """Close the body and its resources exactly once.""" + if (not self._cleanup_callbacks + and self._resource_cleanup_task is not None + and self._resource_cleanup_task.done()): + return + if self._close_task is None: + self._close_task = asyncio.create_task( + self._close(), name='streaming_response_cleanup') + try: + await asyncio.shield(self._close_task) + except (asyncio.CancelledError, GeneratorExit): + raise + except Exception: + logger.exception('Streaming response cleanup failed.') + + async def __call__(self, scope, receive, send) -> None: + try: + await super().__call__(scope, receive, send) + finally: + await self.close() diff --git a/lmdeploy/serve/openai/chat_completions/validation.py b/lmdeploy/serve/openai/chat_completions/validation.py index 13d696b7c9..18ef3571f0 100644 --- a/lmdeploy/serve/openai/chat_completions/validation.py +++ b/lmdeploy/serve/openai/chat_completions/validation.py @@ -4,8 +4,15 @@ from lmdeploy.serve.openai.protocol import ChatCompletionRequest +# Upper bound for `n` (number of choices). Each choice is a separate +# engine.generate() call on the fan-out path, so cap to protect resources. +_MAX_FANOUT_N = 128 -def check_request(request: ChatCompletionRequest, server_context) -> str: + +def check_request(request: ChatCompletionRequest, + server_context, + json_request: dict | None = None) -> str: + """Validate chat-completion options and fan-out compatibility.""" engine_config = server_context.engine_config session_manager = server_context.session_manager try: @@ -35,15 +42,28 @@ def check_request(request: ChatCompletionRequest, server_context) -> str: return f'The session_id {request.session_id!r} is occupied.' # check sampling settings - if request.n <= 0: + if request.n is not None and request.n <= 0: return f'The n {request.n!r} must be a positive int.' + # n > 1 is implemented as server-side fan-out (N independent engine + # generate() calls). Cap it to prevent unbounded resource use. + if request.n is not None and request.n > _MAX_FANOUT_N: + return (f'The n {request.n!r} exceeds the maximum supported ' + f'choices ({_MAX_FANOUT_N}).') + if request.n is not None and request.n > 1 and request.session_id not in ( + None, -1): + return 'n > 1 cannot be used with an explicit session_id.' + if request.n is not None and request.n > 1 and json_request is not None: + if any( + json_request.get(key) + for key in ('migration_request', 'with_cache', + 'preserve_cache')): + return 'n > 1 is not supported with cache migration.' if request.top_p is not None and not (0 < request.top_p <= 1): return f'The top_p {request.top_p!r} must be in (0, 1].' if request.top_k is not None and request.top_k < 0: return f'The top_k {request.top_k!r} cannot be a negative integer.' if request.temperature is not None and not (0 <= request.temperature <= 2): return f'The temperature {request.temperature!r} must be in [0, 2]' - # Validate input_ids and image_data constraints. # messages has higher priority. input_ids and image_data are only used when # messages is empty (None, '', or []). image_data requires input_ids. diff --git a/lmdeploy/serve/openai/endpoints/common.py b/lmdeploy/serve/openai/endpoints/common.py index 8b222ff4f0..41b080a015 100644 --- a/lmdeploy/serve/openai/endpoints/common.py +++ b/lmdeploy/serve/openai/endpoints/common.py @@ -23,7 +23,8 @@ def build_serving_generation_config(request, server_context, ) -def validate_request(request, server_context, request_validator): +def validate_request(request, server_context, request_validator, + **validator_kwargs): """Validate the selected model and endpoint-specific request contract.""" if hasattr( request, @@ -32,7 +33,8 @@ def validate_request(request, server_context, request_validator): HTTPStatus.NOT_FOUND, f'The model {request.model!r} does not exist.') - error_message = request_validator(request, server_context) + error_message = request_validator(request, server_context, + **validator_kwargs) if error_message: return create_error_response(HTTPStatus.BAD_REQUEST, error_message) return None diff --git a/lmdeploy/serve/proxy/proxy.py b/lmdeploy/serve/proxy/proxy.py index 7afd875c2b..015925ad61 100644 --- a/lmdeploy/serve/proxy/proxy.py +++ b/lmdeploy/serve/proxy/proxy.py @@ -587,7 +587,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque probable tokens with probabilities that add up to top_p or higher are kept for generation. - **n** (int): How many chat completion choices to generate for each input - message. **Only support one here**. + message. Accepts values from 1 to 128, except in DistServe mode. - **stream**: whether to stream the results or not. Default to false. - **max_completion_tokens** (int | None): output token nums. Default to None. - **max_tokens** (int | None): output token nums. Default to None. @@ -650,6 +650,11 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque check_response = await node_manager.check_request_model(request.model) if check_response is not None: return check_response + if (node_manager.serving_strategy == ServingStrategy.DistServe + and request.n is not None and request.n > 1): + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'n > 1 is not supported with the DistServe serving strategy.') if node_manager.serving_strategy == ServingStrategy.Hybrid: node_url = node_manager.get_node_url(request.model) diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/conftest.py b/tests/test_lmdeploy/serve/openai/chat_completions/conftest.py new file mode 100644 index 0000000000..64327b28d3 --- /dev/null +++ b/tests/test_lmdeploy/serve/openai/chat_completions/conftest.py @@ -0,0 +1,192 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Shared fakes for ``/v1/chat/completions`` handler tests.""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from lmdeploy.serve.openai.chat_completions import register +from lmdeploy.serve.openai.protocol import DeltaMessage + + +class FakeTokenizer: + model = SimpleNamespace(model='fake-tokenizer') + + +class FakeAsyncEngine: + """Engine fake whose ``generate`` returns distinct outputs per call. + + Each call yields a ``GenOut``-like stream whose text encodes the call + index, so fan-out tests can assert the N choices are distinct. + """ + + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self): + self.session_mgr = FakeSessionManager() + self.tokenizer = SimpleNamespace(model=FakeTokenizer()) + self.call_count = 0 + self.gen_configs = [] + + async def preprocess(self, prompt, session, **kwargs): + """Return the minimal preprocessed input consumed by the fake.""" + self.gen_configs.append(kwargs.get('gen_config')) + return SimpleNamespace(prompt=prompt, session=session) + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + call_index = self.call_count + + async def _generator(): + yield SimpleNamespace( + response=f'choice-{call_index}', + token_ids=[call_index], + input_token_len=4, + generate_token_len=call_index, + finish_reason='stop', + logprobs=None, + cached_tokens=0, + routed_experts=None, + cache_block_ids=None, + ) + + return _generator() + + +class PassthroughResponseParser: + """Stateful passthrough parser mirroring the real ResponseParser API.""" + + tool_parser_cls = None + + def __init__(self, request): + self.request = request + self.tool_parser = None + self._chunks = [] + self.reasoning_tokens = 0 + + def stream_chunk(self, delta_text, delta_token_ids, **kwargs): + if not delta_text: + return [] + return [(DeltaMessage(content=delta_text), False)] + + def parse_complete(self, text, token_ids=None, **kwargs): + return text, None, None + + def validate_complete(self, raw_text=None): + return True + + +class FakeSessionManager: + """Mimics the real SessionManager's id/mapping semantics closely enough + to surface fan-out session bugs: explicit user_session_ids are mapped + one-to-one and a duplicate raises (like map_user_session_id), while + None/-1 auto-generates a fresh internal id.""" + + def __init__(self): + self.removed = [] + self.sessions = {} + self.user_session_id_map = {} + self._next_id = 0 + + def map_user_session_id(self, user_session_id): + if user_session_id in self.user_session_id_map: + raise ValueError( + f'User session id {user_session_id} already exists') + session_id = self._next_id + self._next_id += 1 + self.user_session_id_map[user_session_id] = session_id + return session_id + + def get(self, session_id=None, create_if_not_exists=True, **kwargs): + if not create_if_not_exists: + return self.sessions.get(session_id, None) + if session_id is None: + session_id = self._next_id + self._next_id += 1 + if session_id in self.sessions: + return self.sessions[session_id] + session = FakeSession(session_id) + self.sessions[session_id] = session + return session + + def has(self, session_id): + return session_id in self.sessions + + def remove(self, session): + if session is None: + return + session_id = (session if isinstance(session, int) + else session.session_id) + self.sessions.pop(session_id, None) + # also drop any user mapping pointing at this session_id + for uid, sid in list(self.user_session_id_map.items()): + if sid == session_id: + self.user_session_id_map.pop(uid, None) + self.removed.append(session) + + +class FakeSession: + + def __init__(self, session_id): + self.session_id = session_id + self.epoch = 0 + self.aborted = False + + async def async_abort(self): + self.aborted = True + + +class FakeServerContext: + response_parser_cls = PassthroughResponseParser + + def __init__(self): + self.async_engine = FakeAsyncEngine() + self.default_gen_config = {} + + @property + def engine_config(self): + return self.async_engine.backend_config + + @property + def session_manager(self): + return self.async_engine.session_mgr + + def create_session(self, user_session_id): + # Mirror ServerContext.create_session: None/-1 auto-generates; an + # explicit id maps one-to-one and collides on a second use. + if user_session_id is None or user_session_id == -1: + session = self.session_manager.get() + else: + session_id = self.session_manager.map_user_session_id( + user_session_id) + session = self.session_manager.get(session_id) + session.epoch = 0 + return session + + +class FakeRawRequest: + + def __init__(self, payload=None): + self._payload = payload or {} + + async def json(self): + return self._payload + + async def is_disconnected(self): + return False + + +@pytest.fixture +def chat_endpoint(): + context = FakeServerContext() + from fastapi import APIRouter + r = APIRouter() + register(r, context) + return r.routes[0].endpoint, context + + +@pytest.fixture +def fake_raw_request(): + return FakeRawRequest() diff --git a/tests/test_lmdeploy/serve/openai/chat_completions/test_n_completions.py b/tests/test_lmdeploy/serve/openai/chat_completions/test_n_completions.py new file mode 100644 index 0000000000..ae22efaa1a --- /dev/null +++ b/tests/test_lmdeploy/serve/openai/chat_completions/test_n_completions.py @@ -0,0 +1,539 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Regression tests for multiple chat completion choices.""" +from __future__ import annotations + +import asyncio +import json +from types import SimpleNamespace + +import pytest +from fastapi.responses import JSONResponse + +from lmdeploy.serve.openai.chat_completions.fanout import _batch_stream_payloads +from lmdeploy.serve.openai.protocol import ChatCompletionRequest + + +class _PreprocessingEngine: + """Provide the preprocessing stage expected by the serving endpoint.""" + + async def preprocess(self, prompt, session, **kwargs): + self.gen_configs.append(kwargs.get('gen_config')) + return SimpleNamespace(prompt=prompt, session=session) + + +def _request(**kwargs): + return ChatCompletionRequest( + model='fake-model', + messages=[{ + 'role': 'user', + 'content': 'hi' + }], + **kwargs, + ) + + +def _sse_payloads(text): + payloads = [] + for line in text.splitlines(): + if line.startswith('data: '): + data = line.removeprefix('data: ') + if data != '[DONE]': + payloads.append(json.loads(data)) + return payloads + + +async def _collect_stream(response): + chunks = [] + async for chunk in response.body_iterator: + chunks.append(chunk.decode() if isinstance(chunk, bytes) else chunk) + return ''.join(chunks) + + +def _stream_payload(index, content, **metadata): + return { + 'id': 'chatcmpl-test', + 'object': 'chat.completion.chunk', + 'created': 1, + 'model': 'fake-model', + 'choices': [{ + 'index': index, + 'delta': { + 'content': content + } + }], + **metadata, + } + + +def test_ready_stream_chunks_are_batched_by_choice(): + payloads = [ + _stream_payload(0, '0-a'), + _stream_payload(0, '0-b'), + _stream_payload(1, '1-a'), + ] + + batches = _batch_stream_payloads(payloads) + + assert [[choice['index'] for choice in batch['choices']] + for batch in batches] == [[0, 1], [0]] + assert [choice['delta']['content'] + for batch in batches + for choice in batch['choices'] if choice['index'] == 0 + ] == ['0-a', '0-b'] + + +def test_handler_n3_nonstream_collates_single_choice_path( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + response = asyncio.run(endpoint(_request(n=3, seed=42), fake_raw_request)) + + assert response['object'] == 'chat.completion' + assert [choice['index'] for choice in response['choices']] == [0, 1, 2] + assert {choice['message']['content'] + for choice in response['choices'] + } == {'choice-1', 'choice-2', 'choice-3'} + assert response['usage']['prompt_tokens'] == 4 + assert response['usage']['completion_tokens'] == 6 + assert [config.random_seed + for config in context.async_engine.gen_configs] == [42, 43, 44] + assert context.async_engine.call_count == 3 + assert len(context.session_manager.removed) == 3 + assert context.session_manager.sessions == {} + + +def test_handler_n3_stream_interleaves_indices_and_aggregates_usage( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + request = _request( + n=3, + stream=True, + stream_options={'include_usage': True}, + ) + response = asyncio.run(endpoint(request, fake_raw_request)) + text = asyncio.run(_collect_stream(response)) + payloads = _sse_payloads(text) + + indices = { + choice['index'] + for payload in payloads + for choice in payload.get('choices', []) + } + assert indices == {0, 1, 2} + usage_chunks = [ + payload for payload in payloads if payload.get('usage') is not None + ] + assert len(usage_chunks) == 1 + assert usage_chunks[0]['choices'] == [] + assert usage_chunks[0]['usage']['prompt_tokens'] == 4 + assert usage_chunks[0]['usage']['completion_tokens'] == 6 + assert text.rstrip().endswith('data: [DONE]') + assert len(context.session_manager.removed) == 3 + assert context.session_manager.sessions == {} + + +@pytest.mark.parametrize('n', [None, 1]) +def test_handler_single_choice_keeps_fast_path(n, chat_endpoint, + fake_raw_request): + endpoint, context = chat_endpoint + response = asyncio.run(endpoint(_request(n=n), fake_raw_request)) + assert len(response['choices']) == 1 + assert context.async_engine.call_count == 1 + + +def test_handler_unseeded_choices_leave_seed_resolution_to_engine( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + asyncio.run(endpoint(_request(n=3), fake_raw_request)) + assert [config.random_seed for config in context.async_engine.gen_configs + ] == [None, None, None] + + +def test_handler_negative_seed_is_mapped_to_engine_seed_domain( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + asyncio.run(endpoint(_request(n=2, seed=-1), fake_raw_request)) + assert [config.random_seed + for config in context.async_engine.gen_configs] == [(1 << 64) - 1, + 0] + + +@pytest.mark.parametrize('n, expected', [ + (0, 'positive int'), + (129, 'maximum supported'), +]) +def test_validation_rejects_invalid_n(n, expected, chat_endpoint, + fake_raw_request): + endpoint, context = chat_endpoint + response = asyncio.run(endpoint(_request(n=n), fake_raw_request)) + assert isinstance(response, JSONResponse) + assert response.status_code == 400 + assert expected in response.body.decode() + assert context.async_engine.call_count == 0 + + +def test_handler_rejects_explicit_session_id_for_multiple_choices( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + response = asyncio.run( + endpoint(_request(n=2, session_id=777), fake_raw_request)) + assert isinstance(response, JSONResponse) + assert response.status_code == 400 + assert 'explicit session_id' in response.body.decode() + assert context.async_engine.call_count == 0 + assert context.session_manager.sessions == {} + + +def test_handler_rejects_cache_migration_for_multiple_choices( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + fake_raw_request._payload = {'with_cache': True} + response = asyncio.run(endpoint(_request(n=2), fake_raw_request)) + assert isinstance(response, JSONResponse) + assert response.status_code == 400 + assert 'cache migration' in response.body.decode() + assert context.async_engine.call_count == 0 + + +def test_distserve_proxy_rejects_multiple_choices(monkeypatch): + from lmdeploy.pytorch.disagg.config import ServingStrategy + from lmdeploy.serve.proxy import proxy + + async def model_exists(model): + return None + + monkeypatch.setattr(proxy.node_manager, 'check_request_model', + model_exists) + monkeypatch.setattr(proxy.node_manager, 'serving_strategy', + ServingStrategy.DistServe) + response = asyncio.run(proxy.chat_completions_v1(_request(n=2))) + assert isinstance(response, JSONResponse) + assert response.status_code == 400 + assert 'DistServe' in response.body.decode() + + +def test_prompt_cache_usage_is_counted_once(chat_endpoint, fake_raw_request): + + class CachedEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def generate(): + yield SimpleNamespace( + response=f'choice-{index}', + token_ids=[index], + input_token_len=4, + generate_token_len=1, + finish_reason='stop', + logprobs=None, + cached_tokens=index, + routed_experts=None, + cache_block_ids=None, + ) + + return generate() + + endpoint, context = chat_endpoint + context.async_engine = CachedEngine(context.async_engine) + response = asyncio.run(endpoint(_request(n=2), fake_raw_request)) + assert response['usage']['prompt_tokens'] == 4 + assert response['usage']['completion_tokens'] == 2 + assert response['usage']['prompt_tokens_details']['cached_tokens'] == 1 + + +def test_streaming_multiple_choices_preserves_each_inner_parser( + chat_endpoint, fake_raw_request): + + class MultiChunkEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def generate(): + for suffix in ('a', 'b', 'c'): + yield SimpleNamespace( + response=f'{index}-{suffix}', + token_ids=[index], + input_token_len=3, + generate_token_len=1, + finish_reason=None, + logprobs=None, + cached_tokens=0, + routed_experts=None, + cache_block_ids=None, + ) + yield SimpleNamespace( + response='', + token_ids=[], + input_token_len=3, + generate_token_len=3, + finish_reason='stop', + logprobs=None, + cached_tokens=0, + routed_experts=None, + cache_block_ids=None, + ) + + return generate() + + endpoint, context = chat_endpoint + context.async_engine = MultiChunkEngine(context.async_engine) + response = asyncio.run( + endpoint(_request(n=2, stream=True), fake_raw_request)) + payloads = _sse_payloads(asyncio.run(_collect_stream(response))) + content = {0: '', 1: ''} + for payload in payloads: + for choice in payload.get('choices', []): + content[choice['index']] += choice['delta'].get('content') or '' + assert content == {0: '1-a1-b1-c', 1: '2-a2-b2-c'} + + +def test_fanout_error_cancels_siblings_and_cleans_sessions( + chat_endpoint, fake_raw_request): + + class FailingEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + self.sibling_started = asyncio.Event() + self.sibling_closed = False + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def fail(): + await self.sibling_started.wait() + raise RuntimeError('choice failed') + yield # noqa: unreachable + + async def wait_forever(): + self.sibling_started.set() + try: + await asyncio.Event().wait() + yield # noqa: unreachable + finally: + self.sibling_closed = True + + return fail() if index == 1 else wait_forever() + + endpoint, context = chat_endpoint + engine = FailingEngine(context.async_engine) + context.async_engine = engine + + with pytest.raises(RuntimeError, match='choice failed'): + asyncio.run(endpoint(_request(n=2), fake_raw_request)) + assert engine.sibling_closed + assert context.session_manager.sessions == {} + + +def test_early_stream_close_cleans_all_choice_sessions(chat_endpoint, + fake_raw_request): + endpoint, context = chat_endpoint + response = asyncio.run( + endpoint(_request(n=2, stream=True), fake_raw_request)) + + async def consume_one_chunk(): + iterator = response.body_iterator + await anext(iterator) + await iterator.aclose() + + asyncio.run(consume_one_chunk()) + assert context.session_manager.sessions == {} + + +def test_asgi_disconnect_before_stream_start_cleans_all_choice_sessions( + chat_endpoint, fake_raw_request): + endpoint, context = chat_endpoint + + async def disconnect_before_stream_start(): + response = await endpoint( + _request(n=2, stream=True), fake_raw_request) + + async def receive(): + return {'type': 'http.disconnect'} + + async def send(message): + if message['type'] == 'http.response.start': + await asyncio.Event().wait() + + scope = { + 'type': 'http', + 'asgi': { + 'version': '3.0', + 'spec_version': '2.3' + }, + } + await response(scope, receive, send) + + asyncio.run(asyncio.wait_for(disconnect_before_stream_start(), 1)) + assert context.session_manager.sessions == {} + + +def test_streaming_cancelled_choice_fails_without_hanging( + chat_endpoint, fake_raw_request): + + class CancelledChoiceEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + self.sibling_started = asyncio.Event() + self.sibling_closed = False + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def cancel(): + await self.sibling_started.wait() + raise asyncio.CancelledError + yield # noqa: unreachable + + async def wait_forever(): + self.sibling_started.set() + try: + await asyncio.Event().wait() + yield # noqa: unreachable + finally: + self.sibling_closed = True + + return cancel() if index == 1 else wait_forever() + + endpoint, context = chat_endpoint + engine = CancelledChoiceEngine(context.async_engine) + context.async_engine = engine + + async def collect(): + response = await endpoint( + _request(n=2, stream=True), fake_raw_request) + await asyncio.wait_for(_collect_stream(response), 1) + + with pytest.raises(RuntimeError, match='choice 0 was cancelled'): + asyncio.run(collect()) + assert engine.sibling_closed + assert context.session_manager.sessions == {} + + +def test_streaming_generation_failure_cleans_all_sessions( + chat_endpoint, fake_raw_request): + + class FailingEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + self.sibling_started = asyncio.Event() + self.sibling_closed = False + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def wait_forever(): + self.sibling_started.set() + try: + await asyncio.Event().wait() + yield # noqa: unreachable + finally: + self.sibling_closed = True + + async def fail(): + await self.sibling_started.wait() + raise RuntimeError('choice generation failed') + yield # noqa: unreachable + + return wait_forever() if index == 1 else fail() + + endpoint, context = chat_endpoint + engine = FailingEngine(context.async_engine) + context.async_engine = engine + + async def collect(): + response = await endpoint( + _request(n=2, stream=True), fake_raw_request) + await asyncio.wait_for(_collect_stream(response), 1) + + with pytest.raises(RuntimeError, match='choice generation failed'): + asyncio.run(collect()) + assert engine.sibling_closed + assert context.session_manager.sessions == {} + + +def test_stream_usage_is_omitted_when_a_choice_has_no_usage( + chat_endpoint, fake_raw_request): + + class IncompleteUsageEngine(_PreprocessingEngine): + model_name = 'fake-model' + backend_config = SimpleNamespace(adapters=[], logprobs_mode=None) + + def __init__(self, original_engine): + self.session_mgr = original_engine.session_mgr + self.tokenizer = original_engine.tokenizer + self.call_count = 0 + self.gen_configs = [] + + def generate(self, preprocessed, **kwargs): + self.call_count += 1 + index = self.call_count + + async def generate(): + yield SimpleNamespace( + response=f'choice-{index}', + token_ids=[index], + input_token_len=4, + generate_token_len=1, + finish_reason='stop' if index == 1 else None, + logprobs=None, + cached_tokens=0, + routed_experts=None, + cache_block_ids=None, + ) + + return generate() + + endpoint, context = chat_endpoint + context.async_engine = IncompleteUsageEngine(context.async_engine) + request = _request( + n=2, + stream=True, + stream_options={'include_usage': True}, + ) + response = asyncio.run(endpoint(request, fake_raw_request)) + payloads = _sse_payloads(asyncio.run(_collect_stream(response))) + + assert not [ + payload for payload in payloads if payload.get('usage') is not None + ] + assert context.session_manager.sessions == {}