From 867b2b2ddb6fadeb0554407fde3df35d50f2e2c9 Mon Sep 17 00:00:00 2001 From: svonava Date: Sat, 26 Sep 2026 14:07:32 +0000 Subject: [PATCH] fix(server): limit the label prompt of GLiNER-family extract requests GLiNER, GLiNER2 and GLiREL encode a request's labels, relation types, class labels and schema fields with every document, and bill only the document. Validate that prompt before any model work, as GLiFormer and GLiNER2.5-Decide do, and reject a request with INVALID_INPUT when: - a label, relation type, class label, task name, field name or choice has more than 128 characters (checked before anything is tokenized), or - the prompt takes more than max_prompt_tokens tokens, counted with the model's tokenizer: 1024 for GLiNER and GLiREL, 2048 for GLiNER2, whose schemas carry field descriptions. The option is configurable per model. The defaults sit well above real label sets: a 60-type PII list takes about 230 tokens and a 50-field described schema about 1,300. GLiREL also limits an item to 256 supplied entities, since it scores every pair of them, and reports missing or malformed input (labels, text, entities, entity offsets) as INVALID_INPUT, checked for every item before any is run. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/sie_server/adapters/_prompt_limit.py | 131 +++++++ .../sie_server/adapters/gliner/__init__.py | 42 +++ .../sie_server/adapters/gliner2/adapter.py | 77 ++++- .../sie_server/adapters/glirel/__init__.py | 69 +++- .../tests/adapters/test_gliner2_long_text.py | 15 +- .../tests/adapters/test_prompt_limit.py | 323 ++++++++++++++++++ 6 files changed, 628 insertions(+), 29 deletions(-) create mode 100644 packages/sie_server/src/sie_server/adapters/_prompt_limit.py create mode 100644 packages/sie_server/tests/adapters/test_prompt_limit.py diff --git a/packages/sie_server/src/sie_server/adapters/_prompt_limit.py b/packages/sie_server/src/sie_server/adapters/_prompt_limit.py new file mode 100644 index 000000000..5f3de48b4 --- /dev/null +++ b/packages/sie_server/src/sie_server/adapters/_prompt_limit.py @@ -0,0 +1,131 @@ +"""Limits on the task prompt GLiNER-family adapters encode with each document. + +GLiNER, GLiNER2 and GLiREL encode a request's labels, relation types, class +labels and schema fields (its task prompt) with every document. The prompt is +not billed (input tokens count the document), so it is bounded instead, as +GLiFormer and GLiNER2.5-Decide bound theirs. A request is rejected with +``InvalidInputError`` (HTTP 400 ``INVALID_INPUT``) when + +* a label, relation type, class label, task name, field name or choice has + more than ``MAX_LABEL_CHARS`` characters (checked first, before anything is + tokenized), or +* the prompt takes more than ``max_prompt_tokens`` tokens, counted with the + model's tokenizer. + +The defaults sit far above the label sets these models are used with: the +largest label set among this repository's examples takes about 40 tokens, a +60-type PII list about 230, and a structured-extraction schema of 50 described +fields about 1,200, while ``DEFAULT_MAX_PROMPT_TOKENS`` holds several hundred +entity types and ``DEFAULT_MAX_SCHEMA_PROMPT_TOKENS`` a schema of dozens of +described fields with their choices. + +Descriptions are not limited one by one; the whole prompt is tokenized only +after its characters are checked, so counting costs at most +``MAX_PROMPT_CHARS_PER_TOKEN`` characters of tokenization per allowed token. +""" + +from __future__ import annotations + +import hashlib +from collections import OrderedDict +from collections.abc import Callable, Hashable, Iterable +from typing import Any + +from sie_server.types.inputs import InvalidInputError + +# Characters a label, relation type, class label, task name, field name or choice may have. +MAX_LABEL_CHARS = 128 +# Tokens a request's labels, relation types and class labels may take. +DEFAULT_MAX_PROMPT_TOKENS = 1024 +# Tokens a GLiNER2 request's labels or schema (field names, descriptions and choices) may take. +DEFAULT_MAX_SCHEMA_PROMPT_TOKENS = 2048 +# No token of these tokenizers covers more characters than this. +MAX_PROMPT_CHARS_PER_TOKEN = 32 +_PROMPT_CACHE_SIZE = 256 + + +def validate_max_prompt_tokens(value: object) -> int: + """A ``max_prompt_tokens`` adapter option, checked. + + Raises: + ValueError: The value is not a positive integer. + """ + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + raise ValueError("max_prompt_tokens must be a positive integer") + return value + + +def check_label_chars(model: str, kind: str, values: Iterable[str]) -> None: + """Reject a label (or relation type, class label, field name, choice) of more than ``MAX_LABEL_CHARS`` characters. + + Raises: + InvalidInputError: A value is not a string or is too long. + """ + for value in values: + if not isinstance(value, str): + raise InvalidInputError(f"{model} {kind} must be strings") + if len(value) > MAX_LABEL_CHARS: + raise InvalidInputError(f"{model} {kind} may have at most {MAX_LABEL_CHARS} characters each") + + +class PromptLimit: + """Checks a request's task prompt against ``max_tokens``, remembering recent prompts' sizes.""" + + __slots__ = ("_counts", "max_tokens", "model") + + def __init__(self, model: str, max_tokens: int) -> None: + self.model = model + self.max_tokens = validate_max_prompt_tokens(max_tokens) + self._counts: OrderedDict[bytes, int] = OrderedDict() + + def check(self, texts: Iterable[str], count: Callable[[], int], key: Hashable) -> int: + """The prompt's tokens, from ``count()``, after checking ``texts`` (its strings) and the result. + + ``key`` identifies the prompt (its strings, in order, and anything else + its size depends on); a prompt seen recently is not counted again. + + Raises: + InvalidInputError: The prompt takes more than ``max_tokens`` tokens. + """ + digest = hashlib.sha256(repr(key).encode("utf-8", "surrogatepass")).digest() + tokens = self._counts.get(digest) + if tokens is None: + chars = sum(len(text) for text in texts) + if chars > self.max_tokens * MAX_PROMPT_CHARS_PER_TOKEN: + raise InvalidInputError(self._message(None, chars)) + tokens = int(count()) + self._counts[digest] = tokens + if len(self._counts) > _PROMPT_CACHE_SIZE: + self._counts.popitem(last=False) + else: + self._counts.move_to_end(digest) + if tokens > self.max_tokens: + raise InvalidInputError(self._message(tokens, None)) + return tokens + + def _message(self, tokens: int | None, chars: int | None) -> str: + size = f"{tokens} tokens" if tokens is not None else f"{chars} characters" + return ( + f"{self.model} labels, relation types, class labels and schema fields take {size}; " + f"a request may use at most {self.max_tokens} tokens for them" + ) + + +def gliner_prompt_counter(model: Any) -> Callable[[list[str], list[str]], int]: + """Tokens of the prompt a loaded ``gliner`` model builds for entity and relation types. + + The prompt is built by the model's own processor (``prepare_inputs``, with + no document words) and tokenized as the processor tokenizes it. + """ + processor = model.data_processor + tokenizer = processor.transformer_tokenizer + + def count(entity_types: list[str], relation_types: list[str]) -> int: + kwargs = {"relations": relation_types} if relation_types else {} + (words,), _ = processor.prepare_inputs([[]], entity_types, **kwargs) + if not words: + return 0 + encoding = tokenizer(list(words), is_split_into_words=True, add_special_tokens=False) + return len(encoding["input_ids"]) + + return count diff --git a/packages/sie_server/src/sie_server/adapters/gliner/__init__.py b/packages/sie_server/src/sie_server/adapters/gliner/__init__.py index 4bb803f7c..84a410836 100644 --- a/packages/sie_server/src/sie_server/adapters/gliner/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/gliner/__init__.py @@ -17,6 +17,11 @@ Joint entity-relation ("relex") models also extract relations between the entities they find when a request names relation types in ``options["relation_labels"]``. Without it they return entities only. + +A request's labels and relation types, which GLiNER encodes with every +document and does not bill, may have at most 128 characters each and take at +most ``max_prompt_tokens`` tokens together (default 1024); a longer prompt is +rejected with ``INVALID_INPUT``. """ import math @@ -27,6 +32,12 @@ import torch from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import ( + DEFAULT_MAX_PROMPT_TOKENS, + PromptLimit, + check_label_chars, + gliner_prompt_counter, +) from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import bound_gliner_words, plan_forwards @@ -96,6 +107,7 @@ def __init__( multi_label: bool = False, merge_adjacent_entities: bool = False, relation_threshold: float | None = None, + max_prompt_tokens: int = DEFAULT_MAX_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, # Accept extra args from loader (e.g., pooling) @@ -112,6 +124,9 @@ def __init__( relation_threshold: Minimum relation score (0-1) for joint entity-relation models. None uses the entity threshold, as the gliner library does. + max_prompt_tokens: Most tokens a request's labels and relation + types may take in the prompt encoded with each document (see + ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -133,6 +148,9 @@ def __init__( self._extracts_relations = False # True when the encoder's attention memory grows with the square of a row (see ``_inference``). self._quadratic_attention = False + self._prompt_limit = PromptLimit("GLiNER", max_prompt_tokens) + # Tokens of the label prompt, as the loaded model builds it; None until loaded. + self._count_prompt: Any = None def load(self, device: str) -> None: """Load the model onto the specified device. @@ -173,6 +191,7 @@ def load(self, device: str) -> None: # gliner's max_len counts words, whatever their subwords: read at most a # bounded number of subwords too, with a long word in pieces. self._quadratic_attention = bound_gliner_words(self._model) + self._count_prompt = gliner_prompt_counter(self._model) def extract( self, @@ -227,6 +246,8 @@ def extract( if relation_labels and not self._extracts_relations: raise InvalidInputError(_ERR_NO_RELATIONS) + self._check_prompt(labels, relation_labels) + # Extract texts from all items texts = [self._extract_text(item) for item in items] if any(not text.strip() for text in texts): @@ -303,6 +324,27 @@ def extract( return ExtractOutput(entities=all_entities, relations=all_relations, input_token_counts=input_token_counts) + def _check_prompt(self, labels: list[str], relation_labels: list[str]) -> None: + """Reject a request whose labels and relation types take more than ``max_prompt_tokens``. + + gliner encodes the label prompt with every document, and only the + document is billed. + + Raises: + InvalidInputError: The prompt is too long, or a label is not a string. + """ + check_label_chars("GLiNER", "labels", labels) + check_label_chars("GLiNER", "relation_labels", relation_labels) + entity_types = list(dict.fromkeys(labels)) # gliner drops repeated labels + count = self._count_prompt + + def tokens() -> int: + return count(entity_types, relation_labels) if count is not None else 0 + + self._prompt_limit.check( + [*entity_types, *relation_labels], tokens, (tuple(entity_types), tuple(relation_labels)) + ) + @staticmethod def _validate_relation_labels(value: Any, labels: list[str]) -> list[str]: """Return the requested relation types (empty when none were asked for).""" diff --git a/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py b/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py index 911acb6d8..6623bc2d6 100644 --- a/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py +++ b/packages/sie_server/src/sie_server/adapters/gliner2/adapter.py @@ -13,6 +13,7 @@ from huggingface_hub import snapshot_download from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, PromptLimit, check_label_chars from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import ( @@ -76,6 +77,12 @@ class GLiNER2Adapter(BaseAdapter): - Batch methods cover entities, relations, structured data, and classification - Classification uses ``classify_text()`` / ``batch_classify_text()`` + A request's labels, class labels, relation types or schema fields (field + names, descriptions and choices), which gliner2 encodes with every document + and does not bill, may take at most ``max_prompt_tokens`` tokens (default + 2048), and each label, task name, field name or choice at most 128 + characters; a longer prompt is rejected with ``INVALID_INPUT``. + Reference models: - fastino/gliner2-base-v1 - fastino/gliner2-large-v1 @@ -98,6 +105,7 @@ def __init__( default_labels: list[str] | None = None, multi_label: bool = False, max_seq_length: int | None = None, + max_prompt_tokens: int = DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, @@ -115,6 +123,9 @@ def __init__( multi_label: Whether the configured classification task may return multiple labels. max_seq_length: Maximum document and schema input length. + max_prompt_tokens: Most tokens a request's labels, class labels, + relation types or schema fields may take in the task prompt + encoded with each document (see ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -127,6 +138,7 @@ def __init__( self._default_labels = self._validate_labels(default_labels) if default_labels is not None else None self._multi_label = multi_label self._max_seq_length = max_seq_length + self._prompt_limit = PromptLimit("GLiNER2", max_prompt_tokens) self._compute_precision = compute_precision self._revision = revision @@ -254,8 +266,14 @@ def extract( raise ValueError("GLiNER2 structured extraction does not accept classification_task") structures = self._json_schema_to_structures(output_schema) specs = [spec for fields in structures.values() for spec in fields] - # A field's choices are read twice: in its structure and in a prefix before the document. - rows = self._row_tokens(windows, specs + specs) + # A field's choices are listed in its structure, and each again in a prefix before the document. + choices = [ + choice for definition in output_schema["properties"].values() for choice in definition.get("enum") or [] + ] + check_label_chars("GLiNER2", "output_schema property names", output_schema["properties"]) + check_label_chars("GLiNER2", "output_schema enum values", choices) + prompt = self._prompt_tokens(specs, key=("json", tuple(specs)), extra=len(choices)) + rows = self._row_tokens(windows, prompt) with torch.inference_mode(): raw_results = self._run_planned( model_texts, @@ -286,7 +304,11 @@ def extract( normalized_entities = [ self._normalize_input_entities(item, entities or []) for item, entities in zip(items, relation_entities) ] - rows = self._row_tokens(windows, normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens( + normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION, key=("relations", tuple(normalized_labels)) + ) + rows = self._row_tokens(windows, prompt) with torch.inference_mode(): raw_results = self._run_planned( model_texts, @@ -314,6 +336,12 @@ def extract( if classification_task is not None: if not isinstance(classification_task, str) or not classification_task.strip(): raise ValueError("GLiNER2 classification_task must be a non-empty string") + check_label_chars("GLiNER2", "classification_task", [classification_task]) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens( + [classification_task, *normalized_labels], + key=("classification", classification_task, tuple(normalized_labels)), + ) return self._classify( model_texts, normalized_labels, @@ -321,7 +349,7 @@ def extract( multi_label=multi_label, threshold=effective_threshold, input_token_counts=input_token_counts, - rows=self._row_tokens(windows, [classification_task, *normalized_labels]), + rows=self._row_tokens(windows, prompt), ) def extract_entities(batch: list[str]) -> list[Any]: @@ -345,10 +373,12 @@ def extract_entities(batch: list[str]) -> list[Any]: max_len=self._max_seq_length, ) + check_label_chars("GLiNER2", "labels", normalized_labels) + prompt = self._prompt_tokens(normalized_labels, key=("entities", tuple(normalized_labels))) with torch.inference_mode(): raw_results = self._run_planned( model_texts, - self._row_tokens(windows, normalized_labels), + self._row_tokens(windows, prompt), extract_entities, rows_per_pass=1 if len(texts) == 1 else _PACKAGE_BATCH_SIZE, ) @@ -412,22 +442,43 @@ def classify(batch: list[str]) -> list[Any]: input_token_counts=input_token_counts, ) - def _row_tokens( + def _prompt_tokens( self, - windows: list[tuple[str, int | None]], - prompt_entries: Iterable[str], + entries: list[str], *, + key: tuple[Any, ...], per_entry: int = _PROMPT_TOKENS_PER_ENTRY, - ) -> list[int] | None: + extra: int = 0, + ) -> int | None: + """Estimated tokens of the task prompt gliner2 builds from ``entries``, checked against the limit. + + Each label, class label, relation type and schema field is counted + with the tokens gliner2 adds around it, plus ``extra`` (a token per + field choice, which gliner2 lists again before the document). None when words are + not counted (no bounded splitter is installed); the prompt's + characters are still checked. + + Raises: + InvalidInputError: The prompt takes more than ``max_prompt_tokens``. + """ + count = self._count_subwords + strings = [entry for entry in entries if isinstance(entry, str)] + + def tokens() -> int: + if count is None: + return 0 + return sum(count(strings)) + per_entry * len(strings) + extra + _ROW_OVERHEAD_TOKENS + + prompt = self._prompt_limit.check(strings, tokens, (per_entry, extra, *key)) + return prompt if count is not None else None + + def _row_tokens(self, windows: list[tuple[str, int | None]], prompt: int | None) -> list[int] | None: """Estimated tokens of each item's encoder row: the task prompt, then the words it reads. None when the words were not counted (no bounded splitter is installed). """ - count = self._count_subwords - if count is None or any(subwords is None for _, subwords in windows): + if prompt is None or any(subwords is None for _, subwords in windows): return None - entries = [entry for entry in prompt_entries if isinstance(entry, str)] - prompt = sum(count(entries)) + per_entry * len(entries) + _ROW_OVERHEAD_TOKENS return [prompt + (subwords or 0) for _, subwords in windows] def _run_planned( diff --git a/packages/sie_server/src/sie_server/adapters/glirel/__init__.py b/packages/sie_server/src/sie_server/adapters/glirel/__init__.py index 9d5df2d26..3d520e353 100644 --- a/packages/sie_server/src/sie_server/adapters/glirel/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/glirel/__init__.py @@ -7,6 +7,11 @@ Reference models: - jackboyla/glirel-large-v0 (zero-shot relation extraction) - jackboyla/glirel_re_large-v0 (relation-focused variant) + +A request's relation types, which GLiREL encodes with every text, may have at +most 128 characters each and take at most ``max_prompt_tokens`` tokens together +(default 1024), and an item may carry at most ``MAX_ENTITIES`` (256) entities; +other requests are rejected with ``INVALID_INPUT``. """ import re @@ -17,6 +22,7 @@ import torch from sie_server.adapters._base_adapter import BaseAdapter +from sie_server.adapters._prompt_limit import DEFAULT_MAX_PROMPT_TOKENS, PromptLimit, check_label_chars from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision from sie_server.adapters._word_window import SubwordCounter, WindowedSplitter, split_word_counter, subword_budget @@ -27,6 +33,9 @@ # Error messages _ERR_REQUIRES_LABELS = "GLiREL requires labels parameter for relation extraction" _ERR_REQUIRES_ENTITIES = "GLiREL requires entities in item metadata for relation extraction" +# GLiREL scores every pair of an item's entities (in Python while preparing +# the batch, and on the GPU), so an item may carry at most this many. +MAX_ENTITIES = 256 _TOKEN_PATTERN = re.compile(r"\w+(?:[-_]\w+)*|\S") # Subword tokens a text may take per word GLiREL reads: its checkpoints read # English, whose prose, code and logs run at up to 2.4 subwords per word. @@ -72,6 +81,7 @@ def __init__( model_name_or_path: str | Path, *, threshold: float = 0.3, + max_prompt_tokens: int = DEFAULT_MAX_PROMPT_TOKENS, compute_precision: ComputePrecision = "float16", revision: str | None = None, **kwargs: Any, # Accept extra args from loader @@ -81,6 +91,8 @@ def __init__( Args: model_name_or_path: HuggingFace model ID or local path to GLiREL model. threshold: Minimum confidence score for relation extraction (0-1). + max_prompt_tokens: Most tokens a request's relation types may take + in the prompt encoded with each text (see ``_prompt_limit``). compute_precision: Compute precision for inference. revision: Optional HuggingFace revision/branch/commit SHA to pin when loading model artifacts. @@ -96,6 +108,9 @@ def __init__( self._device: str | None = None # The words of a text GLiREL reads (see ``_tokenize``); None until loaded. self._words: WindowedSplitter | None = None + self._prompt_limit = PromptLimit("GLiREL", max_prompt_tokens) + # Tokens of the relation-type prompt, as the loaded model builds it; None until loaded. + self._count_prompt: Any = None def load(self, device: str) -> None: """Load the model onto the specified device. @@ -126,6 +141,9 @@ def load(self, device: str) -> None: int(self._model.base_config.max_len), getattr(embeddings.model, "config", None), ) + self._count_prompt = _prompt_counter( + embeddings.tokenizer, str(self._model.rel_token), str(self._model.sep_token) + ) def _bound_words(self, tokenizer: Any, max_words: int, encoder_config: Any = None) -> None: """Read at most ``max_words`` words of a text, and a bounded number of their subwords. @@ -177,23 +195,25 @@ def extract( Raises: RuntimeError: If model not loaded. - ValueError: If labels not provided or items lack entities. - InvalidInputError: If an entity is not an object or has invalid offsets. + InvalidInputError: If labels are missing or too long, or an item lacks + text or entities, carries more than ``MAX_ENTITIES`` entities, or + has invalid entity offsets. """ self._check_loaded() if not labels: - raise ValueError(_ERR_REQUIRES_LABELS) + raise InvalidInputError(_ERR_REQUIRES_LABELS) + self._check_prompt(labels) - # Check every item's entities before running any item. + # Check every item before running any, so bad input fails as a 400 before model work. inputs: list[tuple[str, list[dict[str, Any]]]] = [] for item in items: text = self._extract_text(item) entities = self._extract_entities(item) - if not entities: - raise ValueError(_ERR_REQUIRES_ENTITIES) - + raise InvalidInputError(_ERR_REQUIRES_ENTITIES) + if len(entities) > MAX_ENTITIES: + raise InvalidInputError(f"GLiREL items may carry at most {MAX_ENTITIES} entities in metadata") for entity in entities: self._validate_entity_span(entity, text) inputs.append((text, entities)) @@ -256,10 +276,27 @@ def extract( return ExtractOutput(entities=all_entities, relations=all_relations) + def _check_prompt(self, labels: list[str]) -> None: + """Reject a request whose relation types take more than ``max_prompt_tokens``. + + GLiREL encodes the relation types with every text, and GLiREL output + carries no input token count. + + Raises: + InvalidInputError: The prompt is too long, or a relation type is not a string. + """ + check_label_chars("GLiREL", "labels", labels) + count = self._count_prompt + + def tokens() -> int: + return count(labels) if count is not None else 0 + + self._prompt_limit.check(labels, tokens, tuple(labels)) + def _extract_text(self, item: Item) -> str: """Extract text from an item.""" if item.text is None: - raise ValueError(ERR_REQUIRES_TEXT.format(adapter_name="GLiREL adapter")) + raise InvalidInputError(ERR_REQUIRES_TEXT.format(adapter_name="GLiREL adapter")) return item.text def _extract_entities(self, item: Item) -> list[dict[str, Any]]: @@ -267,7 +304,10 @@ def _extract_entities(self, item: Item) -> list[dict[str, Any]]: metadata = item.metadata if metadata is None: return [] - return metadata.get("entities", []) + entities = metadata.get("entities", []) + if not isinstance(entities, list): + raise InvalidInputError("GLiREL item metadata.entities must be a list") + return entities def _tokenize(self, text: str) -> tuple[list[str], list[tuple[int, int]], int]: """Tokenize text like GLiREL, keeping the words it reads: ``(words, offsets, end of the text read)``. @@ -376,6 +416,17 @@ def _relation_entity_text( return str(relation_text) +def _prompt_counter(tokenizer: Any, rel_token: str, sep_token: str) -> Any: + """Tokens of GLiREL's prompt for a list of relation types: ``[REL] type ... [REL] type [SEP]``.""" + + def count(labels: list[str]) -> int: + words = [word for label in labels for word in (rel_token, label)] + [sep_token] + encoding = tokenizer(words, is_split_into_words=True, add_special_tokens=False) + return len(encoding["input_ids"]) + + return count + + def _words(text: str) -> Iterator[tuple[str, int, int]]: """GLiREL's words of ``text`` with their character offsets.""" for match in _TOKEN_PATTERN.finditer(text): diff --git a/packages/sie_server/tests/adapters/test_gliner2_long_text.py b/packages/sie_server/tests/adapters/test_gliner2_long_text.py index 3a5e0b3ec..cb87b2921 100644 --- a/packages/sie_server/tests/adapters/test_gliner2_long_text.py +++ b/packages/sie_server/tests/adapters/test_gliner2_long_text.py @@ -373,14 +373,15 @@ def test_the_prefix_reads_what_gliner2_reads_with_its_sentence_end(text: str, ma def test_relation_rows_run_in_passes_within_the_attention_budget() -> None: # gliner2 builds a structure of about ten tokens around each relation type. - adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}) + adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}, max_prompt_tokens=8192) labels = [f"rel{index}" for index in range(150)] text = " ".join(["abcdefghi"] * 400) entities = [{"text": "abcdefghi", "label": "x", "start": 0, "end": 9}] adapter.extract([Item(text=text, metadata={"entities": entities}) for _ in range(8)], labels=labels) - estimated = adapter._row_tokens([adapter._window(text)] * 1, labels, per_entry=_PROMPT_TOKENS_PER_RELATION) + prompt = adapter._prompt_tokens(labels, per_entry=_PROMPT_TOKENS_PER_RELATION, key=("relations", tuple(labels))) + estimated = adapter._row_tokens([adapter._window(text)], prompt) assert estimated is not None for batch in model.inputs: rows, width = batch.input_ids.shape @@ -389,7 +390,7 @@ def test_relation_rows_run_in_passes_within_the_attention_budget() -> None: def test_structured_rows_are_not_underestimated() -> None: - adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}) + adapter, model = make_adapter(encoder_config={"model_type": "deberta-v2"}, max_prompt_tokens=8192) schema = { "type": "object", "properties": { @@ -401,11 +402,11 @@ def test_structured_rows_are_not_underestimated() -> None: adapter.extract([Item(text=text)], output_schema=schema) - structures = adapter._json_schema_to_structures(schema) - specs = [spec for fields in structures.values() for spec in fields] - estimated = adapter._row_tokens([adapter._window(text)], specs + specs) + (prompt,) = adapter._prompt_limit._counts.values() # the estimate extract() planned with + estimated = adapter._row_tokens([adapter._window(text)], prompt) assert estimated is not None - assert model.inputs[-1].input_ids.shape[1] <= estimated[0] + # The schema is estimated from its field specs, not gliner2's own word split of it. + assert model.inputs[-1].input_ids.shape[1] <= estimated[0] * 1.02 @pytest.mark.parametrize("max_len", [1, 2]) diff --git a/packages/sie_server/tests/adapters/test_prompt_limit.py b/packages/sie_server/tests/adapters/test_prompt_limit.py new file mode 100644 index 000000000..266203cd5 --- /dev/null +++ b/packages/sie_server/tests/adapters/test_prompt_limit.py @@ -0,0 +1,323 @@ +"""The limit on the task prompt (labels, relation types, class labels, schema fields) of GLiNER-family adapters.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import MagicMock + +import pytest +from sie_server.adapters._prompt_limit import ( + DEFAULT_MAX_PROMPT_TOKENS, + DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, + MAX_LABEL_CHARS, + MAX_PROMPT_CHARS_PER_TOKEN, + PromptLimit, +) +from sie_server.adapters.gliner import GLiNERAdapter +from sie_server.adapters.gliner2.classification import GLiNER2ClassificationAdapter +from sie_server.adapters.glirel import MAX_ENTITIES, GLiRELAdapter +from sie_server.types.inputs import InvalidInputError, Item + +from .test_gliner2_long_text import make_adapter as make_gliner2_adapter +from .test_gliner_long_words import DEBERTA, FakeGLiNER, char_tokenizer, encoded, load + +# The largest label set among the repository's examples (12 labels). The stand-in +# tokenizer takes a token per character, so these take about 200 tokens. +EXAMPLE_LABELS = [ + "prompt injection", + "jailbreak", + "role play", + "system prompt leak", + "harmful request", + "data exfiltration", + "instruction override", + "obfuscation", + "social engineering", + "policy evasion", + "privilege escalation", + "benign", +] +# Sixty short types (about 300 tokens with the stand-in tokenizer). +PII_LABELS = [f"pii{index}" for index in range(60)] + + +class Counting: + def __init__(self, tokens: int) -> None: + self.tokens = tokens + self.calls = 0 + + def __call__(self) -> int: + self.calls += 1 + return self.tokens + + +def test_a_prompt_within_the_limit_passes_and_is_counted_once() -> None: + limit = PromptLimit("Model", 100) + count = Counting(40) + + assert limit.check(["a", "b"], count, ("a", "b")) == 40 + assert limit.check(["a", "b"], count, ("a", "b")) == 40 + assert count.calls == 1 + + +def test_a_prompt_over_the_limit_is_rejected_every_time() -> None: + limit = PromptLimit("Model", 100) + count = Counting(101) + + for _ in range(2): + with pytest.raises(InvalidInputError, match="take 101 tokens; a request may use at most 100"): + limit.check(["a"], count, ("a",)) + assert count.calls == 1 + + +def test_a_prompt_of_too_many_characters_is_rejected_before_it_is_tokenized() -> None: + limit = PromptLimit("Model", 10) + count = Counting(1) + + with pytest.raises(InvalidInputError, match="characters"): + limit.check(["x" * (10 * MAX_PROMPT_CHARS_PER_TOKEN + 1)], count, ("long",)) + assert count.calls == 0 + + +@pytest.mark.parametrize("value", [0, -1, True, 1.5, "1024"]) +def test_the_limit_must_be_a_positive_integer(value: Any) -> None: + with pytest.raises(ValueError, match="positive integer"): + PromptLimit("Model", value) + + +def test_gliner_counts_the_prompt_its_processor_builds() -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + labels = ["person", "organization"] + + tokens = adapter._count_prompt(labels, []) + + # The encoded row is [CLS], the prompt, the document's subwords, [SEP]. + (row,) = encoded(model, ["abc"], labels) + assert tokens == len(row) - 3 - 2 + + +@pytest.mark.parametrize("labels", [EXAMPLE_LABELS, PII_LABELS], ids=["example-labels", "pii-list"]) +def test_gliner_accepts_ordinary_label_sets(labels: list[str]) -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=labels) + + assert model.calls == [["Priya Raman works at Novartis."]] + + +@pytest.mark.parametrize( + "labels", + [[f"{index}" + "x" * 120 for index in range(12)], [f"type number {index}" for index in range(100)]], + ids=["long-labels", "many-labels"], +) +def test_gliner_rejects_a_prompt_over_the_limit_before_inference(labels: list[str]) -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + + with pytest.raises(InvalidInputError, match=f"at most {DEFAULT_MAX_PROMPT_TOKENS} tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=labels) + assert model.calls == [] + + +def test_gliner_limit_is_configurable() -> None: + model = FakeGLiNER(DEBERTA) + adapter = GLiNERAdapter("fake/gliner", max_prompt_tokens=20) + adapter._model = model + adapter._count_prompt = load(GLiNERAdapter, FakeGLiNER(DEBERTA))._count_prompt + + with pytest.raises(InvalidInputError, match="at most 20 tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=["organization", "location"]) + + +def glirel_adapter() -> tuple[GLiRELAdapter, MagicMock]: + adapter = GLiRELAdapter("fake/glirel") + adapter._model = MagicMock() + adapter._model.predict_relations.return_value = [] + adapter._bound_words(char_tokenizer(), 64, DEBERTA) + from sie_server.adapters.glirel import _prompt_counter + + adapter._count_prompt = _prompt_counter(char_tokenizer(), "[", "]") + return adapter, adapter._model.predict_relations + + +def test_glirel_rejects_relation_types_over_the_limit() -> None: + adapter, predict = glirel_adapter() + item = Item( + text="Tim Cook leads Apple", + metadata={"entities": [{"text": "Tim Cook", "label": "PERSON", "start": 0, "end": 8}]}, + ) + + adapter.extract([item], labels=["ceo of", "works for"]) + with pytest.raises(InvalidInputError, match="GLiREL"): + adapter.extract([item], labels=[f"{i}" + "z" * 500 for i in range(3)]) + assert predict.call_count == 1 + + +def test_gliner2_accepts_ordinary_label_sets_and_schemas() -> None: + adapter, model = make_gliner2_adapter() + schema = { + "type": "object", + "properties": { + f"field{index}": {"type": "string", "description": "the value of a field on an insurance claim form"} + for index in range(30) + }, + } + + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=PII_LABELS) + adapter._model.batch_extract_json = MagicMock(return_value=[{}]) + adapter.extract([Item(text="Priya Raman works at Novartis")], output_schema=schema) + + assert len(model.inputs) == 1 + adapter._model.batch_extract_json.assert_called_once() + + +@pytest.mark.parametrize( + ("labels", "output_schema"), + [ + ([f"{index:03d}" + "w" * 117 for index in range(60)], None), + (None, {"type": "object", "properties": {"f": {"type": "string", "description": "d" * 70_000}}}), + ( + None, + { + "type": "object", + "properties": { + f"field{index}": {"type": "string", "enum": [f"choice {index} {option}" for option in range(40)]} + for index in range(40) + }, + }, + ), + ], + ids=["long-labels", "long-description", "many-choices"], +) +def test_gliner2_rejects_a_prompt_over_the_limit_before_inference( + labels: list[str] | None, output_schema: dict[str, Any] | None +) -> None: + adapter, model = make_gliner2_adapter() + adapter._model.batch_extract_json = MagicMock(return_value=[{}]) + + with pytest.raises(InvalidInputError, match=f"at most {DEFAULT_MAX_SCHEMA_PROMPT_TOKENS} tokens"): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=labels, output_schema=output_schema) + assert model.inputs == [] + adapter._model.batch_extract_json.assert_not_called() + + +def test_gliner2_rejects_a_classification_prompt_over_the_limit() -> None: + adapter, model = make_gliner2_adapter(classification_task="topic") + + with pytest.raises(InvalidInputError, match="GLiNER2"): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=[f"t{index}" * 40 for index in range(80)]) + assert model.inputs == [] + + +class NoCounting: + def __call__(self, *args: Any, **kwargs: Any) -> Any: + raise AssertionError("a label over the character limit must be rejected before it is tokenized") + + +def test_gliner_rejects_a_long_label_before_tokenizing_it() -> None: + model = FakeGLiNER(DEBERTA) + adapter = load(GLiNERAdapter, model) + adapter._count_prompt = NoCounting() + + with pytest.raises(InvalidInputError, match=f"at most {MAX_LABEL_CHARS} characters"): + adapter.extract([Item(text="Priya Raman works at Novartis.")], labels=["q" * (4 * 1024 * 1024)]) + assert model.calls == [] + + +def test_gliner_rejects_long_relation_types() -> None: + adapter = load(GLiNERAdapter, FakeGLiNER(DEBERTA)) + adapter._extracts_relations = True + + with pytest.raises(InvalidInputError, match="relation_labels may have at most"): + adapter.extract( + [Item(text="Priya Raman works at Novartis.")], + labels=["person"], + options={"relation_labels": ["r" * (MAX_LABEL_CHARS + 1)]}, + ) + + +@pytest.mark.parametrize( + ("labels", "output_schema", "match"), + [ + (["x" * (MAX_LABEL_CHARS + 1)], None, "labels may have at most"), + (None, {"type": "object", "properties": {"n" * (MAX_LABEL_CHARS + 1): {"type": "string"}}}, "property names"), + ( + None, + {"type": "object", "properties": {"f": {"type": "string", "enum": ["c" * (MAX_LABEL_CHARS + 1)]}}}, + "enum values", + ), + ], + ids=["label", "field-name", "choice"], +) +def test_gliner2_rejects_long_labels_field_names_and_choices( + labels: list[str] | None, output_schema: dict[str, Any] | None, match: str +) -> None: + adapter, model = make_gliner2_adapter() + adapter._count_subwords = NoCounting() + + with pytest.raises(InvalidInputError, match=match): + adapter.extract([Item(text="Priya Raman works at Novartis")], labels=labels, output_schema=output_schema) + assert model.inputs == [] + + +def test_gliner2_classifier_rejects_a_long_task_name_and_long_class_labels() -> None: + # The GLiGuard route (gliner2 2.x in the transformers5 bundle) uses the same checks. + adapter, model = make_gliner2_adapter(GLiNER2ClassificationAdapter, classification_task="prompt_safety") + + with pytest.raises(InvalidInputError, match="classification_task may have at most"): + adapter.extract([Item(text="Hello")], labels=["safe", "unsafe"], options={"classification_task": "t" * 200}) + with pytest.raises(InvalidInputError, match="GLiNER2"): + adapter.extract([Item(text="Hello")], labels=[f"class {index:03d} " + "k" * 110 for index in range(60)]) + adapter.extract([Item(text="Hello")], labels=["safe", "unsafe"]) + assert len(model.inputs) == 1 + + +def glirel_item(entities: list[Any], text: str = "Tim Cook leads Apple") -> Item: + return Item(text=text, metadata={"entities": entities}) + + +def test_glirel_caps_the_entities_of_an_item() -> None: + adapter, predict = glirel_adapter() + text = "word " * (MAX_ENTITIES + 1) + entities = [ + {"text": "word", "label": "X", "start": 5 * index, "end": 5 * index + 4} for index in range(MAX_ENTITIES + 1) + ] + + with pytest.raises(InvalidInputError, match=f"at most {MAX_ENTITIES} entities"): + adapter.extract([glirel_item(entities, text)], labels=["ceo of"]) + adapter.extract([glirel_item(entities[:MAX_ENTITIES], text)], labels=["ceo of"]) + assert predict.call_count == 1 + + +@pytest.mark.parametrize( + "entity", + [ + {"text": "x", "label": "ORG", "start": 5, "end": 2}, + {"text": "x", "label": "ORG", "start": "0", "end": 3}, + {"text": "x", "label": "ORG", "start": 500, "end": 505}, + "Tim Cook", + ], + ids=["reversed", "non-integer", "past-the-text", "not-an-object"], +) +def test_glirel_rejects_bad_entities_as_invalid_input_before_any_model_work(entity: Any) -> None: + adapter, predict = glirel_adapter() + good = glirel_item([{"text": "Tim Cook", "label": "PERSON", "start": 0, "end": 8}]) + + with pytest.raises(InvalidInputError): + adapter.extract([good, glirel_item([entity])], labels=["ceo of"]) + predict.assert_not_called() + + +def test_glirel_rejects_missing_input_as_invalid_input() -> None: + adapter, _ = glirel_adapter() + + with pytest.raises(InvalidInputError, match="requires labels"): + adapter.extract([glirel_item([{"text": "Tim", "label": "P", "start": 0, "end": 3}])], labels=[]) + with pytest.raises(InvalidInputError, match="requires entities"): + adapter.extract([Item(text="Tim Cook leads Apple")], labels=["ceo of"]) + with pytest.raises(InvalidInputError, match="must be a list"): + adapter.extract([Item(text="Tim Cook", metadata={"entities": "Tim"})], labels=["ceo of"]) + with pytest.raises(InvalidInputError, match="labels may have at most"): + adapter.extract([glirel_item([{"text": "Tim", "label": "P", "start": 0, "end": 3}])], labels=["r" * 200])