diff --git a/docs/en/advance/structed_output.md b/docs/en/advance/structed_output.md index 1d16d9eb82..4d072955fc 100644 --- a/docs/en/advance/structed_output.md +++ b/docs/en/advance/structed_output.md @@ -106,3 +106,16 @@ response = client.chat.completions.create( top_p=0.8) print(response) ``` + +## Limits + +To protect the server's availability, grammar sources supplied through `response_format` are +validated before compilation, and requests exceeding the following limits are rejected with +HTTP 400: + +- the serialized grammar source (JSON schema, regex, or structural tag) may not exceed 16 KiB; +- a JSON schema may not exceed a nesting depth of 128 raw JSON levels (a nested object spans + two levels through its `properties` wrapper), since grammar compile cost grows exponentially + with depth; +- validation runs off the event loop with a 5-second budget; a schema that still takes longer + is rejected. diff --git a/docs/zh_cn/advance/structed_output.md b/docs/zh_cn/advance/structed_output.md index 81352c6fc6..fc506011ff 100644 --- a/docs/zh_cn/advance/structed_output.md +++ b/docs/zh_cn/advance/structed_output.md @@ -108,3 +108,13 @@ print(response) ``` 输出结果是一个 json 格式的回答。 + +## 限制 + +为保障服务的可用性,通过 `response_format` 提供的语法源在编译前会先做校验,超出以下限制的请求 +将以 HTTP 400 拒绝: + +- 序列化后的语法源(JSON schema、正则或 structural tag)不超过 16 KiB; +- JSON schema 的原始 JSON 嵌套深度不超过 128 层(嵌套对象经由 `properties` 包装,每层占两级), + 因为语法编译开销随深度呈指数增长; +- 校验在事件循环之外的工作线程中执行,预算为 5 秒,超时的 schema 同样会被拒绝。 diff --git a/lmdeploy/_guided_decoding.py b/lmdeploy/_guided_decoding.py index 50736e2e99..4c0dd9fc76 100644 --- a/lmdeploy/_guided_decoding.py +++ b/lmdeploy/_guided_decoding.py @@ -1,11 +1,73 @@ # Copyright (c) OpenMMLab. All rights reserved. """Backend-neutral guided-decoding grammar compilation.""" + from __future__ import annotations +import asyncio import json +from concurrent.futures import ThreadPoolExecutor from functools import lru_cache from typing import Any +# Bounds for client-supplied grammar sources. The compile cost of a nested +# JSON schema grows exponentially with its depth and XGrammar exposes no +# depth or time limit on that path (as of 0.2.x), so unbounded schemas must be +# rejected before compilation. +MAX_GRAMMAR_SOURCE_BYTES = 16 * 1024 +# Raw JSON nesting depth of a grammar source (a nested object spans two JSON +# levels per schema level through its "properties" wrapper). XGrammar's +# compile cost grows exponentially with this depth, so keep the worst +# accepted compile well under a second. Realistic structured-output schemas +# stay far below the limit. +MAX_JSON_NESTING_DEPTH = 128 +# Hard budget for one grammar validation. The work itself runs in a worker +# thread, so this only bounds how long a request may wait on it. +GRAMMAR_COMPILE_TIMEOUT = 5.0 + + +def _max_nesting_depth(value: Any) -> int: + """Return the nesting depth of a JSON-like structure. + + Raises ValueError when the depth exceeds MAX_JSON_NESTING_DEPTH or the + structure contains a cycle. Every container type json.dumps serializes + (dict, list, tuple) counts toward the depth, and scalars do not. + + Recursive within the cap on purpose: the walk stops at + MAX_JSON_NESTING_DEPTH + 1 frames, far below Python's recursion limit, + and at the first cycle instead of spinning. + """ + active: set[int] = set() + + def visit(node: Any, depth: int) -> int: + if not isinstance(node, (dict, list, tuple)): + return depth - 1 + if depth > MAX_JSON_NESTING_DEPTH: + raise ValueError( + f'json_schema exceeds the maximum nesting depth of {MAX_JSON_NESTING_DEPTH}.') + if id(node) in active: + raise ValueError('circular reference detected in grammar source') + active.add(id(node)) + try: + children = node.values() if isinstance(node, dict) else node + return max((visit(child, depth + 1) for child in children), default=depth) + finally: + active.discard(id(node)) + + return visit(value, 1) + + +def _check_source_size(source: str) -> None: + # Count UTF-8 bytes, not code points: a serialized source keeps non-ASCII + # characters (json.dumps with ensure_ascii=False), so code points can + # understate the real size by up to 4x. + size = len(source.encode('utf-8')) + if size > MAX_GRAMMAR_SOURCE_BYTES: + raise ValueError(f'grammar source exceeds the maximum size of {MAX_GRAMMAR_SOURCE_BYTES} bytes.') + + +def _check_schema_depth(schema: Any) -> None: + _max_nesting_depth(schema) + def _json_schema_from_response_format(response_format: dict[str, Any]) -> str: schema: Any = response_format['json_schema'] @@ -15,29 +77,47 @@ def _json_schema_from_response_format(response_format: dict[str, Any]) -> str: schema = schema[key] break if isinstance(schema, (dict, bool)): + _check_schema_depth(schema) return json.dumps(schema, ensure_ascii=False) if isinstance(schema, str): + _check_source_size(schema) + try: + parsed = json.loads(schema) + except RecursionError as err: + raise ValueError(f'json_schema exceeds the maximum nesting depth of {MAX_JSON_NESTING_DEPTH}.') from err + _check_schema_depth(parsed) return schema raise ValueError( f'Cannot parse schema {schema}. The schema must be either a dictionary ' - 'or a string that contains the JSON Schema specification') + 'or a string that contains the JSON Schema specification' + ) def _grammar_source(response_format: dict[str, Any]) -> tuple[str, str]: """Return the XGrammar entry point and its serialized input.""" schema_type = response_format.get('type', 'text') if schema_type == 'json_schema': - return schema_type, _json_schema_from_response_format(response_format) - if schema_type == 'regex_schema': - return schema_type, str(response_format.get('regex_schema', '')) - if schema_type == 'json_object': - schema = json.dumps({'type': 'object', 'additionalProperties': True}) - return 'json_schema', schema - if schema_type == 'structural_tag': - return schema_type, json.dumps(response_format, ensure_ascii=False) - if schema_type == 'text': + source = _json_schema_from_response_format(response_format) + elif schema_type == 'regex_schema': + source = str(response_format.get('regex_schema', '')) + elif schema_type == 'json_object': + source = json.dumps({'type': 'object', 'additionalProperties': True}) + schema_type = 'json_schema' + elif schema_type == 'structural_tag': + # Built by the tool-call parsers and the gpt-oss response parser (the + # latter wraps client-supplied schemas), and also accepted directly by + # GenerationConfig. Only the size bound applies here: XGrammar guards + # the tag's own format tree with its recursion guard, pydantic + # serialization caps how deep the gpt-oss parser can wrap a schema, + # and an embedded schema additionally passes the EBNF parser's fixed + # nest limit. + source = json.dumps(response_format, ensure_ascii=False) + elif schema_type == 'text': return schema_type, '' - raise ValueError(f'unsupported format type: {schema_type}') + else: + raise ValueError(f'unsupported format type: {schema_type}') + _check_source_size(source) + return schema_type, source @lru_cache(maxsize=128) @@ -57,15 +137,38 @@ def _check_response_format(serialized_format: str) -> None: xgr.Grammar.from_structural_tag(source) -def ensure_response_format_compilable(response_format: dict[str, Any]) -> None: - """Reject response formats that XGrammar cannot compile.""" +def _ensure_response_format_compilable(response_format: dict[str, Any]) -> None: + """Synchronous validation body; must not run on an event loop.""" try: serialized_format = json.dumps(response_format, ensure_ascii=False, separators=(',', ':')) _check_response_format(serialized_format) - except (KeyError, RuntimeError, TypeError) as err: + except (KeyError, RecursionError, RuntimeError, TypeError) as err: raise ValueError(f'Unsupported response format: {err}') from err +VALIDATION_EXECUTOR = ThreadPoolExecutor(max_workers=2, thread_name_prefix='grammar-validate') + + +async def ensure_response_format_compilable(response_format: dict[str, Any]) -> None: + """Reject response formats that XGrammar cannot compile. + + The compile runs in a dedicated worker pool under a hard timeout: + XGrammar has no compile-time bound of its own, and a synchronous compile + on the event loop would stall every request, stream, and health check in + the process. A compile that overruns the timeout cannot be cancelled, but + it only ever occupies this small pool, never the interpreter's default + executor shared by other offloaded work. + """ + loop = asyncio.get_running_loop() + try: + await asyncio.wait_for( + loop.run_in_executor(VALIDATION_EXECUTOR, _ensure_response_format_compilable, response_format), + timeout=GRAMMAR_COMPILE_TIMEOUT, + ) + except asyncio.TimeoutError as err: + raise ValueError(f'Response format validation timed out after {GRAMMAR_COMPILE_TIMEOUT}s.') from err + + def compile_response_format(compiler, response_format: dict[str, Any]): """Compile one public or internal response format with XGrammar. diff --git a/lmdeploy/serve/core/async_engine.py b/lmdeploy/serve/core/async_engine.py index 7d8577f082..4510b55835 100644 --- a/lmdeploy/serve/core/async_engine.py +++ b/lmdeploy/serve/core/async_engine.py @@ -530,8 +530,6 @@ def remove_session(): if (messages is not None) ^ (input_ids is None): raise RequestError(ErrorCode.INVALID_REQUEST, 'You must specify exactly one of messages or input_ids.') - if gen_config is not None and gen_config.response_format is not None: - ensure_response_format_compilable(gen_config.response_format) if isinstance(session_id, Session): session = session_id elif isinstance(session_id, int): @@ -545,6 +543,11 @@ def remove_session(): raise RequestError( ErrorCode.REQUEST_CONFLICT, f'Session {session.session_id} already has an active request.') + if gen_config is not None and gen_config.response_format is not None: + # Resolve the session first: a rejected response format raises + # ValueError below, and remove_session() can only clean up a + # session that this method has already taken charge of. + await ensure_response_format_compilable(gen_config.response_format) chat_template_kwargs = chat_template_kwargs or {} if enable_thinking is not None: diff --git a/tests/test_lmdeploy/test_guided_decoding.py b/tests/test_lmdeploy/test_guided_decoding.py new file mode 100644 index 0000000000..68dd47b21b --- /dev/null +++ b/tests/test_lmdeploy/test_guided_decoding.py @@ -0,0 +1,216 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Unit tests for guided-decoding grammar validation bounds. + +The response_format validation must reject unbounded grammar sources before +they reach XGrammar: the compile cost of a nested JSON schema grows +exponentially with depth and XGrammar has no depth or time limit on that +path. These tests run offline; no model, tokenizer, or GPU is required. +""" +import asyncio +import json +import threading +from types import SimpleNamespace +from typing import Any + +import pytest + +from lmdeploy import _guided_decoding as gd +from lmdeploy._guided_decoding import ( + MAX_GRAMMAR_SOURCE_BYTES, + MAX_JSON_NESTING_DEPTH, + _grammar_source, + compile_response_format, + ensure_response_format_compilable, +) +from lmdeploy.serve.core.async_engine import AsyncEngine, RequestError +from lmdeploy.serve.managers.session_manager import SessionManager + + +def nested_object_schema(raw_depth: int, key: str = 'a') -> dict: + """Build a response_format whose raw JSON nesting is exactly ``raw_depth``. + + The common ``properties`` shape spans two JSON levels per schema level; + remaining levels are padded with a plain key. + """ + schema: Any = {'type': 'object'} + node = schema + for _ in range((raw_depth - 1) // 2): + child: dict = {'type': 'object'} + node['properties'] = {key: child} + node = child + for _ in range((raw_depth - 1) % 2): + child: dict = {'type': 'object'} + node['pad'] = child + node = child + return {'type': 'json_schema', 'json_schema': {'name': 't', 'schema': schema}} + + +def json_schema_format(schema) -> dict: + return {'type': 'json_schema', 'json_schema': {'name': 't', 'schema': schema}} + + +def string_schema_format(schema: str) -> dict: + return {'type': 'json_schema', 'json_schema': {'name': 't', 'schema': schema}} + + +class TestGrammarSourceBounds: + def test_schema_depth_boundary(self): + """A schema at the cap passes; one level deeper is rejected in the + bounds check, however deep the input is.""" + _grammar_source(nested_object_schema(MAX_JSON_NESTING_DEPTH)) + with pytest.raises(ValueError, match='nesting depth'): + _grammar_source(nested_object_schema(MAX_JSON_NESTING_DEPTH + 1)) + + def test_deep_string_schema_rejected(self): + # A pre-serialized schema string with parseable over-cap depth. + schema = json.dumps(nested_object_schema(MAX_JSON_NESTING_DEPTH + 1)['json_schema']['schema']) + with pytest.raises(ValueError, match='nesting depth'): + _grammar_source(string_schema_format(schema)) + # Built as a literal so the test itself never recurses: far beyond + # both the depth cap and the JSON parser recursion limit, rejected + # cleanly whichever check trips first. + literal = '{"a":' * 2000 + '{}' + '}' * 2000 + with pytest.raises(ValueError): + _grammar_source(string_schema_format(literal)) + + def test_tuple_nesting_counts_toward_depth(self): + # json.dumps serializes tuples as arrays, so in-process callers can + # smuggle depth through tuples; they must count like lists. + schema: Any = {'type': 'object'} + for _ in range(MAX_JSON_NESTING_DEPTH // 2 + 10): + schema = {'allOf': (schema,)} + with pytest.raises(ValueError, match='nesting depth'): + _grammar_source(json_schema_format(schema)) + + def test_cyclic_schema_rejected_not_hung(self): + # A cyclic dict cannot arrive over the JSON API but can be passed by + # in-process callers; it must fail fast with ValueError on both the + # validation and the engine-side compile path, not spin forever. + schema: dict = {} + schema['properties'] = {'a': schema} + with pytest.raises(ValueError): + _grammar_source(json_schema_format(schema)) + with pytest.raises(ValueError): + compile_response_format(None, json_schema_format(schema)) + + def test_oversized_sources_rejected(self): + with pytest.raises(ValueError, match='maximum size'): + _grammar_source({'type': 'regex_schema', 'regex_schema': 'a' * (MAX_GRAMMAR_SOURCE_BYTES + 1)}) + schema = {'type': 'object', 'properties': {f'k{i}': {'type': 'string'} for i in range(1200)}} + with pytest.raises(ValueError, match='maximum size'): + _grammar_source(json_schema_format(schema)) + # The string branch applies the same byte bound before parsing. + with pytest.raises(ValueError, match='maximum size'): + _grammar_source(string_schema_format(json.dumps('x' * (MAX_GRAMMAR_SOURCE_BYTES + 1)))) + + def test_multibyte_source_counted_as_utf8_bytes(self): + # The size limit is UTF-8 bytes: 6000 CJK characters are only 6000 + # code points but 18000 bytes, so the source must be rejected even + # though a code-point count would pass it. + schema = {'type': 'object', 'properties': {'名字': {'description': '描' * 6000}}} + with pytest.raises(ValueError, match='maximum size'): + _grammar_source(json_schema_format(schema)) + + def test_regular_formats_pass(self): + assert _grammar_source({'type': 'text'}) == ('text', '') + assert _grammar_source({'type': 'json_object'}) == ( + 'json_schema', + json.dumps({'type': 'object', 'additionalProperties': True}), + ) + schema_type, source = _grammar_source({'type': 'regex_schema', 'regex_schema': r'\d{3}-\d{4}'}) + assert (schema_type, source) == ('regex_schema', r'\d{3}-\d{4}') + schema = { + 'type': 'object', + 'properties': { + 'name': {'type': 'string'}, + 'items': { + 'type': 'array', + 'items': {'type': 'number'}, + }, + 'extra': {}, + }, + } + _grammar_source(json_schema_format(schema)) + + def test_engine_compile_path_enforces_bounds(self): + # compile_response_format runs in the engine process; the bounds + # check in _grammar_source must reject before the compiler is used. + with pytest.raises(ValueError, match='nesting depth'): + compile_response_format(None, nested_object_schema(10000)) + + +class TestAsyncValidation: + def test_valid_and_invalid_formats(self): + schema = {'type': 'object', 'properties': {'name': {'type': 'string'}}} + asyncio.run(ensure_response_format_compilable(json_schema_format(schema))) + with pytest.raises(ValueError): + asyncio.run(ensure_response_format_compilable(nested_object_schema(MAX_JSON_NESTING_DEPTH + 1))) + + def test_validation_runs_in_dedicated_pool(self, monkeypatch): + seen_threads = [] + original = gd._check_response_format.__wrapped__ + + def spy(serialized_format): + seen_threads.append(threading.current_thread()) + return original(serialized_format) + + monkeypatch.setattr(gd, '_check_response_format', spy) + schema = {'type': 'object', 'properties': {'name': {'type': 'string'}}} + asyncio.run(ensure_response_format_compilable(json_schema_format(schema))) + assert seen_threads, 'validation did not reach the grammar check' + assert all(t.name.startswith('grammar-validate') for t in seen_threads), ( + 'validation must run in the dedicated grammar pool, not the default executor' + ) + + def test_validation_timeout_becomes_value_error(self, monkeypatch): + monkeypatch.setattr(gd, 'GRAMMAR_COMPILE_TIMEOUT', 0.2) + + def slow_check(_serialized_format): + threading.Event().wait(5) + + monkeypatch.setattr(gd, '_check_response_format', slow_check) + schema = {'type': 'object', 'properties': {'name': {'type': 'string'}}} + with pytest.raises(ValueError, match='timed out'): + asyncio.run(ensure_response_format_compilable(json_schema_format(schema))) + + def test_event_loop_stays_responsive_during_validation(self): + """A real uncached validation must not stall the loop (GIL released by + XGrammar, compile runs in a worker thread).""" + schema = nested_object_schema(120, key='responsive') + + async def main(): + delays = [] + + async def heartbeat(): + while True: + start = asyncio.get_running_loop().time() + await asyncio.sleep(0.02) + delays.append(asyncio.get_running_loop().time() - start - 0.02) + + hb = asyncio.create_task(heartbeat()) + await asyncio.sleep(0.2) + delays.clear() + await ensure_response_format_compilable(schema) + await asyncio.sleep(0.1) + hb.cancel() + return delays + + delays = asyncio.run(main()) + assert len(delays) > 3, 'heartbeat stopped ticking during validation' + assert max(delays) < 1.0, f'event loop stalled for {max(delays):.3f}s during validation' + + def test_failed_validation_removes_passed_session(self): + """A rejected response_format must still clean up a session the caller + created and registered before preprocess ran (api_server hands Session + objects to the engine).""" + engine = AsyncEngine.__new__(AsyncEngine) + engine.session_mgr = SessionManager() + session = engine.session_mgr.get(42) + gen_config = SimpleNamespace(response_format=nested_object_schema(MAX_JSON_NESTING_DEPTH + 1)) + with pytest.raises(RequestError) as exc_info: + asyncio.run( + engine.preprocess(messages=None, input_ids=[1, 2], gen_config=gen_config, session_id=session)) + assert 'nesting depth' in str(exc_info.value) + assert engine.session_mgr.get(42, create_if_not_exists=False) is None, ( + 'rejected request leaked its session in the SessionManager' + )