Compare commits

...
2 Commits
Author SHA1 Message Date
shamoon 71af44be2c one-liner 2026-08-16 20:15:53 -07:00
shamoon 38df497760 Fix: DocumentClassifierSchema bounds 2026-08-16 20:15:52 -07:00
2 changed files with 128 additions and 4 deletions
+44 -4
View File
@@ -1,7 +1,34 @@
from typing import Any
from typing import Final
from typing import TypedDict from typing import TypedDict
from pydantic import BaseModel from pydantic import BaseModel
from pydantic import Field from pydantic import Field
from pydantic import ValidationInfo
from pydantic import field_validator
from pydantic.fields import FieldInfo
# taxonomy.py MAX_TAG_CANDIDATES = 10, prompt is "up to 3 relevant dates"
MAX_EXISTING_IDS: Final = 10
MAX_NEW_NAMES: Final = 8
MAX_DATES: Final = 3
# Matches documents.models.Document.title's CharField(max_length=128).
MAX_TITLE_LENGTH: Final = 128
def _truncate_to_field_limit(value: Any, field: FieldInfo) -> Any:
"""
Clip down to its it's declared maximum. Run as a `mode="before"` validator.
"""
limit = next(
(m.max_length for m in field.metadata if hasattr(m, "max_length")),
None,
)
return (
value
if (limit is None or not isinstance(value, (list, str)))
else value[:limit]
)
class TaxonomyChoice(BaseModel): class TaxonomyChoice(BaseModel):
@@ -14,19 +41,32 @@ class TaxonomyChoice(BaseModel):
TaxonomyChoiceDict below. TaxonomyChoiceDict below.
""" """
existing_ids: list[int] = Field(default_factory=list) existing_ids: list[int] = Field(
new_names: list[str] = Field(default_factory=list) default_factory=list,
max_length=MAX_EXISTING_IDS,
)
new_names: list[str] = Field(default_factory=list, max_length=MAX_NEW_NAMES)
@field_validator("existing_ids", "new_names", mode="before")
@classmethod
def _truncate(cls, value: Any, info: ValidationInfo) -> Any:
return _truncate_to_field_limit(value, cls.model_fields[info.field_name])
class DocumentClassifierSchema(BaseModel): class DocumentClassifierSchema(BaseModel):
"""Schema for document classification suggestions.""" """Schema for document classification suggestions."""
title: str title: str = Field(max_length=MAX_TITLE_LENGTH)
tags: TaxonomyChoice = Field(default_factory=TaxonomyChoice) tags: TaxonomyChoice = Field(default_factory=TaxonomyChoice)
correspondents: TaxonomyChoice = Field(default_factory=TaxonomyChoice) correspondents: TaxonomyChoice = Field(default_factory=TaxonomyChoice)
document_types: TaxonomyChoice = Field(default_factory=TaxonomyChoice) document_types: TaxonomyChoice = Field(default_factory=TaxonomyChoice)
storage_paths: TaxonomyChoice = Field(default_factory=TaxonomyChoice) storage_paths: TaxonomyChoice = Field(default_factory=TaxonomyChoice)
dates: list[str] = Field(default_factory=list) dates: list[str] = Field(default_factory=list, max_length=MAX_DATES)
@field_validator("title", "dates", mode="before")
@classmethod
def _truncate(cls, value: Any, info: ValidationInfo) -> Any:
return _truncate_to_field_limit(value, cls.model_fields[info.field_name])
class TaxonomyChoiceDict(TypedDict): class TaxonomyChoiceDict(TypedDict):
+84
View File
@@ -1,3 +1,7 @@
from paperless_ai.base_model import MAX_DATES
from paperless_ai.base_model import MAX_EXISTING_IDS
from paperless_ai.base_model import MAX_NEW_NAMES
from paperless_ai.base_model import MAX_TITLE_LENGTH
from paperless_ai.base_model import ClassificationSuggestions from paperless_ai.base_model import ClassificationSuggestions
from paperless_ai.base_model import DocumentClassifierSchema from paperless_ai.base_model import DocumentClassifierSchema
from paperless_ai.base_model import TaxonomyChoice from paperless_ai.base_model import TaxonomyChoice
@@ -63,6 +67,86 @@ def test_document_classifier_schema_json_schema_is_self_contained():
assert set(taxonomy_choice_properties.keys()) == {"existing_ids", "new_names"} assert set(taxonomy_choice_properties.keys()) == {"existing_ids", "new_names"}
def test_every_sequence_in_the_emitted_schema_is_bounded():
"""
GIVEN:
- The DocumentClassifierSchema pydantic model
WHEN:
- Its JSON schema is generated via model_json_schema()
THEN:
- Every array property in the schema, including those on the
referenced TaxonomyChoice definition, carries a maxItems
"""
schema = DocumentClassifierSchema.model_json_schema()
unbounded = [
f"{owner}.{name}"
for owner, definition in [
("DocumentClassifierSchema", schema),
*schema.get("$defs", {}).items(),
]
for name, prop in definition.get("properties", {}).items()
if prop.get("type") == "array" and "maxItems" not in prop
]
assert unbounded == []
def test_dates_bound_matches_what_the_prompt_asks_for():
"""
GIVEN:
- The DocumentClassifierSchema pydantic model
WHEN:
- The emitted maxItems for dates is inspected
THEN:
- It equals the 3 that build_prompt_without_rag asks the model for
"""
dates_schema = DocumentClassifierSchema.model_json_schema()["properties"]["dates"]
assert dates_schema["maxItems"] == MAX_DATES == 3
def test_over_long_response_is_truncated_rather_than_rejected():
"""
GIVEN:
- An LLM response overshooting every declared bound
WHEN:
- DocumentClassifierSchema is constructed from it
THEN:
- Each field is clipped to its maximum, with no ValidationError
"""
parsed = DocumentClassifierSchema(
title="T" * (MAX_TITLE_LENGTH + 50),
tags=TaxonomyChoice(
existing_ids=list(range(MAX_EXISTING_IDS + 20)),
new_names=["n"] * (MAX_NEW_NAMES + 20),
),
dates=[f"2016-{month:02d}-01" for month in range(1, 13)],
)
assert len(parsed.title) == MAX_TITLE_LENGTH
assert len(parsed.dates) == MAX_DATES
assert len(parsed.tags.existing_ids) == MAX_EXISTING_IDS
assert len(parsed.tags.new_names) == MAX_NEW_NAMES
def test_truncation_keeps_the_earliest_entries():
"""
GIVEN:
- An over-long dates list from an LLM response
WHEN:
- DocumentClassifierSchema is constructed from it
THEN:
- The kept entries are the first ones the model emitted
"""
parsed = DocumentClassifierSchema(
title="T",
dates=["2016-10-01", "2016-09-01", "2016-08-01", "2016-07-01", "2016-06-01"],
)
assert parsed.dates == ["2016-10-01", "2016-09-01", "2016-08-01"]
def test_model_dump_matches_typed_dict_keys(): def test_model_dump_matches_typed_dict_keys():
""" """
GIVEN: GIVEN: