mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-08-17 16:23:18 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
71af44be2c | ||
|
|
38df497760 |
@@ -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):
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user