document classifier agent that runs as a policy and assigns a category and sub-category docType

This commit is contained in:
EthanHealy01
2026-06-29 21:58:03 +01:00
parent 501a7199e0
commit fed7ad300f
37 changed files with 2309 additions and 100 deletions
+2
View File
@@ -1,5 +1,6 @@
"""Agent modules for Stirling AI reasoning flows."""
from .document_classifier import DocumentClassifierAgent
from .execution import ExecutionPlanningAgent
from .orchestrator import OrchestratorAgent
from .pdf_create import PdfCreateAgent
@@ -9,6 +10,7 @@ from .pdf_review import PdfReviewAgent
from .user_spec import UserSpecAgent
__all__ = [
"DocumentClassifierAgent",
"ExecutionPlanningAgent",
"OrchestratorAgent",
"PdfCreateAgent",
@@ -0,0 +1,568 @@
{
"categories": [
{
"id": "invoice",
"label": "Invoice",
"docTypes": [
{
"id": "invoice",
"label": "Invoice"
},
{
"id": "receipt",
"label": "Receipt"
},
{
"id": "credit_note",
"label": "Credit note"
},
{
"id": "purchase_order",
"label": "Purchase order"
},
{
"id": "quote",
"label": "Quote"
}
]
},
{
"id": "contract",
"label": "Contract",
"docTypes": [
{
"id": "nda",
"label": "Non-disclosure agreement"
},
{
"id": "employment_agreement",
"label": "Employment agreement"
},
{
"id": "service_agreement",
"label": "Service agreement"
},
{
"id": "lease_agreement",
"label": "Lease agreement"
},
{
"id": "master_service_agreement",
"label": "Master service agreement"
},
{
"id": "statement_of_work",
"label": "Statement of work"
},
{
"id": "terms_of_service",
"label": "Terms of service"
}
]
},
{
"id": "financial_statement",
"label": "Financial statement",
"docTypes": [
{
"id": "balance_sheet",
"label": "Balance sheet"
},
{
"id": "income_statement",
"label": "Income statement"
},
{
"id": "cash_flow_statement",
"label": "Cash flow statement"
},
{
"id": "bank_statement",
"label": "Bank statement"
},
{
"id": "annual_report",
"label": "Annual report"
}
]
},
{
"id": "report",
"label": "Report",
"docTypes": [
{
"id": "business_report",
"label": "Business report"
},
{
"id": "project_report",
"label": "Project report"
},
{
"id": "research_report",
"label": "Research report"
},
{
"id": "status_report",
"label": "Status report"
},
{
"id": "incident_report",
"label": "Incident report"
}
]
},
{
"id": "letter",
"label": "Letter",
"docTypes": [
{
"id": "business_letter",
"label": "Business letter"
},
{
"id": "cover_letter",
"label": "Cover letter"
},
{
"id": "recommendation_letter",
"label": "Recommendation letter"
},
{
"id": "complaint_letter",
"label": "Complaint letter"
},
{
"id": "demand_letter",
"label": "Demand letter"
}
]
},
{
"id": "form",
"label": "Form",
"docTypes": [
{
"id": "application_form",
"label": "Application form"
},
{
"id": "registration_form",
"label": "Registration form"
},
{
"id": "consent_form",
"label": "Consent form"
},
{
"id": "survey",
"label": "Survey"
},
{
"id": "questionnaire",
"label": "Questionnaire"
}
]
},
{
"id": "resume",
"label": "Resume",
"docTypes": [
{
"id": "resume",
"label": "Resume"
},
{
"id": "curriculum_vitae",
"label": "Curriculum vitae"
},
{
"id": "portfolio",
"label": "Portfolio"
},
{
"id": "reference_sheet",
"label": "Reference sheet"
}
]
},
{
"id": "tax_form",
"label": "Tax form",
"docTypes": [
{
"id": "tax_return",
"label": "Tax return"
},
{
"id": "w2",
"label": "W-2"
},
{
"id": "w9",
"label": "W-9"
},
{
"id": "form_1099",
"label": "Form 1099"
},
{
"id": "vat_return",
"label": "VAT return"
}
]
},
{
"id": "expense_report",
"label": "Expense report",
"docTypes": [
{
"id": "expense_report",
"label": "Expense report"
},
{
"id": "reimbursement_request",
"label": "Reimbursement request"
},
{
"id": "mileage_log",
"label": "Mileage log"
},
{
"id": "per_diem_claim",
"label": "Per diem claim"
}
]
},
{
"id": "presentation",
"label": "Presentation",
"docTypes": [
{
"id": "slide_deck",
"label": "Slide deck"
},
{
"id": "pitch_deck",
"label": "Pitch deck"
},
{
"id": "training_deck",
"label": "Training deck"
},
{
"id": "webinar_deck",
"label": "Webinar deck"
}
]
},
{
"id": "medical_record",
"label": "Medical record",
"docTypes": [
{
"id": "lab_result",
"label": "Lab result"
},
{
"id": "prescription",
"label": "Prescription"
},
{
"id": "discharge_summary",
"label": "Discharge summary"
},
{
"id": "medical_history",
"label": "Medical history"
},
{
"id": "imaging_report",
"label": "Imaging report"
},
{
"id": "vaccination_record",
"label": "Vaccination record"
}
]
},
{
"id": "legal_filing",
"label": "Legal filing",
"docTypes": [
{
"id": "court_filing",
"label": "Court filing"
},
{
"id": "complaint",
"label": "Complaint"
},
{
"id": "motion",
"label": "Motion"
},
{
"id": "subpoena",
"label": "Subpoena"
},
{
"id": "affidavit",
"label": "Affidavit"
},
{
"id": "deposition",
"label": "Deposition"
}
]
},
{
"id": "identity_document",
"label": "Identity document",
"docTypes": [
{
"id": "passport",
"label": "Passport"
},
{
"id": "drivers_license",
"label": "Driver's license"
},
{
"id": "national_id",
"label": "National ID"
},
{
"id": "birth_certificate",
"label": "Birth certificate"
},
{
"id": "visa",
"label": "Visa"
}
]
},
{
"id": "insurance",
"label": "Insurance",
"docTypes": [
{
"id": "insurance_policy",
"label": "Insurance policy"
},
{
"id": "insurance_claim",
"label": "Insurance claim"
},
{
"id": "certificate_of_insurance",
"label": "Certificate of insurance"
},
{
"id": "explanation_of_benefits",
"label": "Explanation of benefits"
}
]
},
{
"id": "real_estate",
"label": "Real estate",
"docTypes": [
{
"id": "deed",
"label": "Deed"
},
{
"id": "mortgage_agreement",
"label": "Mortgage agreement"
},
{
"id": "property_appraisal",
"label": "Property appraisal"
},
{
"id": "closing_disclosure",
"label": "Closing disclosure"
},
{
"id": "title_report",
"label": "Title report"
}
]
},
{
"id": "shipping",
"label": "Shipping",
"docTypes": [
{
"id": "bill_of_lading",
"label": "Bill of lading"
},
{
"id": "packing_slip",
"label": "Packing slip"
},
{
"id": "customs_declaration",
"label": "Customs declaration"
},
{
"id": "delivery_note",
"label": "Delivery note"
},
{
"id": "air_waybill",
"label": "Air waybill"
}
]
},
{
"id": "hr_document",
"label": "HR document",
"docTypes": [
{
"id": "offer_letter",
"label": "Offer letter"
},
{
"id": "performance_review",
"label": "Performance review"
},
{
"id": "payslip",
"label": "Payslip"
},
{
"id": "employee_handbook",
"label": "Employee handbook"
},
{
"id": "termination_letter",
"label": "Termination letter"
},
{
"id": "timesheet",
"label": "Timesheet"
}
]
},
{
"id": "academic_record",
"label": "Academic record",
"docTypes": [
{
"id": "transcript",
"label": "Transcript"
},
{
"id": "diploma",
"label": "Diploma"
},
{
"id": "certificate",
"label": "Certificate"
},
{
"id": "syllabus",
"label": "Syllabus"
},
{
"id": "thesis",
"label": "Thesis"
},
{
"id": "report_card",
"label": "Report card"
}
]
},
{
"id": "marketing_material",
"label": "Marketing material",
"docTypes": [
{
"id": "brochure",
"label": "Brochure"
},
{
"id": "flyer",
"label": "Flyer"
},
{
"id": "case_study",
"label": "Case study"
},
{
"id": "white_paper",
"label": "White paper"
},
{
"id": "press_release",
"label": "Press release"
}
]
},
{
"id": "technical_document",
"label": "Technical document",
"docTypes": [
{
"id": "user_manual",
"label": "User manual"
},
{
"id": "specification",
"label": "Specification"
},
{
"id": "api_documentation",
"label": "API documentation"
},
{
"id": "installation_guide",
"label": "Installation guide"
},
{
"id": "datasheet",
"label": "Datasheet"
},
{
"id": "release_notes",
"label": "Release notes"
}
]
}
],
"tags": [
"finance",
"legal",
"medical",
"hr",
"tax",
"insurance",
"marketing",
"technical",
"operations",
"academic",
"government",
"draft",
"final",
"signed",
"unsigned",
"executed",
"expired",
"amended",
"void",
"confidential",
"internal",
"public",
"pii",
"phi",
"certified",
"notarized",
"scanned",
"redacted",
"template",
"urgent"
]
}
@@ -0,0 +1,179 @@
from __future__ import annotations
import logging
from pathlib import Path
from pydantic import Field
from pydantic_ai import Agent
from pydantic_ai.output import NativeOutput
from stirling.contracts import (
ClassificationTaxonomy,
ClassifyDocumentRequest,
ClassifyDocumentResponse,
DocumentClassificationResponse,
PageText,
)
from stirling.models import ApiModel
from stirling.services import AppRuntime
logger = logging.getLogger(__name__)
# Sentinel for a label that fell outside the supplied vocabulary.
UNKNOWN_LABEL = "unknown"
# An off-list answer can never be reported as more confident than this, so a
# confident-but-wrong model answer can't clear an organisation's accept
# threshold downstream. See the design doc's "Validate" step.
UNKNOWN_MAX_CONFIDENCE = 0.2
# Pages read from each end of the document. A document's type is evident from
# its opening (and closing) pages, so a fixed window keeps cost and latency flat
# regardless of length. Promote to AppSettings if it ever needs tuning.
WINDOW_PAGES = 2
# The built-in vocabulary the classifier falls back to when a request doesn't
# supply its own. GENERATED from the TS source of truth
# (frontend/editor/src/proprietary/data/classificationTaxonomy.ts) via
# `task frontend:taxonomy` — edit that file, not this JSON. Validated into the
# typed contract on import, so a malformed entry fails fast.
_DEFAULT_TAXONOMY_PATH = Path(__file__).with_name("default_taxonomy.generated.json")
DEFAULT_TAXONOMY = ClassificationTaxonomy.model_validate_json(_DEFAULT_TAXONOMY_PATH.read_text(encoding="utf-8"))
_SYSTEM_PROMPT = (
"You identify what a document is, choosing only from a fixed vocabulary you "
"are given. Decide along three axes:\n"
"- category: the document's structural family. Choose EXACTLY ONE category id.\n"
"- doc_type: the specific instrument within that family. Choose EXACTLY ONE "
"doc_type id listed under the category you chose.\n"
"- tags: zero or more descriptor ids from the tag list.\n"
"\n"
"Rules:\n"
"- Use only ids from the supplied vocabulary. If nothing fits, return "
f'"{UNKNOWN_LABEL}" for category and/or doc_type.\n'
"- The doc_type you pick must belong to the category you pick.\n"
"- type_confidence (0.0-1.0) is how sure you are that doc_type is correct.\n"
"- Judge from the document's content and structure, not from keywords alone. "
"The document may be in any language.\n"
"- You are shown only the first and last pages; that is enough to identify the type."
)
class _ClassifierOutput(ApiModel):
"""Raw model answer, before it is validated against the taxonomy."""
category: str = Field(description="A category id from the vocabulary, or 'unknown'.")
doc_type: str = Field(description="A doc_type id belonging to the chosen category, or 'unknown'.")
type_confidence: float = Field(ge=0.0, le=1.0, description="Confidence that doc_type is correct.")
tags: list[str] = Field(default_factory=list, description="Descriptor ids drawn from the tag list.")
def render_taxonomy(taxonomy: ClassificationTaxonomy) -> str:
"""Render the vocabulary for the prompt, ids first so the model echoes them."""
lines = ["Categories (id (label): doc_types as id (label)):"]
for category in taxonomy.categories:
types = ", ".join(f"{doc_type.id} ({doc_type.label})" for doc_type in category.doc_types) or "(none)"
lines.append(f"- {category.id} ({category.label}): {types}")
lines.append(f"Tags: {', '.join(taxonomy.tags) or '(none)'}")
return "\n".join(lines)
def select_window(pages: list[PageText], window: int = WINDOW_PAGES) -> list[PageText]:
"""Return the first and last ``window`` pages, never overlapping.
Documents short enough that the two ends would meet are returned whole. The
caller usually sends just the window already; this is a defensive trim in
case it sends more.
"""
if window <= 0 or len(pages) <= window * 2:
return list(pages)
return [*pages[:window], *pages[-window:]]
def format_window(pages: list[PageText]) -> str:
if not pages:
return "(no extractable text)"
return "\n\n".join(f"[Page {page.page_number}]\n{page.text}" for page in pages)
def validate_against_taxonomy(
output: _ClassifierOutput,
taxonomy: ClassificationTaxonomy,
) -> DocumentClassificationResponse:
"""Coerce a raw model answer onto the supplied vocabulary.
An off-list category collapses both axes to ``unknown``; a doc_type that
isn't a child of its (valid) category collapses the type alone. Either
collapse caps confidence. Tags are filtered to the known set, de-duplicated,
and returned in the model's order. The model identifies; these rules decide
what is allowed to stand.
"""
categories_by_id = {category.id.lower(): category for category in taxonomy.categories}
allowed_tags = {tag.lower(): tag for tag in taxonomy.tags}
kept_tags: list[str] = []
for tag in output.tags:
canonical = allowed_tags.get(tag.strip().lower())
if canonical is not None and canonical not in kept_tags:
kept_tags.append(canonical)
category = categories_by_id.get(output.category.strip().lower())
if category is None:
return DocumentClassificationResponse(
category=UNKNOWN_LABEL,
doc_type=UNKNOWN_LABEL,
type_confidence=min(output.type_confidence, UNKNOWN_MAX_CONFIDENCE),
tags=kept_tags,
)
types_by_id = {doc_type.id.lower(): doc_type for doc_type in category.doc_types}
doc_type = types_by_id.get(output.doc_type.strip().lower())
if doc_type is None:
return DocumentClassificationResponse(
category=category.id,
doc_type=UNKNOWN_LABEL,
type_confidence=min(output.type_confidence, UNKNOWN_MAX_CONFIDENCE),
tags=kept_tags,
)
return DocumentClassificationResponse(
category=category.id,
doc_type=doc_type.id,
type_confidence=output.type_confidence,
tags=kept_tags,
)
class DocumentClassifierAgent:
"""Identifies a document's category, type, and tags against a taxonomy.
Reads the bounded page window supplied on the request (first/last
``WINDOW_PAGES``) and runs a single fast-model pass, then validates the
answer against the vocabulary so nothing off-list survives.
"""
def __init__(self, runtime: AppRuntime) -> None:
self.runtime = runtime
self._agent = Agent(
model=runtime.fast_model,
output_type=NativeOutput(_ClassifierOutput),
system_prompt=_SYSTEM_PROMPT,
model_settings=runtime.fast_model_settings,
)
async def classify(self, request: ClassifyDocumentRequest) -> ClassifyDocumentResponse:
# Override point: a request-supplied taxonomy (e.g. a future per-org / DB
# vocabulary the backend resolves) wins; the generated default is the fallback.
taxonomy = request.taxonomy or DEFAULT_TAXONOMY
window = select_window(request.pages)
prompt = self._build_prompt(request.file_name, taxonomy, window)
logger.debug("[classify] prompt:\n%s", prompt)
result = await self._agent.run(prompt)
return validate_against_taxonomy(result.output, taxonomy)
@staticmethod
def _build_prompt(file_name: str, taxonomy: ClassificationTaxonomy, window: list[PageText]) -> str:
return (
f"{render_taxonomy(taxonomy)}\n\n"
f"Document file name: {file_name}\n"
f"Document content (first and last pages):\n{format_window(window)}"
)
+4
View File
@@ -10,6 +10,7 @@ from pydantic_ai import Agent
from pydantic_ai.models.instrumented import InstrumentationSettings
from stirling.agents import (
DocumentClassifierAgent,
ExecutionPlanningAgent,
OrchestratorAgent,
PdfEditAgent,
@@ -24,6 +25,7 @@ from stirling.api.middleware import UserIdMiddleware
from stirling.api.routes import (
agent_capabilities_router,
agent_draft_router,
document_classifier_router,
document_router,
execution_router,
ledger_router,
@@ -95,6 +97,7 @@ async def lifespan(fast_api: FastAPI):
fast_api.state.execution_planning_agent = ExecutionPlanningAgent(runtime)
fast_api.state.math_auditor_agent = MathAuditorAgent(runtime)
fast_api.state.pdf_comment_agent = PdfCommentAgent(runtime)
fast_api.state.document_classifier_agent = DocumentClassifierAgent(runtime)
tracer_provider = setup_posthog_tracking(settings)
if tracer_provider:
Agent.instrument_all(InstrumentationSettings(tracer_provider=tracer_provider))
@@ -131,6 +134,7 @@ app.include_router(document_router, dependencies=_user_gate)
app.include_router(ledger_router, dependencies=_user_gate)
app.include_router(pdf_comments_router, dependencies=_user_gate)
app.include_router(agent_capabilities_router, dependencies=_user_gate)
app.include_router(document_classifier_router, dependencies=_user_gate)
@app.get("/health", response_model=HealthResponse)
+5
View File
@@ -5,6 +5,7 @@ from typing import Annotated
from fastapi import Depends, HTTPException, Request, status
from stirling.agents import (
DocumentClassifierAgent,
ExecutionPlanningAgent,
OrchestratorAgent,
PdfEditAgent,
@@ -55,6 +56,10 @@ def get_pdf_comment_agent(request: Request) -> PdfCommentAgent:
return request.app.state.pdf_comment_agent
def get_document_classifier_agent(request: Request) -> DocumentClassifierAgent:
return request.app.state.document_classifier_agent
def require_user_id() -> UserId:
"""FastAPI dependency for routes that touch per-user storage.
@@ -1,5 +1,6 @@
from .agent_capabilities import router as agent_capabilities_router
from .agent_drafts import router as agent_draft_router
from .document_classifier import router as document_classifier_router
from .documents import router as document_router
from .execution import router as execution_router
from .ledger import router as ledger_router
@@ -11,6 +12,7 @@ from .pdf_questions import router as pdf_question_router
__all__ = [
"agent_capabilities_router",
"agent_draft_router",
"document_classifier_router",
"document_router",
"execution_router",
"ledger_router",
@@ -0,0 +1,24 @@
from __future__ import annotations
from typing import Annotated
from fastapi import APIRouter, Depends
from stirling.agents import DocumentClassifierAgent
from stirling.api.dependencies import get_document_classifier_agent
from stirling.contracts import ClassifyDocumentRequest, ClassifyDocumentResponse
router = APIRouter(prefix="/api/v1/documents/classify", tags=["document-classifier"])
@router.post("", response_model=ClassifyDocumentResponse)
async def classify_document(
request: ClassifyDocumentRequest,
agent: Annotated[DocumentClassifierAgent, Depends(get_document_classifier_agent)],
) -> ClassifyDocumentResponse:
"""Classify a document from its supplied page text against the default taxonomy.
The caller sends the bounded page window inline, so no per-user document
storage is touched here — the request is self-contained.
"""
return await agent.classify(request)
+14
View File
@@ -36,6 +36,14 @@ from .contradiction import (
ContradictionReport,
ContradictionSeverity,
)
from .document_classifier import (
ClassificationTaxonomy,
ClassifyDocumentRequest,
ClassifyDocumentResponse,
DocumentCategory,
DocumentClassificationResponse,
DocumentType,
)
from .documents import (
DeleteDocumentResponse,
IngestDocumentRequest,
@@ -129,6 +137,9 @@ __all__ = [
"AiToolAgentStep",
"ArtifactKind",
"CannotContinueExecutionAction",
"ClassificationTaxonomy",
"ClassifyDocumentRequest",
"ClassifyDocumentResponse",
"Claim",
"CommentSpec",
"CompletedExecutionAction",
@@ -139,8 +150,11 @@ __all__ = [
"DeleteDocumentResponse",
"PurgeOwnerResponse",
"Discrepancy",
"DocumentCategory",
"DocumentClassificationResponse",
"DocumentMeta",
"DocumentSections",
"DocumentType",
"DiscrepancyKind",
"EditCannotDoResponse",
"EditClarificationRequest",
+1
View File
@@ -63,6 +63,7 @@ class WorkflowOutcome(StrEnum):
UNSUPPORTED_CAPABILITY = "unsupported_capability"
GENERATE_FILE = "generate_file"
CONVERT_MARKDOWN = "convert_markdown"
CLASSIFICATION = "classification"
class ArtifactKind(StrEnum):
@@ -0,0 +1,70 @@
from __future__ import annotations
from typing import Literal
from pydantic import Field
from stirling.models import ApiModel
from .common import WorkflowOutcome
from .documents import PageText
class DocumentType(ApiModel):
"""A specific instrument within a category (e.g. ``nda`` inside ``contract``)."""
id: str = Field(min_length=1)
label: str = Field(min_length=1)
class DocumentCategory(ApiModel):
"""A structural family of documents, owning the doc_types shaped like it."""
id: str = Field(min_length=1)
label: str = Field(min_length=1)
doc_types: list[DocumentType] = Field(default_factory=list)
class ClassificationTaxonomy(ApiModel):
"""The vocabulary a document is classified against.
Supplied per request by the backend. When omitted, the engine falls back to
its small built-in default (see ``DEFAULT_TAXONOMY``). Tags are free-standing
descriptors that never own doc_types.
"""
categories: list[DocumentCategory] = Field(min_length=1)
tags: list[str] = Field(default_factory=list)
class ClassifyDocumentRequest(ApiModel):
"""Classify one document from its page text.
The caller sends the page text directly — typically just the bounded window
(first/last pages), since the classifier reads no more than that. There is no
ingestion or RAG step.
"""
file_name: str = Field(min_length=1)
pages: list[PageText] = Field(default_factory=list)
taxonomy: ClassificationTaxonomy | None = None
class DocumentClassificationResponse(ApiModel):
"""Terminal classification result.
``category`` and ``doc_type`` are ids drawn from the taxonomy, or the
sentinel ``"unknown"`` when the model's answer fell outside it. ``tags`` are
the subset of the model's tags that exist in the taxonomy.
"""
outcome: Literal[WorkflowOutcome.CLASSIFICATION] = WorkflowOutcome.CLASSIFICATION
category: str
doc_type: str
type_confidence: float = Field(ge=0.0, le=1.0)
tags: list[str] = Field(default_factory=list)
# Only one response shape today; kept as a named alias so routes and agents have
# a stable response type to import.
ClassifyDocumentResponse = DocumentClassificationResponse
+177
View File
@@ -0,0 +1,177 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from stirling.agents.document_classifier import (
DEFAULT_TAXONOMY,
UNKNOWN_LABEL,
UNKNOWN_MAX_CONFIDENCE,
DocumentClassifierAgent,
_ClassifierOutput,
render_taxonomy,
select_window,
validate_against_taxonomy,
)
from stirling.contracts import (
ClassificationTaxonomy,
ClassifyDocumentRequest,
DocumentCategory,
DocumentClassificationResponse,
PageText,
)
from stirling.services.runtime import AppRuntime
def _page(number: int, text: str = "x") -> PageText:
return PageText(page_number=number, text=text)
# ── select_window ───────────────────────────────────────────────────────────
def test_select_window_returns_short_documents_whole() -> None:
pages = [_page(1), _page(2), _page(3), _page(4)]
assert select_window(pages, window=2) == pages
def test_select_window_takes_both_ends_without_overlap() -> None:
pages = [_page(n) for n in range(1, 6)] # 5 pages
selected = select_window(pages, window=2)
assert [p.page_number for p in selected] == [1, 2, 4, 5]
def test_select_window_handles_empty() -> None:
assert select_window([], window=2) == []
def test_select_window_zero_returns_all() -> None:
pages = [_page(1), _page(2), _page(3)]
assert select_window(pages, window=0) == pages
# ── validate_against_taxonomy ────────────────────────────────────────────────
def test_valid_classification_is_preserved() -> None:
output = _ClassifierOutput(category="contract", doc_type="nda", type_confidence=0.95, tags=["legal", "signed"])
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert isinstance(result, DocumentClassificationResponse)
assert result.category == "contract"
assert result.doc_type == "nda"
assert result.type_confidence == 0.95
assert result.tags == ["legal", "signed"]
def test_off_list_category_collapses_to_unknown_with_capped_confidence() -> None:
output = _ClassifierOutput(category="spaceship", doc_type="warp_core", type_confidence=0.99)
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.category == UNKNOWN_LABEL
assert result.doc_type == UNKNOWN_LABEL
assert result.type_confidence == UNKNOWN_MAX_CONFIDENCE
def test_off_list_type_keeps_category_but_unknown_type() -> None:
output = _ClassifierOutput(category="contract", doc_type="invoice", type_confidence=0.9)
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.category == "contract"
assert result.doc_type == UNKNOWN_LABEL
assert result.type_confidence == UNKNOWN_MAX_CONFIDENCE
def test_type_from_a_different_category_is_not_a_child() -> None:
# "lab_result" is a valid type, but only under medical_record, not contract.
output = _ClassifierOutput(category="contract", doc_type="lab_result", type_confidence=0.8)
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.category == "contract"
assert result.doc_type == UNKNOWN_LABEL
def test_matching_is_case_insensitive_and_returns_canonical_ids() -> None:
output = _ClassifierOutput(category="Contract", doc_type="NDA", type_confidence=0.7, tags=["LEGAL"])
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.category == "contract"
assert result.doc_type == "nda"
assert result.tags == ["legal"]
def test_unknown_tags_dropped_and_deduplicated_in_order() -> None:
output = _ClassifierOutput(
category="invoice",
doc_type="invoice",
type_confidence=0.9,
tags=["finance", "made-up", "finance", "legal"],
)
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.tags == ["finance", "legal"]
def test_low_confidence_is_not_raised_when_collapsing() -> None:
output = _ClassifierOutput(category="nope", doc_type="nope", type_confidence=0.05)
result = validate_against_taxonomy(output, DEFAULT_TAXONOMY)
assert result.type_confidence == 0.05 # min(0.05, 0.2)
# ── render_taxonomy ──────────────────────────────────────────────────────────
def test_render_taxonomy_lists_ids_and_tags() -> None:
rendered = render_taxonomy(DEFAULT_TAXONOMY)
assert "contract" in rendered
assert "nda" in rendered
assert "finance" in rendered
def test_render_taxonomy_handles_category_without_types() -> None:
taxonomy = ClassificationTaxonomy(
categories=[DocumentCategory(id="memo", label="Memo", doc_types=[])],
tags=[],
)
rendered = render_taxonomy(taxonomy)
assert "(none)" in rendered
# ── DocumentClassifierAgent (inline page text) ───────────────────────────────
@pytest.mark.anyio
async def test_classify_validates_model_output_against_default_taxonomy(runtime: AppRuntime) -> None:
agent = DocumentClassifierAgent(runtime)
agent._agent.run = AsyncMock(
return_value=SimpleNamespace(
output=_ClassifierOutput(category="invoice", doc_type="invoice", type_confidence=0.97, tags=["finance"])
)
)
result = await agent.classify(
ClassifyDocumentRequest(
file_name="invoice.pdf",
pages=[PageText(page_number=1, text="Invoice INV-1 total due 100.00")],
)
)
assert isinstance(result, DocumentClassificationResponse)
assert result.category == "invoice"
assert result.doc_type == "invoice"
assert result.tags == ["finance"]
@pytest.mark.anyio
async def test_classify_collapses_off_list_model_answer(runtime: AppRuntime) -> None:
agent = DocumentClassifierAgent(runtime)
agent._agent.run = AsyncMock(
return_value=SimpleNamespace(
output=_ClassifierOutput(category="boarding_pass", doc_type="seat", type_confidence=0.9)
)
)
result = await agent.classify(
ClassifyDocumentRequest(file_name="weird.pdf", pages=[PageText(page_number=1, text="Some text")])
)
assert isinstance(result, DocumentClassificationResponse)
assert result.category == UNKNOWN_LABEL
assert result.doc_type == UNKNOWN_LABEL
assert result.type_confidence == UNKNOWN_MAX_CONFIDENCE
@@ -0,0 +1,67 @@
from __future__ import annotations
from collections.abc import Iterator
import pytest
from fastapi.testclient import TestClient
from stirling.api import app
from stirling.api.dependencies import get_document_classifier_agent
from stirling.contracts import (
ClassifyDocumentRequest,
ClassifyDocumentResponse,
DocumentClassificationResponse,
)
class StubClassifierAgent:
"""Stands in for DocumentClassifierAgent so route tests don't call a model."""
def __init__(self, response: ClassifyDocumentResponse) -> None:
self._response = response
async def classify(self, _request: ClassifyDocumentRequest) -> ClassifyDocumentResponse:
return self._response
@pytest.fixture
def classification_client() -> Iterator[TestClient]:
app.dependency_overrides[get_document_classifier_agent] = lambda: StubClassifierAgent(
DocumentClassificationResponse(
category="contract", doc_type="nda", type_confidence=0.96, tags=["legal", "signed"]
)
)
try:
yield TestClient(app)
finally:
app.dependency_overrides.pop(get_document_classifier_agent, None)
def test_classify_returns_camel_cased_result(classification_client: TestClient) -> None:
response = classification_client.post(
"/api/v1/documents/classify",
json={"fileName": "nda.pdf", "pages": [{"pageNumber": 1, "text": "Mutual NDA between A and B."}]},
)
assert response.status_code == 200
body = response.json()
assert body["outcome"] == "classification"
assert body["category"] == "contract"
assert body["docType"] == "nda"
assert body["typeConfidence"] == 0.96
assert body["tags"] == ["legal", "signed"]
def test_classify_accepts_empty_pages(classification_client: TestClient) -> None:
response = classification_client.post(
"/api/v1/documents/classify",
json={"fileName": "blank.pdf", "pages": []},
)
assert response.status_code == 200
def test_classify_rejects_empty_file_name(classification_client: TestClient) -> None:
response = classification_client.post(
"/api/v1/documents/classify",
json={"fileName": "", "pages": []},
)
assert response.status_code == 422