document classifier agent that runs as a policy and assigns a category and sub-category docType
This commit is contained in:
@@ -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)}"
|
||||
)
|
||||
@@ -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,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)
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user