File size: 3,722 Bytes
d86db02
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
import time
from dataclasses import dataclass, field
from pathlib import Path

import easyocr
from google import genai

from src.kyc_detector import detect_kyc
from src.llm_classifier import classify_with_llm
from src.monitoring import trace_classify, trace_kyc_ocr, trace_llm_classify, trace_rules_engine
from src.rules_engine import apply_rules


_SUB_TYPE_DEFAULTS = {"bill": "Bills", "kyc": "KYC"}


def _default_sub_type(doc_type: str) -> str | None:
    return _SUB_TYPE_DEFAULTS.get(doc_type)


@dataclass
class ClassificationResult:
    filename: str
    doc_type: str           # "bill" | "kyc" | "image"
    sub_type: str | None    # Gemini category for "image" docs, else None
    method: str             # "rules" | "ocr" | "llm"
    latency_ms: int
    input_tokens: int = 0
    output_tokens: int = 0


def classify(
    filename: str,
    image_bytes: bytes,
    ocr_reader: easyocr.Reader | None = None,
    llm_client: genai.Client | None = None,
) -> ClassificationResult:
    """Run a single image through the full classification pipeline."""
    start = time.monotonic()

    # Stage 1 — rules engine (filename-based, no ML)
    rules_result = apply_rules(filename)
    trace_rules_engine(
        filename=filename,
        doc_type=rules_result.doc_type,
        send_to_llm=rules_result.send_to_llm,
    )
    if not rules_result.send_to_llm:
        result = ClassificationResult(
            filename=filename,
            doc_type=rules_result.doc_type,
            sub_type=_default_sub_type(rules_result.doc_type),
            method="rules",
            latency_ms=int((time.monotonic() - start) * 1000),
        )
        trace_classify(**result.__dict__)
        return result

    # Stage 2 — KYC OCR detector
    kyc_result = detect_kyc(filename, image_bytes, reader=ocr_reader)
    trace_kyc_ocr(
        filename=filename,
        doc_type=kyc_result.doc_type,
        send_to_llm=kyc_result.send_to_llm,
        ocr_text=kyc_result.ocr_text,
    )
    if not kyc_result.send_to_llm:
        result = ClassificationResult(
            filename=filename,
            doc_type=kyc_result.doc_type,
            sub_type=_default_sub_type(kyc_result.doc_type),
            method="ocr",
            latency_ms=int((time.monotonic() - start) * 1000),
        )
        trace_classify(**result.__dict__)
        return result

    # Stage 3 — Gemini LLM classifier
    llm_result = classify_with_llm(filename, image_bytes, client=llm_client)
    trace_llm_classify(
        filename=filename,
        sub_type=llm_result.sub_type,
        input_tokens=llm_result.input_tokens,
        output_tokens=llm_result.output_tokens,
        raw_response=llm_result.raw_response,
    )
    result = ClassificationResult(
        filename=filename,
        doc_type=llm_result.doc_type,
        sub_type=llm_result.sub_type,
        method="llm",
        latency_ms=int((time.monotonic() - start) * 1000),
        input_tokens=llm_result.input_tokens,
        output_tokens=llm_result.output_tokens,
    )
    trace_classify(**result.__dict__)
    return result


def classify_dataset(
    dataset_dir: str | Path = "dataset",
    ocr_reader: easyocr.Reader | None = None,
    llm_client: genai.Client | None = None,
) -> list[ClassificationResult]:
    """Classify all PNG/JPEG images in a directory."""
    dataset_path = Path(dataset_dir)
    results = []

    for image_path in sorted(dataset_path.glob("*.png")):
        image_bytes = image_path.read_bytes()
        result = classify(
            filename=image_path.name,
            image_bytes=image_bytes,
            ocr_reader=ocr_reader,
            llm_client=llm_client,
        )
        results.append(result)

    return results