class ServingRender(BaseServing):
def __init__(
self,
models: OpenAIServingModels | OpenAIModelRegistry,
online_renderer: "OnlineRenderer",
*,
request_logger: RequestLogger | None = None,
tool_server: "ToolServer | None" = None,
) -> None:
super().__init__(
models=models,
model_config=online_renderer.model_config,
request_logger=request_logger,
)
self.online_renderer = online_renderer
self.tool_server = tool_server
self._merge_inline_system = (
AnthropicServingMessages._detect_merge_inline_system(
online_renderer.chat_template
)
)
self._placeholder_metadata_parser: MultiModalDataParser | None = None
self._placeholder_metadata_parser_failed = False
self.default_sampling_params = (
online_renderer.model_config.get_diff_sampling_param()
)
mc = online_renderer.model_config
self.override_max_tokens = (
self.default_sampling_params.get("max_tokens")
if mc.generation_config not in ("auto", "vllm")
else getattr(mc, "override_generation_config", {}).get("max_new_tokens")
)
async def render_chat_request(
self,
request: ChatCompletionRequest,
) -> GenerateRequest | ErrorResponse:
"""Validate the model and preprocess a chat completion request.
This is the authoritative implementation used directly by the
GPU-less render server and delegated to by OpenAIServingChat.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
logger.error("Error with model %s", error_check_ret)
return error_check_ret
if request.use_beam_search:
return self.create_error_response(
"Beam search is not supported by the render endpoint"
)
result = await self.online_renderer.render_chat(request, skip_mm_cache=True)
if isinstance(result, ErrorResponse):
return result
_, engine_inputs = result
if len(engine_inputs) != 1:
return self.create_error_response(
f"Expected exactly 1 engine prompt, got {len(engine_inputs)}"
)
engine_input = engine_inputs[0]
prompt_components = extract_prompt_components(self.model_config, engine_input)
token_ids = prompt_components.token_ids
if not token_ids:
return self.create_error_response("No token_ids rendered")
token_ids = list(token_ids)
input_length = extract_prompt_len(self.model_config, engine_input)
max_tokens = get_max_tokens(
self.model_config.max_model_len,
request.max_completion_tokens
if request.max_completion_tokens is not None
else request.max_tokens,
input_length,
self.default_sampling_params,
self.override_max_tokens,
truncate_prompt_tokens=request.truncate_prompt_tokens,
)
params = request.to_sampling_params(max_tokens, self.default_sampling_params)
assistant_tokens_mask: list[int] | None = engine_input.get( # type: ignore[assignment]
"assistant_tokens_mask"
)
if assistant_tokens_mask is not None and len(assistant_tokens_mask) != len(
token_ids
):
logger.warning(
"assistant_tokens_mask length (%d) != token_ids length (%d); "
"this can happen with multimodal inputs where "
"placeholder expansion changes the token count. "
"The mask may be positionally misaligned.",
len(assistant_tokens_mask),
len(token_ids),
)
if len(assistant_tokens_mask) < len(token_ids):
assistant_tokens_mask.extend(
[0] * (len(token_ids) - len(assistant_tokens_mask))
)
else:
assistant_tokens_mask = assistant_tokens_mask[: len(token_ids)]
request_id = f"chatcmpl-{random_uuid()}"
return GenerateRequest(
request_id=request_id,
token_ids=token_ids,
assistant_tokens_mask=assistant_tokens_mask,
features=self._extract_mm_features(engine_input),
sampling_params=params,
model=request.model,
stream=bool(request.stream),
stream_options=(request.stream_options if request.stream else None),
cache_salt=request.cache_salt,
priority=request.priority,
token_offsets=engine_input.get("prompt_token_offsets"),
)
async def render_messages_request(
self,
request: AnthropicMessagesRequest,
) -> GenerateRequest | ErrorResponse:
"""Validate the model and preprocess an Anthropic Messages request.
Converts the request to the OpenAI chat format using the same
conversion as the /v1/messages server path, then delegates to
render_chat_request so the rendered tokens match the server exactly.
"""
chat_req = AnthropicServingMessages.to_chat_completion_request(
request, merge_inline_system=self._merge_inline_system
)
return await self.render_chat_request(chat_req)
async def render_completion_request(
self,
request: CompletionRequest,
) -> list[GenerateRequest] | ErrorResponse:
"""Validate the model and preprocess a completion request.
This is the authoritative implementation used directly by the
GPU-less render server and delegated to by OpenAIServingCompletion.
"""
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
result = await self.online_renderer.render_completion(
request, skip_mm_cache=True
)
if isinstance(result, ErrorResponse):
return result
generate_requests: list[GenerateRequest] = []
for engine_input in result:
prompt_components = extract_prompt_components(
self.model_config, engine_input
)
token_ids = prompt_components.token_ids
if not token_ids:
return self.create_error_response("No token_ids rendered")
token_ids = list(token_ids)
input_length = extract_prompt_len(self.model_config, engine_input)
max_tokens = get_max_tokens(
self.model_config.max_model_len,
request.max_tokens,
input_length,
self.default_sampling_params,
self.override_max_tokens,
truncate_prompt_tokens=request.truncate_prompt_tokens,
)
params = request.to_sampling_params(
max_tokens, self.default_sampling_params
)
request_id = f"cmpl-{random_uuid()}"
generate_requests.append(
GenerateRequest(
request_id=request_id,
token_ids=token_ids,
features=self._extract_mm_features(engine_input),
sampling_params=params,
model=request.model,
stream=bool(request.stream),
stream_options=(request.stream_options if request.stream else None),
cache_salt=request.cache_salt,
priority=request.priority,
token_offsets=engine_input.get("prompt_token_offsets"),
)
)
return generate_requests
async def render_responses_request(
self,
request: ResponsesRequest,
) -> GenerateRequest | ErrorResponse:
error_check_ret = await self._check_model(request)
if error_check_ret is not None:
return error_check_ret
if request.previous_response_id is not None:
return self.create_error_response(
message=(
"previous_response_id is not supported by the stateless "
"render endpoint."
),
err_type="invalid_request_error",
param="previous_response_id",
)
result = await self.online_renderer.render_responses(
request,
previous_messages=None,
previous_response_outputs=None,
tool_server=self.tool_server,
skip_mm_cache=True,
)
if isinstance(result, ErrorResponse):
return result
engine_input = result.engine_input
prompt_components = extract_prompt_components(self.model_config, engine_input)
token_ids = prompt_components.token_ids
if not token_ids:
return self.create_error_response("No token_ids rendered")
input_length = extract_prompt_len(self.model_config, engine_input)
max_tokens = get_max_tokens(
self.model_config.max_model_len,
request.max_output_tokens,
input_length,
self.default_sampling_params,
self.override_max_tokens,
truncate_prompt_tokens=(-1 if request.truncation != "disabled" else None),
)
params = request.to_sampling_params(max_tokens, self.default_sampling_params)
return GenerateRequest(
request_id=request.request_id,
token_ids=list(token_ids),
features=self._extract_mm_features(engine_input),
sampling_params=params,
model=request.model,
stream=bool(request.stream),
cache_salt=request.cache_salt,
priority=request.priority,
kv_transfer_params=request.kv_transfer_params,
ec_transfer_params=request.ec_transfer_params,
token_offsets=engine_input.get("prompt_token_offsets"),
)
def _placeholder_metadata_fields(self, modality: str) -> set[str]:
"""Fields the EC consumer still requires after embeddings transfer."""
if self._placeholder_metadata_parser_failed:
return set()
if self._placeholder_metadata_parser is None:
try:
from vllm.multimodal import MULTIMODAL_REGISTRY
self._placeholder_metadata_parser = (
MULTIMODAL_REGISTRY.create_processor(
self.model_config
).info.data_parser
)
except Exception:
logger.debug(
"Could not load placeholder metadata fields; "
"mm_metadata will use keep_on_cpu fields only."
)
self._placeholder_metadata_parser_failed = True
return set()
return set(
self._placeholder_metadata_parser.placeholder_metadata_fields(modality)
)
def _extract_mm_features(
self,
engine_input: EngineInput,
) -> MultiModalFeatures | None:
return extract_mm_features(
engine_input,
metadata_fields_for=self._placeholder_metadata_fields,
)