Skip to content

Commit 5aa0287

Browse files
svonavaclaude
andauthored
fix(server): limit the label prompt of GLiNER-family extract requests (#381)
GLiNER, GLiNER2 and GLiREL now validate the label prompt before any model work. Each label, relation type, class label, task name, schema field and choice is limited to 128 characters, and the whole prompt to a token limit counted with the model's tokenizer: 1024 for GLiNER and GLiREL, 2048 for GLiNER2. Requests over a limit get 400 INVALID_INPUT. The limit is configurable per model with max_prompt_tokens. GLiREL accepts at most 256 supplied entities per item, and rejects malformed input with INVALID_INPUT before any item runs. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
1 parent 894a017 commit 5aa0287

6 files changed

Lines changed: 628 additions & 29 deletions

File tree

Lines changed: 131 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,131 @@
1+
"""Limits on the task prompt GLiNER-family adapters encode with each document.
2+
3+
GLiNER, GLiNER2 and GLiREL encode a request's labels, relation types, class
4+
labels and schema fields (its task prompt) with every document. The prompt is
5+
not billed (input tokens count the document), so it is bounded instead, as
6+
GLiFormer and GLiNER2.5-Decide bound theirs. A request is rejected with
7+
``InvalidInputError`` (HTTP 400 ``INVALID_INPUT``) when
8+
9+
* a label, relation type, class label, task name, field name or choice has
10+
more than ``MAX_LABEL_CHARS`` characters (checked first, before anything is
11+
tokenized), or
12+
* the prompt takes more than ``max_prompt_tokens`` tokens, counted with the
13+
model's tokenizer.
14+
15+
The defaults sit far above the label sets these models are used with: the
16+
largest label set among this repository's examples takes about 40 tokens, a
17+
60-type PII list about 230, and a structured-extraction schema of 50 described
18+
fields about 1,200, while ``DEFAULT_MAX_PROMPT_TOKENS`` holds several hundred
19+
entity types and ``DEFAULT_MAX_SCHEMA_PROMPT_TOKENS`` a schema of dozens of
20+
described fields with their choices.
21+
22+
Descriptions are not limited one by one; the whole prompt is tokenized only
23+
after its characters are checked, so counting costs at most
24+
``MAX_PROMPT_CHARS_PER_TOKEN`` characters of tokenization per allowed token.
25+
"""
26+
27+
from __future__ import annotations
28+
29+
import hashlib
30+
from collections import OrderedDict
31+
from collections.abc import Callable, Hashable, Iterable
32+
from typing import Any
33+
34+
from sie_server.types.inputs import InvalidInputError
35+
36+
# Characters a label, relation type, class label, task name, field name or choice may have.
37+
MAX_LABEL_CHARS = 128
38+
# Tokens a request's labels, relation types and class labels may take.
39+
DEFAULT_MAX_PROMPT_TOKENS = 1024
40+
# Tokens a GLiNER2 request's labels or schema (field names, descriptions and choices) may take.
41+
DEFAULT_MAX_SCHEMA_PROMPT_TOKENS = 2048
42+
# No token of these tokenizers covers more characters than this.
43+
MAX_PROMPT_CHARS_PER_TOKEN = 32
44+
_PROMPT_CACHE_SIZE = 256
45+
46+
47+
def validate_max_prompt_tokens(value: object) -> int:
48+
"""A ``max_prompt_tokens`` adapter option, checked.
49+
50+
Raises:
51+
ValueError: The value is not a positive integer.
52+
"""
53+
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
54+
raise ValueError("max_prompt_tokens must be a positive integer")
55+
return value
56+
57+
58+
def check_label_chars(model: str, kind: str, values: Iterable[str]) -> None:
59+
"""Reject a label (or relation type, class label, field name, choice) of more than ``MAX_LABEL_CHARS`` characters.
60+
61+
Raises:
62+
InvalidInputError: A value is not a string or is too long.
63+
"""
64+
for value in values:
65+
if not isinstance(value, str):
66+
raise InvalidInputError(f"{model} {kind} must be strings")
67+
if len(value) > MAX_LABEL_CHARS:
68+
raise InvalidInputError(f"{model} {kind} may have at most {MAX_LABEL_CHARS} characters each")
69+
70+
71+
class PromptLimit:
72+
"""Checks a request's task prompt against ``max_tokens``, remembering recent prompts' sizes."""
73+
74+
__slots__ = ("_counts", "max_tokens", "model")
75+
76+
def __init__(self, model: str, max_tokens: int) -> None:
77+
self.model = model
78+
self.max_tokens = validate_max_prompt_tokens(max_tokens)
79+
self._counts: OrderedDict[bytes, int] = OrderedDict()
80+
81+
def check(self, texts: Iterable[str], count: Callable[[], int], key: Hashable) -> int:
82+
"""The prompt's tokens, from ``count()``, after checking ``texts`` (its strings) and the result.
83+
84+
``key`` identifies the prompt (its strings, in order, and anything else
85+
its size depends on); a prompt seen recently is not counted again.
86+
87+
Raises:
88+
InvalidInputError: The prompt takes more than ``max_tokens`` tokens.
89+
"""
90+
digest = hashlib.sha256(repr(key).encode("utf-8", "surrogatepass")).digest()
91+
tokens = self._counts.get(digest)
92+
if tokens is None:
93+
chars = sum(len(text) for text in texts)
94+
if chars > self.max_tokens * MAX_PROMPT_CHARS_PER_TOKEN:
95+
raise InvalidInputError(self._message(None, chars))
96+
tokens = int(count())
97+
self._counts[digest] = tokens
98+
if len(self._counts) > _PROMPT_CACHE_SIZE:
99+
self._counts.popitem(last=False)
100+
else:
101+
self._counts.move_to_end(digest)
102+
if tokens > self.max_tokens:
103+
raise InvalidInputError(self._message(tokens, None))
104+
return tokens
105+
106+
def _message(self, tokens: int | None, chars: int | None) -> str:
107+
size = f"{tokens} tokens" if tokens is not None else f"{chars} characters"
108+
return (
109+
f"{self.model} labels, relation types, class labels and schema fields take {size}; "
110+
f"a request may use at most {self.max_tokens} tokens for them"
111+
)
112+
113+
114+
def gliner_prompt_counter(model: Any) -> Callable[[list[str], list[str]], int]:
115+
"""Tokens of the prompt a loaded ``gliner`` model builds for entity and relation types.
116+
117+
The prompt is built by the model's own processor (``prepare_inputs``, with
118+
no document words) and tokenized as the processor tokenizes it.
119+
"""
120+
processor = model.data_processor
121+
tokenizer = processor.transformer_tokenizer
122+
123+
def count(entity_types: list[str], relation_types: list[str]) -> int:
124+
kwargs = {"relations": relation_types} if relation_types else {}
125+
(words,), _ = processor.prepare_inputs([[]], entity_types, **kwargs)
126+
if not words:
127+
return 0
128+
encoding = tokenizer(list(words), is_split_into_words=True, add_special_tokens=False)
129+
return len(encoding["input_ids"])
130+
131+
return count

‎packages/sie_server/src/sie_server/adapters/gliner/__init__.py‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,11 @@
1717
Joint entity-relation ("relex") models also extract relations between the
1818
entities they find when a request names relation types in
1919
``options["relation_labels"]``. Without it they return entities only.
20+
21+
A request's labels and relation types, which GLiNER encodes with every
22+
document and does not bill, may have at most 128 characters each and take at
23+
most ``max_prompt_tokens`` tokens together (default 1024); a longer prompt is
24+
rejected with ``INVALID_INPUT``.
2025
"""
2126

2227
import math
@@ -27,6 +32,12 @@
2732
import torch
2833

2934
from sie_server.adapters._base_adapter import BaseAdapter
35+
from sie_server.adapters._prompt_limit import (
36+
DEFAULT_MAX_PROMPT_TOKENS,
37+
PromptLimit,
38+
check_label_chars,
39+
gliner_prompt_counter,
40+
)
3041
from sie_server.adapters._spec import AdapterSpec
3142
from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision
3243
from sie_server.adapters._word_window import bound_gliner_words, plan_forwards
@@ -96,6 +107,7 @@ def __init__(
96107
multi_label: bool = False,
97108
merge_adjacent_entities: bool = False,
98109
relation_threshold: float | None = None,
110+
max_prompt_tokens: int = DEFAULT_MAX_PROMPT_TOKENS,
99111
compute_precision: ComputePrecision = "float16",
100112
revision: str | None = None,
101113
**kwargs: Any, # Accept extra args from loader (e.g., pooling)
@@ -112,6 +124,9 @@ def __init__(
112124
relation_threshold: Minimum relation score (0-1) for joint
113125
entity-relation models. None uses the entity threshold, as the
114126
gliner library does.
127+
max_prompt_tokens: Most tokens a request's labels and relation
128+
types may take in the prompt encoded with each document (see
129+
``_prompt_limit``).
115130
compute_precision: Compute precision for inference.
116131
revision: Optional HuggingFace revision/branch/commit SHA to pin when
117132
loading model artifacts.
@@ -133,6 +148,9 @@ def __init__(
133148
self._extracts_relations = False
134149
# True when the encoder's attention memory grows with the square of a row (see ``_inference``).
135150
self._quadratic_attention = False
151+
self._prompt_limit = PromptLimit("GLiNER", max_prompt_tokens)
152+
# Tokens of the label prompt, as the loaded model builds it; None until loaded.
153+
self._count_prompt: Any = None
136154

137155
def load(self, device: str) -> None:
138156
"""Load the model onto the specified device.
@@ -173,6 +191,7 @@ def load(self, device: str) -> None:
173191
# gliner's max_len counts words, whatever their subwords: read at most a
174192
# bounded number of subwords too, with a long word in pieces.
175193
self._quadratic_attention = bound_gliner_words(self._model)
194+
self._count_prompt = gliner_prompt_counter(self._model)
176195

177196
def extract(
178197
self,
@@ -227,6 +246,8 @@ def extract(
227246
if relation_labels and not self._extracts_relations:
228247
raise InvalidInputError(_ERR_NO_RELATIONS)
229248

249+
self._check_prompt(labels, relation_labels)
250+
230251
# Extract texts from all items
231252
texts = [self._extract_text(item) for item in items]
232253
if any(not text.strip() for text in texts):
@@ -303,6 +324,27 @@ def extract(
303324

304325
return ExtractOutput(entities=all_entities, relations=all_relations, input_token_counts=input_token_counts)
305326

327+
def _check_prompt(self, labels: list[str], relation_labels: list[str]) -> None:
328+
"""Reject a request whose labels and relation types take more than ``max_prompt_tokens``.
329+
330+
gliner encodes the label prompt with every document, and only the
331+
document is billed.
332+
333+
Raises:
334+
InvalidInputError: The prompt is too long, or a label is not a string.
335+
"""
336+
check_label_chars("GLiNER", "labels", labels)
337+
check_label_chars("GLiNER", "relation_labels", relation_labels)
338+
entity_types = list(dict.fromkeys(labels)) # gliner drops repeated labels
339+
count = self._count_prompt
340+
341+
def tokens() -> int:
342+
return count(entity_types, relation_labels) if count is not None else 0
343+
344+
self._prompt_limit.check(
345+
[*entity_types, *relation_labels], tokens, (tuple(entity_types), tuple(relation_labels))
346+
)
347+
306348
@staticmethod
307349
def _validate_relation_labels(value: Any, labels: list[str]) -> list[str]:
308350
"""Return the requested relation types (empty when none were asked for)."""

‎packages/sie_server/src/sie_server/adapters/gliner2/adapter.py‎

Lines changed: 64 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from huggingface_hub import snapshot_download
1414

1515
from sie_server.adapters._base_adapter import BaseAdapter
16+
from sie_server.adapters._prompt_limit import DEFAULT_MAX_SCHEMA_PROMPT_TOKENS, PromptLimit, check_label_chars
1617
from sie_server.adapters._spec import AdapterSpec
1718
from sie_server.adapters._types import ERR_REQUIRES_TEXT, ComputePrecision
1819
from sie_server.adapters._word_window import (
@@ -76,6 +77,12 @@ class GLiNER2Adapter(BaseAdapter):
7677
- Batch methods cover entities, relations, structured data, and classification
7778
- Classification uses ``classify_text()`` / ``batch_classify_text()``
7879
80+
A request's labels, class labels, relation types or schema fields (field
81+
names, descriptions and choices), which gliner2 encodes with every document
82+
and does not bill, may take at most ``max_prompt_tokens`` tokens (default
83+
2048), and each label, task name, field name or choice at most 128
84+
characters; a longer prompt is rejected with ``INVALID_INPUT``.
85+
7986
Reference models:
8087
- fastino/gliner2-base-v1
8188
- fastino/gliner2-large-v1
@@ -98,6 +105,7 @@ def __init__(
98105
default_labels: list[str] | None = None,
99106
multi_label: bool = False,
100107
max_seq_length: int | None = None,
108+
max_prompt_tokens: int = DEFAULT_MAX_SCHEMA_PROMPT_TOKENS,
101109
compute_precision: ComputePrecision = "float16",
102110
revision: str | None = None,
103111
**kwargs: Any,
@@ -115,6 +123,9 @@ def __init__(
115123
multi_label: Whether the configured classification task may return
116124
multiple labels.
117125
max_seq_length: Maximum document and schema input length.
126+
max_prompt_tokens: Most tokens a request's labels, class labels,
127+
relation types or schema fields may take in the task prompt
128+
encoded with each document (see ``_prompt_limit``).
118129
compute_precision: Compute precision for inference.
119130
revision: Optional HuggingFace revision/branch/commit SHA to pin when
120131
loading model artifacts.
@@ -127,6 +138,7 @@ def __init__(
127138
self._default_labels = self._validate_labels(default_labels) if default_labels is not None else None
128139
self._multi_label = multi_label
129140
self._max_seq_length = max_seq_length
141+
self._prompt_limit = PromptLimit("GLiNER2", max_prompt_tokens)
130142
self._compute_precision = compute_precision
131143
self._revision = revision
132144

@@ -254,8 +266,14 @@ def extract(
254266
raise ValueError("GLiNER2 structured extraction does not accept classification_task")
255267
structures = self._json_schema_to_structures(output_schema)
256268
specs = [spec for fields in structures.values() for spec in fields]
257-
# A field's choices are read twice: in its structure and in a prefix before the document.
258-
rows = self._row_tokens(windows, specs + specs)
269+
# A field's choices are listed in its structure, and each again in a prefix before the document.
270+
choices = [
271+
choice for definition in output_schema["properties"].values() for choice in definition.get("enum") or []
272+
]
273+
check_label_chars("GLiNER2", "output_schema property names", output_schema["properties"])
274+
check_label_chars("GLiNER2", "output_schema enum values", choices)
275+
prompt = self._prompt_tokens(specs, key=("json", tuple(specs)), extra=len(choices))
276+
rows = self._row_tokens(windows, prompt)
259277
with torch.inference_mode():
260278
raw_results = self._run_planned(
261279
model_texts,
@@ -286,7 +304,11 @@ def extract(
286304
normalized_entities = [
287305
self._normalize_input_entities(item, entities or []) for item, entities in zip(items, relation_entities)
288306
]
289-
rows = self._row_tokens(windows, normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION)
307+
check_label_chars("GLiNER2", "labels", normalized_labels)
308+
prompt = self._prompt_tokens(
309+
normalized_labels, per_entry=_PROMPT_TOKENS_PER_RELATION, key=("relations", tuple(normalized_labels))
310+
)
311+
rows = self._row_tokens(windows, prompt)
290312
with torch.inference_mode():
291313
raw_results = self._run_planned(
292314
model_texts,
@@ -314,14 +336,20 @@ def extract(
314336
if classification_task is not None:
315337
if not isinstance(classification_task, str) or not classification_task.strip():
316338
raise ValueError("GLiNER2 classification_task must be a non-empty string")
339+
check_label_chars("GLiNER2", "classification_task", [classification_task])
340+
check_label_chars("GLiNER2", "labels", normalized_labels)
341+
prompt = self._prompt_tokens(
342+
[classification_task, *normalized_labels],
343+
key=("classification", classification_task, tuple(normalized_labels)),
344+
)
317345
return self._classify(
318346
model_texts,
319347
normalized_labels,
320348
task=classification_task,
321349
multi_label=multi_label,
322350
threshold=effective_threshold,
323351
input_token_counts=input_token_counts,
324-
rows=self._row_tokens(windows, [classification_task, *normalized_labels]),
352+
rows=self._row_tokens(windows, prompt),
325353
)
326354

327355
def extract_entities(batch: list[str]) -> list[Any]:
@@ -345,10 +373,12 @@ def extract_entities(batch: list[str]) -> list[Any]:
345373
max_len=self._max_seq_length,
346374
)
347375

376+
check_label_chars("GLiNER2", "labels", normalized_labels)
377+
prompt = self._prompt_tokens(normalized_labels, key=("entities", tuple(normalized_labels)))
348378
with torch.inference_mode():
349379
raw_results = self._run_planned(
350380
model_texts,
351-
self._row_tokens(windows, normalized_labels),
381+
self._row_tokens(windows, prompt),
352382
extract_entities,
353383
rows_per_pass=1 if len(texts) == 1 else _PACKAGE_BATCH_SIZE,
354384
)
@@ -412,22 +442,43 @@ def classify(batch: list[str]) -> list[Any]:
412442
input_token_counts=input_token_counts,
413443
)
414444

415-
def _row_tokens(
445+
def _prompt_tokens(
416446
self,
417-
windows: list[tuple[str, int | None]],
418-
prompt_entries: Iterable[str],
447+
entries: list[str],
419448
*,
449+
key: tuple[Any, ...],
420450
per_entry: int = _PROMPT_TOKENS_PER_ENTRY,
421-
) -> list[int] | None:
451+
extra: int = 0,
452+
) -> int | None:
453+
"""Estimated tokens of the task prompt gliner2 builds from ``entries``, checked against the limit.
454+
455+
Each label, class label, relation type and schema field is counted
456+
with the tokens gliner2 adds around it, plus ``extra`` (a token per
457+
field choice, which gliner2 lists again before the document). None when words are
458+
not counted (no bounded splitter is installed); the prompt's
459+
characters are still checked.
460+
461+
Raises:
462+
InvalidInputError: The prompt takes more than ``max_prompt_tokens``.
463+
"""
464+
count = self._count_subwords
465+
strings = [entry for entry in entries if isinstance(entry, str)]
466+
467+
def tokens() -> int:
468+
if count is None:
469+
return 0
470+
return sum(count(strings)) + per_entry * len(strings) + extra + _ROW_OVERHEAD_TOKENS
471+
472+
prompt = self._prompt_limit.check(strings, tokens, (per_entry, extra, *key))
473+
return prompt if count is not None else None
474+
475+
def _row_tokens(self, windows: list[tuple[str, int | None]], prompt: int | None) -> list[int] | None:
422476
"""Estimated tokens of each item's encoder row: the task prompt, then the words it reads.
423477
424478
None when the words were not counted (no bounded splitter is installed).
425479
"""
426-
count = self._count_subwords
427-
if count is None or any(subwords is None for _, subwords in windows):
480+
if prompt is None or any(subwords is None for _, subwords in windows):
428481
return None
429-
entries = [entry for entry in prompt_entries if isinstance(entry, str)]
430-
prompt = sum(count(entries)) + per_entry * len(entries) + _ROW_OVERHEAD_TOKENS
431482
return [prompt + (subwords or 0) for _, subwords in windows]
432483

433484
def _run_planned(

0 commit comments

Comments
 (0)