diff --git a/renderers/__init__.py b/renderers/__init__.py index 5baf242..7f7a75e 100644 --- a/renderers/__init__.py +++ b/renderers/__init__.py @@ -39,7 +39,7 @@ reject_assistant_in_extension, trim_to_turn_close, ) -from renderers.client import OverlongPromptError +from renderers.client import MalformedGenerateResponseError, OverlongPromptError from renderers.configs import ( AutoRendererConfig, BaseRendererConfig, @@ -152,6 +152,7 @@ def __dir__() -> list[str]: "Llama3Renderer", "Llama3RendererConfig", "MULTIMODAL_MODELS", + "MalformedGenerateResponseError", "Message", "MiniMaxM2Renderer", "MiniMaxM2RendererConfig", diff --git a/renderers/client.py b/renderers/client.py index 026df02..801645f 100644 --- a/renderers/client.py +++ b/renderers/client.py @@ -14,6 +14,7 @@ import asyncio import json import logging +import math from collections.abc import Mapping from typing import Any, cast @@ -33,6 +34,9 @@ _request_logger = logging.getLogger("renderers.client") ROUTED_EXPERTS_DATA_PREFIX = b'"routed_experts":{"data":"' KEPT_TOKENS_IDS_PREFIX = b'"kept_tokens":{"ids":"' +# vLLM uses this value both when sampled-token evidence is missing and as a +# lower-bound clamp, so receiving it cannot prove the real logprob was returned. +VLLM_LOGPROB_SENTINEL = -9999.0 class OverlongPromptError(Exception): @@ -59,6 +63,10 @@ def __init__(self, *, prompt_len: int, max_prompt_len: int) -> None: ) +class MalformedGenerateResponseError(ValueError): + """The generate endpoint returned unusable sampled-token evidence.""" + + # Per-process cache of resolved engine context-length caps, keyed by # ``(base_url, model)``. ``None`` is the "we asked the engine and it didn't # tell us" sentinel — distinct from "key missing" (haven't asked yet). The @@ -150,6 +158,66 @@ def parse_generate_response(raw: bytes) -> dict[str, Any]: return payload +def _parse_completion_logprobs( + choice: Mapping[str, Any], completion_ids: list[int] +) -> list[float]: + raw_logprobs = choice.get("logprobs") + if not isinstance(raw_logprobs, Mapping): + raise MalformedGenerateResponseError( + "Engine response choice.logprobs must be an object." + ) + + content = raw_logprobs.get("content") + if not isinstance(content, list): + raise MalformedGenerateResponseError( + "Engine response choice.logprobs.content must be a list." + ) + if len(content) != len(completion_ids): + raise MalformedGenerateResponseError( + "Engine response completion token count " + f"({len(completion_ids)}) does not match logprob count ({len(content)})." + ) + + completion_logprobs: list[float] = [] + for index, entry in enumerate(content): + if not isinstance(entry, Mapping): + raise MalformedGenerateResponseError( + f"Engine response choice.logprobs.content[{index}] must be an object." + ) + expected_token = f"token_id:{completion_ids[index]}" + if entry.get("token") != expected_token: + raise MalformedGenerateResponseError( + "Engine response " + f"choice.logprobs.content[{index}].token must be {expected_token!r}." + ) + raw_logprob = entry.get("logprob") + if isinstance(raw_logprob, bool) or not isinstance(raw_logprob, (int, float)): + raise MalformedGenerateResponseError( + "Engine response " + f"choice.logprobs.content[{index}].logprob must be a number." + ) + try: + logprob = float(raw_logprob) + except OverflowError as exc: + raise MalformedGenerateResponseError( + "Engine response " + f"choice.logprobs.content[{index}].logprob must be finite." + ) from exc + if not math.isfinite(logprob): + raise MalformedGenerateResponseError( + "Engine response " + f"choice.logprobs.content[{index}].logprob must be finite." + ) + if logprob == VLLM_LOGPROB_SENTINEL: + raise MalformedGenerateResponseError( + "Engine response " + f"choice.logprobs.content[{index}].logprob does not contain " + "sampling evidence." + ) + completion_logprobs.append(logprob) + return completion_logprobs + + async def generate( *, client: AsyncOpenAI, @@ -292,15 +360,12 @@ def _prepare(): choice = (data.get("choices") or [{}])[0] completion_ids = choice.get("token_ids") or [] + completion_logprobs = _parse_completion_logprobs(choice, completion_ids) + parsed = await _maybe_offload( renderer, lambda: renderer.parse_response(completion_ids, tools=tools) ) - # ChatCompletionLogProbs flatten: {"content": [{"logprob": ...}, ...]} - raw_logprobs = choice.get("logprobs") or {} - content_lp = raw_logprobs.get("content") if isinstance(raw_logprobs, dict) else None - completion_logprobs = [float(c.get("logprob") or 0.0) for c in content_lp or []] - routed_experts = choice.get("routed_experts") kept_tokens = choice.get("kept_tokens") diff --git a/tests/test_client.py b/tests/test_client.py index 1cc1000..3a0becd 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -5,6 +5,7 @@ import httpx import numpy as np import pytest +from renderers import MalformedGenerateResponseError from renderers.base import ( ParsedResponse, ParsedToolCall, @@ -67,33 +68,30 @@ class _FakeClient: def __init__(self): self.calls = [] self.base_url = "http://fake-host:8000/v1" + routed_experts = np.array([[[1]], [[2]]], dtype=np.uint8) + self.choice = { + "index": 0, + "token_ids": [7, 8], + "logprobs": { + "content": [ + {"token": "token_id:7", "logprob": -0.1}, + {"token": "token_id:8", "logprob": -0.2}, + ] + }, + "finish_reason": "stop", + "routed_experts": { + "data": base64.b64encode(routed_experts.tobytes()).decode("ascii"), + "shape": list(routed_experts.shape), + }, + } async def post(self, path, *, cast_to=dict, body=None, options=None): self.calls.append( {"path": path, "cast_to": cast_to, "body": body, "options": options} ) - routed_experts = np.array([[[1]], [[2]]], dtype=np.uint8) payload = { "request_id": "gen-test", - "choices": [ - { - "index": 0, - "token_ids": [7, 8], - "logprobs": { - "content": [ - {"token": "token_id:7", "logprob": -0.1}, - {"token": "token_id:8", "logprob": -0.2}, - ] - }, - "finish_reason": "stop", - "routed_experts": { - "data": base64.b64encode(routed_experts.tobytes()).decode( - "ascii" - ), - "shape": list(routed_experts.shape), - }, - } - ], + "choices": [self.choice], } return httpx.Response( 200, @@ -101,6 +99,18 @@ async def post(self, path, *, cast_to=dict, body=None, options=None): ) +def _run_generate(client, renderer=None): + return asyncio.run( + generate( + client=client, + renderer=renderer or _FakeRenderer(), + messages=[{"role": "user", "content": "hi"}], + model="test-model", + tools=[{"type": "function", "function": {"name": "echo"}}], + ) + ) + + def test_generate_builds_request_body_and_parses_response(): client = _FakeClient() renderer = _FakeRenderer() @@ -172,6 +182,95 @@ def test_generate_builds_request_body_and_parses_response(): assert tc.status == ToolCallParseStatus.OK +def test_generate_rejects_missing_completion_logprobs_before_parsing(): + client = _FakeClient() + client.choice.pop("logprobs") + renderer = _FakeRenderer() + + with pytest.raises( + MalformedGenerateResponseError, + match=r"choice\.logprobs must be an object", + ): + _run_generate(client, renderer) + + assert not hasattr(renderer, "_last_parse_tools") + + +@pytest.mark.parametrize( + "entry", + [ + {"token": "token_id:7"}, + {"token": "token_id:7", "logprob": None}, + {"token": "token_id:7", "logprob": "-0.1"}, + {"token": "token_id:7", "logprob": True}, + ], + ids=["missing", "null", "string", "boolean"], +) +def test_generate_rejects_non_numeric_completion_logprobs(entry): + client = _FakeClient() + client.choice["logprobs"]["content"][0] = entry + + with pytest.raises( + MalformedGenerateResponseError, + match=r"content\[0\]\.logprob must be a number", + ): + _run_generate(client) + + +def test_generate_rejects_completion_logprob_count_mismatch(): + client = _FakeClient() + client.choice["logprobs"]["content"] = [{"token": "token_id:7", "logprob": -0.1}] + + with pytest.raises( + MalformedGenerateResponseError, + match=r"completion token count \(2\) does not match logprob count \(1\)", + ): + _run_generate(client) + + +@pytest.mark.parametrize("logprob", [float("nan"), float("inf"), float("-inf")]) +def test_generate_rejects_non_finite_completion_logprobs(logprob): + client = _FakeClient() + client.choice["logprobs"]["content"][0]["logprob"] = logprob + + with pytest.raises( + MalformedGenerateResponseError, + match=r"content\[0\]\.logprob must be finite", + ): + _run_generate(client) + + +def test_generate_rejects_vllm_missing_logprob_sentinel(): + client = _FakeClient() + client.choice["logprobs"]["content"][0]["logprob"] = -9999.0 + + with pytest.raises( + MalformedGenerateResponseError, + match=r"does not contain sampling evidence", + ): + _run_generate(client) + + +def test_generate_rejects_logprob_token_id_mismatch(): + client = _FakeClient() + client.choice["logprobs"]["content"][0]["token"] = "token_id:8" + + with pytest.raises( + MalformedGenerateResponseError, + match=r"content\[0\]\.token must be 'token_id:7'", + ): + _run_generate(client) + + +def test_generate_preserves_zero_completion_logprob(): + client = _FakeClient() + client.choice["logprobs"]["content"][0]["logprob"] = 0.0 + + result = _run_generate(client) + + assert result["completion_logprobs"] == [0.0, -0.2] + + class _MalformedToolRenderer(_FakeRenderer): """Returns only a malformed tool-call attempt — finish_reason must stay "stop"."""