feat(classify): add rule-based document classification engine, task, and API

Implements the classify pipeline step with:
- Classification rules engine (app/utils/classification_rules.py) with
  pre-built categories (invoice, contract, receipt, letter, report,
  bank_statement, tax_document, insurance, payslip) and support for
  filename patterns, content keywords, and metadata matching rules
- Celery task (app/tasks/classify_document.py) that runs as a pipeline step
- CRUD API (app/api/classification_rules.py) for managing custom rules
- ClassificationRuleModel in app/models.py with migration 027
- Updated pipeline step config_schema and stage mapping
- Comprehensive tests for engine, API, and task

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-03-09 23:38:05 +00:00
parent e6dfa079cc
commit df051e8b81
13 changed files with 1930 additions and 5 deletions
+2
View File
@@ -11,6 +11,7 @@ from app.api.api_tokens import router as api_tokens_router
from app.api.azure import router as azure_router
from app.api.backup import router as backup_router
from app.api.billing import router as billing_router
from app.api.classification_rules import router as classification_rules_router
from app.api.database import router as database_router
from app.api.diagnostic import router as diagnostic_router
from app.api.dropbox import router as dropbox_router
@@ -82,3 +83,4 @@ router.include_router(imap_accounts_router)
router.include_router(integrations_router)
router.include_router(notifications_router)
router.include_router(scheduled_jobs_router)
router.include_router(classification_rules_router)
+325
View File
@@ -0,0 +1,325 @@
"""Classification Rules API endpoints.
Provides CRUD operations for managing custom document classification rules.
System-wide rules (``owner_id IS NULL``) can only be managed by admins.
"""
from __future__ import annotations
import logging
from typing import Annotated, Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from app.auth import require_login
from app.database import get_db
from app.models import ClassificationRuleModel
from app.utils.classification_rules import (
BUILTIN_CATEGORIES,
RULE_TYPE_CONTENT,
RULE_TYPE_FILENAME,
RULE_TYPE_METADATA,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/classification-rules", tags=["classification"])
DbSession = Annotated[Session, Depends(get_db)]
_VALID_RULE_TYPES = {RULE_TYPE_FILENAME, RULE_TYPE_CONTENT, RULE_TYPE_METADATA}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _get_user_id(request: Request) -> str:
"""Extract the user identifier from the request session."""
user = getattr(request.state, "user", None)
if user and hasattr(user, "get"):
return user.get("sub") or user.get("email") or "anonymous"
return "anonymous"
def _is_admin(request: Request) -> bool:
"""Check whether the current user is an admin."""
user = getattr(request.state, "user", None)
if user and hasattr(user, "get"):
groups = user.get("groups", [])
return "admin" in groups or "Admin" in groups
return False
# ---------------------------------------------------------------------------
# Pydantic schemas
# ---------------------------------------------------------------------------
class RuleCreate(BaseModel):
"""Schema for creating a classification rule."""
name: str = Field(..., min_length=1, max_length=255)
category: str = Field(..., min_length=1, max_length=100)
rule_type: str = Field(..., description="One of: filename_pattern, content_keyword, metadata_match")
pattern: str = Field(..., min_length=1, max_length=1000)
priority: int = Field(default=0, ge=0, le=1000)
case_sensitive: bool = False
enabled: bool = True
class RuleUpdate(BaseModel):
"""Schema for updating a classification rule."""
name: str | None = Field(default=None, min_length=1, max_length=255)
category: str | None = Field(default=None, min_length=1, max_length=100)
rule_type: str | None = Field(default=None)
pattern: str | None = Field(default=None, min_length=1, max_length=1000)
priority: int | None = Field(default=None, ge=0, le=1000)
case_sensitive: bool | None = None
enabled: bool | None = None
class RuleResponse(BaseModel):
"""Schema for a classification rule response."""
id: int
owner_id: str | None
name: str
category: str
rule_type: str
pattern: str
priority: int
case_sensitive: bool
enabled: bool
model_config = {"from_attributes": True}
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.get("/categories")
@require_login
async def list_categories(request: Request) -> dict[str, str]:
"""Return all built-in classification categories.
Custom categories created via rules are not included here; they are
discovered dynamically when rules are evaluated.
"""
return BUILTIN_CATEGORIES
@router.get("/rule-types")
@require_login
async def list_rule_types(request: Request) -> list[dict[str, str]]:
"""Return the supported rule types with descriptions."""
return [
{
"type": RULE_TYPE_FILENAME,
"label": "Filename Pattern",
"description": "Regex pattern matched against the original filename.",
},
{
"type": RULE_TYPE_CONTENT,
"label": "Content Keyword",
"description": "Pipe-separated keywords matched against the OCR text.",
},
{
"type": RULE_TYPE_METADATA,
"label": "Metadata Match",
"description": "field=value pattern matched against existing AI metadata.",
},
]
@router.get("/")
@require_login
async def list_rules(request: Request, db: DbSession) -> list[dict[str, Any]]:
"""List classification rules visible to the current user.
Returns both system rules (``owner_id IS NULL``) and the user's own rules.
"""
user_id = _get_user_id(request)
rules = (
db.query(ClassificationRuleModel)
.filter((ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == user_id))
.order_by(ClassificationRuleModel.priority.desc(), ClassificationRuleModel.id)
.all()
)
return [
{
"id": r.id,
"owner_id": r.owner_id,
"name": r.name,
"category": r.category,
"rule_type": r.rule_type,
"pattern": r.pattern,
"priority": r.priority,
"case_sensitive": r.case_sensitive,
"enabled": r.enabled,
}
for r in rules
]
@router.post("/", status_code=status.HTTP_201_CREATED)
@require_login
async def create_rule(request: Request, body: RuleCreate, db: DbSession) -> dict[str, Any]:
"""Create a new custom classification rule.
The rule is owned by the current user. Admins may create system-wide
rules by setting ``owner_id`` to ``null`` (not yet exposed).
"""
if body.rule_type not in _VALID_RULE_TYPES:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
)
user_id = _get_user_id(request)
# Check for duplicate name within the user's scope
existing = (
db.query(ClassificationRuleModel)
.filter(ClassificationRuleModel.owner_id == user_id, ClassificationRuleModel.name == body.name)
.first()
)
if existing:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=f"A rule named '{body.name}' already exists.",
)
rule = ClassificationRuleModel(
owner_id=user_id,
name=body.name,
category=body.category,
rule_type=body.rule_type,
pattern=body.pattern,
priority=body.priority,
case_sensitive=body.case_sensitive,
enabled=body.enabled,
)
try:
db.add(rule)
db.commit()
db.refresh(rule)
except Exception:
db.rollback()
raise
logger.info("Classification rule created: id=%s, user=%s", rule.id, user_id)
return {
"id": rule.id,
"owner_id": rule.owner_id,
"name": rule.name,
"category": rule.category,
"rule_type": rule.rule_type,
"pattern": rule.pattern,
"priority": rule.priority,
"case_sensitive": rule.case_sensitive,
"enabled": rule.enabled,
}
@router.get("/{rule_id}")
@require_login
async def get_rule(request: Request, rule_id: int, db: DbSession) -> dict[str, Any]:
"""Get a single classification rule by ID."""
user_id = _get_user_id(request)
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
if rule is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
# Users can see system rules and their own rules
if rule.owner_id is not None and rule.owner_id != user_id and not _is_admin(request):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
return {
"id": rule.id,
"owner_id": rule.owner_id,
"name": rule.name,
"category": rule.category,
"rule_type": rule.rule_type,
"pattern": rule.pattern,
"priority": rule.priority,
"case_sensitive": rule.case_sensitive,
"enabled": rule.enabled,
}
@router.put("/{rule_id}")
@require_login
async def update_rule(request: Request, rule_id: int, body: RuleUpdate, db: DbSession) -> dict[str, Any]:
"""Update an existing classification rule.
Users can only update their own rules. Admins can update any rule.
"""
user_id = _get_user_id(request)
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
if rule is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
if rule.owner_id != user_id and not _is_admin(request):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this rule")
if body.rule_type is not None and body.rule_type not in _VALID_RULE_TYPES:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
)
update_data = body.model_dump(exclude_unset=True)
for field_name, value in update_data.items():
setattr(rule, field_name, value)
try:
db.commit()
db.refresh(rule)
except Exception:
db.rollback()
raise
logger.info("Classification rule updated: id=%s, user=%s", rule.id, user_id)
return {
"id": rule.id,
"owner_id": rule.owner_id,
"name": rule.name,
"category": rule.category,
"rule_type": rule.rule_type,
"pattern": rule.pattern,
"priority": rule.priority,
"case_sensitive": rule.case_sensitive,
"enabled": rule.enabled,
}
@router.delete("/{rule_id}", status_code=status.HTTP_204_NO_CONTENT)
@require_login
async def delete_rule(request: Request, rule_id: int, db: DbSession) -> None:
"""Delete a classification rule.
Users can only delete their own rules. Admins can delete any rule.
"""
user_id = _get_user_id(request)
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
if rule is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
if rule.owner_id != user_id and not _is_admin(request):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot delete this rule")
try:
db.delete(rule)
db.commit()
except Exception:
db.rollback()
raise
logger.info("Classification rule deleted: id=%s, user=%s", rule_id, user_id)
+8 -2
View File
@@ -117,8 +117,14 @@ PIPELINE_STEP_TYPES: dict[str, dict[str, Any]] = {
},
"classify": {
"label": "Document Classification",
"description": "Classify the document type using AI without full metadata extraction.",
"config_schema": {},
"description": "Classify the document type using built-in and custom rules (filename patterns, content keywords, metadata matching).",
"config_schema": {
"use_builtin_rules": {
"type": "boolean",
"default": True,
"description": "Include the pre-built classification rules (invoice, contract, receipt, etc.).",
},
},
},
}
+1
View File
@@ -22,6 +22,7 @@ from app.tasks.batch_tasks import ( # noqa: F401
sync_search_index,
)
from app.tasks.check_credentials import check_credentials
from app.tasks.classify_document import classify_document_task # noqa: F401
from app.tasks.compute_embedding import backfill_missing_embeddings, compute_document_embedding # noqa: F401
from app.tasks.convert_to_pdf import convert_to_pdf # noqa: F401
from app.tasks.convert_to_pdfa import convert_to_pdfa # noqa: F401
+45
View File
@@ -833,3 +833,48 @@ class ScheduledJob(Base):
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
class ClassificationRuleModel(Base):
"""Custom document classification rule.
Rules are evaluated during the ``classify`` pipeline step to assign a
category to a document. System-wide rules have ``owner_id IS NULL``;
user-specific rules belong to a single owner.
"""
__tablename__ = "classification_rules"
id = Column(Integer, primary_key=True, index=True)
# NULL = system-wide rule visible to all users.
owner_id = Column(String, nullable=True, index=True)
# Human-readable rule name (unique per owner).
name = Column(String(255), nullable=False)
# Target category (e.g. "invoice", "contract", "receipt").
category = Column(String(100), nullable=False, index=True)
# Rule type: "filename_pattern", "content_keyword", or "metadata_match".
rule_type = Column(String(50), nullable=False)
# The matching pattern:
# - filename_pattern: a regex
# - content_keyword: pipe-separated keywords
# - metadata_match: "field=value"
pattern = Column(String(1000), nullable=False)
# Higher priority rules are evaluated first (default 0).
priority = Column(Integer, nullable=False, default=0)
# Whether pattern matching is case-sensitive.
case_sensitive = Column(Boolean, nullable=False, default=False)
# Disabled rules are skipped during classification.
enabled = Column(Boolean, nullable=False, default=True)
created_at = Column(DateTime(timezone=True), server_default=func.now())
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
__table_args__ = (UniqueConstraint("owner_id", "name", name="uq_classification_rules_owner_name"),)
+173
View File
@@ -0,0 +1,173 @@
"""Celery task for rule-based document classification.
This task is executed as a pipeline step (``step_type="classify"``). It
applies built-in and user-defined classification rules against the document's
filename, OCR text, and existing AI metadata to assign a ``document_type``
category.
The result is stored in the ``ai_metadata`` JSON blob on the
:class:`~app.models.FileRecord` (field ``classification``).
"""
from __future__ import annotations
import json
import logging
from typing import Any
from app.celery_app import celery
from app.database import SessionLocal
from app.models import ClassificationRuleModel, FileRecord
from app.tasks.retry_config import BaseTaskWithRetry
from app.utils import log_task_progress
from app.utils.classification_rules import (
ClassificationResult,
classify_document,
db_rule_to_engine_rule,
)
logger = logging.getLogger(__name__)
STEP_NAME = "classify_document"
def _load_custom_rules(owner_id: str | None) -> list[Any]:
"""Load enabled custom classification rules from the database.
Returns engine-level :class:`ClassificationRule` dataclass instances.
Rules are loaded in priority-descending order. System rules
(``owner_id IS NULL``) and the user's own rules are both included.
"""
with SessionLocal() as db:
query = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.enabled.is_(True))
if owner_id:
query = query.filter(
(ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == owner_id)
)
else:
query = query.filter(ClassificationRuleModel.owner_id.is_(None))
rules = query.order_by(ClassificationRuleModel.priority.desc()).all()
return [db_rule_to_engine_rule(r) for r in rules]
@celery.task(base=BaseTaskWithRetry, bind=True)
def classify_document_task(
self: Any,
file_id: int,
owner_id: str | None = None,
) -> dict[str, Any]:
"""Classify a document using rule-based matching.
This task:
1. Loads the :class:`FileRecord` from the database.
2. Gathers filename, OCR text, and existing AI metadata.
3. Loads built-in + user-defined classification rules.
4. Runs the classification engine.
5. Persists the result into ``ai_metadata.classification``.
Args:
file_id: Primary key of the :class:`FileRecord` to classify.
owner_id: Owner identifier for loading user-specific rules.
Returns:
Dict with ``category``, ``confidence``, and ``matched_rules``.
"""
task_id = self.request.id
log_task_progress(
task_id,
STEP_NAME,
"in_progress",
f"Starting classification for file {file_id}",
file_id=file_id,
)
try:
with SessionLocal() as db:
file_record: FileRecord | None = db.query(FileRecord).filter(FileRecord.id == file_id).first()
if file_record is None:
log_task_progress(
task_id,
STEP_NAME,
"failure",
f"FileRecord {file_id} not found",
file_id=file_id,
)
return {"status": "error", "detail": "File not found"}
# Gather inputs
filename = file_record.original_filename or ""
text = file_record.ocr_text or ""
existing_metadata: dict[str, Any] = {}
if file_record.ai_metadata:
try:
existing_metadata = json.loads(file_record.ai_metadata)
except (json.JSONDecodeError, TypeError):
existing_metadata = {}
# Load custom rules
effective_owner = owner_id or file_record.owner_id
custom_rules = _load_custom_rules(effective_owner)
# Run classification engine
result: ClassificationResult = classify_document(
filename=filename,
text=text,
metadata=existing_metadata,
custom_rules=custom_rules,
)
# Persist result into ai_metadata
classification_data = {
"category": result.category,
"confidence": result.confidence,
"matched_rules": [
{
"rule_name": m.rule_name,
"rule_type": m.rule_type,
"category": m.category,
"confidence": m.confidence,
}
for m in result.matched_rules
],
}
existing_metadata["classification"] = classification_data
# If no document_type was set yet, populate it from the classification
if not existing_metadata.get("document_type"):
from app.utils.classification_rules import BUILTIN_CATEGORIES
existing_metadata["document_type"] = BUILTIN_CATEGORIES.get(
result.category, result.category.replace("_", " ").title()
)
file_record.ai_metadata = json.dumps(existing_metadata, ensure_ascii=False)
db.commit()
log_task_progress(
task_id,
STEP_NAME,
"success",
f"Classified as '{result.category}' with confidence {result.confidence}",
file_id=file_id,
detail=f"Matched {len(result.matched_rules)} rule(s)",
)
return {
"status": "success",
"category": result.category,
"confidence": result.confidence,
"matched_rules": len(result.matched_rules),
}
except Exception as e:
logger.exception("Classification failed for file %s: %s", file_id, e)
log_task_progress(
task_id,
STEP_NAME,
"failure",
f"Classification failed: {e}",
file_id=file_id,
)
raise
+384
View File
@@ -0,0 +1,384 @@
"""
Rule-based document classification engine.
Provides pre-built categories and a rule matcher that classifies documents
using filename patterns, content keywords, and metadata fields. Custom
rules stored in the database are evaluated alongside the built-in defaults.
Usage::
from app.utils.classification_rules import classify_document
result = classify_document(
filename="2024-03-01_Invoice_Acme.pdf",
text="Invoice total: $1,234.56",
metadata={"absender": "Acme Corp"},
custom_rules=custom_rules_from_db,
)
# result -> ClassificationResult(category="invoice", confidence=85, matched_rules=[...])
"""
from __future__ import annotations
import logging
import re
from dataclasses import dataclass, field
from typing import Any
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Pre-built categories
# ---------------------------------------------------------------------------
#: Canonical category names recognised by the system. Users may also define
#: their own categories via custom rules.
BUILTIN_CATEGORIES: dict[str, str] = {
"invoice": "Invoice",
"contract": "Contract",
"receipt": "Receipt",
"letter": "Letter",
"report": "Report",
"bank_statement": "Bank Statement",
"tax_document": "Tax Document",
"insurance": "Insurance Document",
"payslip": "Payslip",
"unknown": "Unknown",
}
# ---------------------------------------------------------------------------
# Rule type constants
# ---------------------------------------------------------------------------
RULE_TYPE_FILENAME = "filename_pattern"
RULE_TYPE_CONTENT = "content_keyword"
RULE_TYPE_METADATA = "metadata_match"
# ---------------------------------------------------------------------------
# Data classes
# ---------------------------------------------------------------------------
@dataclass
class ClassificationRule:
"""A single classification rule."""
name: str
category: str
rule_type: str # filename_pattern | content_keyword | metadata_match
pattern: str # regex for filename, keyword(s) for content, "field=value" for metadata
priority: int = 0 # higher = evaluated first
case_sensitive: bool = False
def __post_init__(self) -> None:
if self.rule_type not in (RULE_TYPE_FILENAME, RULE_TYPE_CONTENT, RULE_TYPE_METADATA):
raise ValueError(f"Invalid rule_type: {self.rule_type!r}")
@dataclass
class MatchedRule:
"""Records which rule matched and why."""
rule_name: str
rule_type: str
category: str
confidence: int
@dataclass
class ClassificationResult:
"""The outcome of running the classification engine on a document."""
category: str
confidence: int # 0 100
matched_rules: list[MatchedRule] = field(default_factory=list)
# ---------------------------------------------------------------------------
# Built-in rules
# ---------------------------------------------------------------------------
BUILTIN_RULES: list[ClassificationRule] = [
# ── Invoice ───────────────────────────────────────────────────────────
ClassificationRule("builtin_invoice_filename", "invoice", RULE_TYPE_FILENAME, r"(?i)invoice|rechnung|facture"),
ClassificationRule(
"builtin_invoice_content",
"invoice",
RULE_TYPE_CONTENT,
"invoice number|invoice total|amount due|rechnung|rechnungsnummer|total amount|bill to",
),
ClassificationRule("builtin_invoice_metadata", "invoice", RULE_TYPE_METADATA, "document_type=Invoice"),
ClassificationRule(
"builtin_invoice_kommunikationsart", "invoice", RULE_TYPE_METADATA, "kommunikationsart=Rechnung"
),
# ── Contract ──────────────────────────────────────────────────────────
ClassificationRule("builtin_contract_filename", "contract", RULE_TYPE_FILENAME, r"(?i)contract|vertrag|agreement"),
ClassificationRule(
"builtin_contract_content",
"contract",
RULE_TYPE_CONTENT,
"hereby agrees|terms and conditions|vertrag|agreement between|party agrees|effective date",
),
ClassificationRule("builtin_contract_metadata", "contract", RULE_TYPE_METADATA, "document_type=Contract"),
ClassificationRule(
"builtin_contract_kommunikationsart", "contract", RULE_TYPE_METADATA, "kommunikationsart=Vertrag"
),
# ── Receipt ───────────────────────────────────────────────────────────
ClassificationRule("builtin_receipt_filename", "receipt", RULE_TYPE_FILENAME, r"(?i)receipt|quittung|beleg"),
ClassificationRule(
"builtin_receipt_content",
"receipt",
RULE_TYPE_CONTENT,
"receipt|quittung|payment received|thank you for your purchase|transaction id",
),
ClassificationRule("builtin_receipt_metadata", "receipt", RULE_TYPE_METADATA, "document_type=Receipt"),
ClassificationRule(
"builtin_receipt_kommunikationsart", "receipt", RULE_TYPE_METADATA, "kommunikationsart=Quittung"
),
# ── Letter ────────────────────────────────────────────────────────────
ClassificationRule("builtin_letter_filename", "letter", RULE_TYPE_FILENAME, r"(?i)letter|brief|schreiben"),
ClassificationRule(
"builtin_letter_content",
"letter",
RULE_TYPE_CONTENT,
"dear sir|dear madam|sehr geehrte|to whom it may concern|sincerely|mit freundlichen",
),
# ── Report ────────────────────────────────────────────────────────────
ClassificationRule("builtin_report_filename", "report", RULE_TYPE_FILENAME, r"(?i)report|bericht"),
ClassificationRule(
"builtin_report_content",
"report",
RULE_TYPE_CONTENT,
"executive summary|table of contents|annual report|quarterly report|findings",
),
# ── Bank statement ────────────────────────────────────────────────────
ClassificationRule(
"builtin_bank_filename",
"bank_statement",
RULE_TYPE_FILENAME,
r"(?i)bank.?statement|kontoauszug",
),
ClassificationRule(
"builtin_bank_content",
"bank_statement",
RULE_TYPE_CONTENT,
"account statement|kontoauszug|opening balance|closing balance|account number",
),
ClassificationRule(
"builtin_bank_kommunikationsart", "bank_statement", RULE_TYPE_METADATA, "kommunikationsart=Kontoauszug"
),
# ── Tax document ──────────────────────────────────────────────────────
ClassificationRule("builtin_tax_filename", "tax_document", RULE_TYPE_FILENAME, r"(?i)tax|steuer|steuerbescheid"),
ClassificationRule(
"builtin_tax_content",
"tax_document",
RULE_TYPE_CONTENT,
"tax return|steuerbescheid|taxable income|finanzamt|tax assessment",
),
# ── Insurance ─────────────────────────────────────────────────────────
ClassificationRule(
"builtin_insurance_filename", "insurance", RULE_TYPE_FILENAME, r"(?i)insurance|versicherung|police"
),
ClassificationRule(
"builtin_insurance_content",
"insurance",
RULE_TYPE_CONTENT,
"insurance policy|versicherung|policennummer|coverage|premium|deductible",
),
# ── Payslip ───────────────────────────────────────────────────────────
ClassificationRule(
"builtin_payslip_filename", "payslip", RULE_TYPE_FILENAME, r"(?i)payslip|gehaltsabrechnung|lohnabrechnung"
),
ClassificationRule(
"builtin_payslip_content",
"payslip",
RULE_TYPE_CONTENT,
"gross salary|net salary|gehaltsabrechnung|lohnabrechnung|bruttolohn|nettolohn",
),
]
# ---------------------------------------------------------------------------
# Confidence scoring
# ---------------------------------------------------------------------------
#: Base confidence for each rule type when it matches.
_CONFIDENCE_MAP: dict[str, int] = {
RULE_TYPE_FILENAME: 60,
RULE_TYPE_CONTENT: 70,
RULE_TYPE_METADATA: 90,
}
#: Extra confidence per additional matching rule of the same category (capped).
_CONFIDENCE_BONUS_PER_EXTRA_RULE = 10
# ---------------------------------------------------------------------------
# Matching helpers
# ---------------------------------------------------------------------------
def _match_filename(rule: ClassificationRule, filename: str) -> bool:
"""Return True if *rule.pattern* (regex) matches anywhere in *filename*."""
if not filename:
return False
flags = 0 if rule.case_sensitive else re.IGNORECASE
return bool(re.search(rule.pattern, filename, flags))
def _match_content(rule: ClassificationRule, text: str) -> bool:
"""Return True if any keyword in *rule.pattern* appears in *text*.
Keywords are separated by ``|`` (pipe).
"""
if not text:
return False
keywords = [kw.strip() for kw in rule.pattern.split("|") if kw.strip()]
text_lower = text if rule.case_sensitive else text.lower()
return any((kw if rule.case_sensitive else kw.lower()) in text_lower for kw in keywords)
def _match_metadata(rule: ClassificationRule, metadata: dict[str, Any] | None) -> bool:
"""Return True if *rule.pattern* (``field=value``) matches *metadata*.
Pattern format: ``field_name=expected_value``.
"""
if not metadata:
return False
if "=" not in rule.pattern:
return False
field_name, expected_value = rule.pattern.split("=", 1)
actual = metadata.get(field_name.strip())
if actual is None:
return False
if rule.case_sensitive:
return str(actual) == expected_value.strip()
return str(actual).lower() == expected_value.strip().lower()
_MATCHERS = {
RULE_TYPE_FILENAME: _match_filename,
RULE_TYPE_CONTENT: _match_content,
RULE_TYPE_METADATA: _match_metadata,
}
def _evaluate_rule(
rule: ClassificationRule,
filename: str,
text: str,
metadata: dict[str, Any] | None,
) -> MatchedRule | None:
"""Evaluate a single rule against the document. Return a :class:`MatchedRule` on match."""
matcher = _MATCHERS.get(rule.rule_type)
if matcher is None:
return None
# Dispatch to the appropriate matcher based on rule type
if rule.rule_type == RULE_TYPE_FILENAME:
matched = matcher(rule, filename)
elif rule.rule_type == RULE_TYPE_CONTENT:
matched = matcher(rule, text)
elif rule.rule_type == RULE_TYPE_METADATA:
matched = matcher(rule, metadata)
else:
matched = False
if matched:
return MatchedRule(
rule_name=rule.name,
rule_type=rule.rule_type,
category=rule.category,
confidence=_CONFIDENCE_MAP.get(rule.rule_type, 50),
)
return None
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def classify_document(
filename: str = "",
text: str = "",
metadata: dict[str, Any] | None = None,
custom_rules: list[ClassificationRule] | None = None,
) -> ClassificationResult:
"""Classify a document by evaluating built-in and custom rules.
Rules are evaluated in priority order (highest first, then built-in before
custom for the same priority). The category with the most rule matches
wins; ties are broken by cumulative confidence.
Args:
filename: Original filename of the document.
text: Extracted / OCR text of the document.
metadata: Previously-extracted AI metadata dict (e.g. from ``ai_metadata``).
custom_rules: Optional list of user-defined :class:`ClassificationRule` objects.
Returns:
A :class:`ClassificationResult` with the best matching category,
overall confidence score, and the list of rules that fired.
"""
all_rules = list(BUILTIN_RULES)
if custom_rules:
all_rules.extend(custom_rules)
# Sort by priority descending (higher priority first)
all_rules.sort(key=lambda r: r.priority, reverse=True)
matches: list[MatchedRule] = []
for rule in all_rules:
result = _evaluate_rule(rule, filename, text, metadata)
if result is not None:
matches.append(result)
if not matches:
return ClassificationResult(category="unknown", confidence=0, matched_rules=[])
# Aggregate by category: pick the one with the most matches, then highest
# cumulative confidence as tiebreaker.
category_scores: dict[str, list[MatchedRule]] = {}
for m in matches:
category_scores.setdefault(m.category, []).append(m)
best_category = max(
category_scores,
key=lambda cat: (len(category_scores[cat]), sum(m.confidence for m in category_scores[cat])),
)
best_matches = category_scores[best_category]
base_confidence = max(m.confidence for m in best_matches)
bonus = min(
(len(best_matches) - 1) * _CONFIDENCE_BONUS_PER_EXTRA_RULE,
100 - base_confidence,
)
final_confidence = min(base_confidence + bonus, 100)
return ClassificationResult(
category=best_category,
confidence=final_confidence,
matched_rules=best_matches,
)
def db_rule_to_engine_rule(db_rule: Any) -> ClassificationRule:
"""Convert a database ``ClassificationRuleModel`` row to an engine :class:`ClassificationRule`.
Args:
db_rule: A SQLAlchemy model instance with ``name``, ``category``,
``rule_type``, ``pattern``, ``priority``, and ``case_sensitive`` attributes.
Returns:
A :class:`ClassificationRule` dataclass instance.
"""
return ClassificationRule(
name=db_rule.name,
category=db_rule.category,
rule_type=db_rule.rule_type,
pattern=db_rule.pattern,
priority=db_rule.priority,
case_sensitive=getattr(db_rule, "case_sensitive", False),
)
+1 -3
View File
@@ -396,9 +396,7 @@ _STEP_TYPE_TO_STAGES: dict[str, list[str]] = {
"embed_metadata": ["embed_metadata_into_pdf"],
"compute_embedding": ["compute_embedding"],
"send_to_destinations": ["finalize_document_storage", "send_to_all_destinations"],
# "classify" is defined in PIPELINE_STEP_TYPES but has no Celery log stages yet.
# When a classify task is implemented, add its stage key(s) here.
"classify": [],
"classify": ["classify_document"],
}
# These internal bookkeeping stages are always shown in the flow regardless of