Files
doczyai-pipelines/src/document_classification/main.py
T
Rahul Ailaboina afc73987aa Merged in dtc_report (pull request #902)
Dtc report

* dtc_report_added

* lint format fixed

* Changed regex location call


Approved-by: Katon Minhas
2026-03-09 18:51:53 +00:00

715 lines
26 KiB
Python

import concurrent.futures
import json
import logging
import os
import re
from datetime import datetime
from threading import Lock
import pandas as pd
from src.constants.regex_patterns import (
TIN_PATTERN,
EFFECTIVE_DATE_KEYWORD_PATTERNS,
EFFECTIVE_DATE_FORMAT_PATTERNS,
)
from src.pipelines.shared.preprocessing.preprocessing_funcs import (
clean_newlines,
split_text,
clean_law_symbols,
)
from src.prompts.fieldset import Field
from src.utils import io_utils, llm_utils
from src import config
# Load prompts from JSON using Field class (with thread-safe caching)
# Load prompts at module level (cached for reuse across threads)
DOCUMENT_TYPE_CLASSIFICATION_PROMPT = Field.load_from_file(
config.DTC_PROMPTS_JSON_PATH, "document_type_classification"
).prompt
CONTRACT_TYPE_PROMPT = Field.load_from_file(
config.DTC_PROMPTS_JSON_PATH, "contract_type"
).prompt
NON_CONTRACT_TYPE_PROMPT = Field.load_from_file(
config.DTC_PROMPTS_JSON_PATH, "non_contract_type"
).prompt
# Thread-safe counter for progress tracking
progress_lock = Lock()
progress_counter = {"completed": 0, "total": 0}
# ── Pre-compiled patterns for DTC Report keyword search ─────────────────────
_TIN_COMPILED = re.compile(TIN_PATTERN)
_EFFECTIVE_DATE_COMPILED = [
re.compile(p)
for p in EFFECTIVE_DATE_KEYWORD_PATTERNS + EFFECTIVE_DATE_FORMAT_PATTERNS
]
# ── LOB keywords loaded from JSON mappings (single source of truth) ─────────
_MAPPINGS_DIR = os.path.join(os.path.dirname(__file__), "..", "constants", "mappings")
_LOB_JSON_FILES = [
"crosswalk_lob.json",
"crosswalk_program_lob.json",
"crosswalk_product_lob.json",
]
def _load_lob_keywords_from_mappings():
"""Load LOB keywords and their normalized LOB values from mapping JSON files.
Reads crosswalk_lob.json, crosswalk_program_lob.json, and crosswalk_product_lob.json
and extracts all keyword -> normalized_LOB pairs from their mapping, state_mapping,
and client_mapping sections.
Returns:
dict: Mapping of keyword (str) -> normalized LOB category (str).
e.g. {"Medicaid": "Medicaid", "CHIP": "Medicaid", "TENNCARE": "Medicaid", ...}
"""
keyword_to_lob = {}
for json_file in _LOB_JSON_FILES:
filepath = os.path.join(_MAPPINGS_DIR, json_file)
if not os.path.exists(filepath):
logging.warning(f"LOB mapping file not found: {filepath}")
continue
with open(filepath, "r") as f:
data = json.load(f)
# Extract from "mapping" section (key=keyword, value=normalized LOB)
if "mapping" in data:
for keyword, normalized_lob in data["mapping"].items():
if keyword and normalized_lob:
keyword_to_lob[keyword] = normalized_lob
# Extract from "state_mapping" section (state -> {keyword: normalized LOB})
if "state_mapping" in data:
for _state, state_keywords in data["state_mapping"].items():
if isinstance(state_keywords, dict):
for keyword, normalized_lob in state_keywords.items():
if keyword and normalized_lob:
keyword_to_lob[keyword] = normalized_lob
# Extract from "client_mapping" section (client -> {keyword: normalized LOB})
if "client_mapping" in data:
for _client, client_keywords in data["client_mapping"].items():
if isinstance(client_keywords, dict):
for keyword, normalized_lob in client_keywords.items():
if keyword and normalized_lob:
keyword_to_lob[keyword] = normalized_lob
logging.info(
f"Loaded {len(keyword_to_lob)} LOB keywords from "
f"{len(_LOB_JSON_FILES)} mapping files"
)
return keyword_to_lob
# Load at module level: keyword -> normalized LOB
_LOB_KEYWORD_TO_NORMALIZED = _load_lob_keywords_from_mappings()
# Pre-compile regex for each LOB keyword (case-insensitive, word-boundary)
_LOB_COMPILED = [
(kw, normalized_lob, re.compile(r"(?i)\b" + re.escape(kw) + r"\b"))
for kw, normalized_lob in _LOB_KEYWORD_TO_NORMALIZED.items()
]
# Separate progress tracker for the DTC report (avoids conflicts with main())
_report_progress_lock = Lock()
_report_progress = {"completed": 0, "total": 0}
def call_llm(filename, context, question):
"""Call Claude through llm_utils.invoke_claude (unified LLM interface)
Args:
filename: Name of the file being processed (for logging/tracking)
context: Document text context
question: The prompt/question to ask the LLM
Returns:
str: The LLM's response, or "ERROR" if the call fails
"""
# Prepare the full prompt with context
prompt = f"Document text:\n{context}\n\n{question}"
try:
# Use the unified llm_utils interface (supports both local and EC2 modes)
answer = llm_utils.invoke_claude(
prompt=prompt,
model_id="sonnet_latest", # Uses config.MODEL_ALIASES["sonnet_latest"]
filename=filename,
max_tokens=8192,
)
return answer.strip()
except Exception as e:
logging.error(f"Error calling Claude for {filename}: {str(e)}")
return "ERROR"
def process_single_file(filename, contract_text, json_folder):
"""Process a single file to determine if it's a contract and classify it.
Args:
filename: Name of the file being processed
contract_text: Raw text content of the file
json_folder: Path to folder for saving individual JSON results
Returns:
tuple: (filename, result_dict)
"""
try:
contract_text = clean_newlines(contract_text)
contract_text = clean_law_symbols(contract_text)
text_dict = split_text(contract_text)
# Loop through n pages
is_contract = "NO"
contract_type = "N/A"
non_contract_type = "N/A"
found_on_page = "N/A"
for i in range(config.DTC_MAX_PAGES_TO_CHECK):
page_key = str(i + 1)
# Handle edge case: document has fewer pages than config.DTC_MAX_PAGES_TO_CHECK
if page_key not in text_dict:
break
context = text_dict[page_key]
# Ask if this page is a contract
is_contract_answer = call_llm(
filename,
context,
DOCUMENT_TYPE_CLASSIFICATION_PROMPT,
)
if is_contract_answer.strip().upper() == "YES":
# If contract found, get contract type
type_of_contract_answer = call_llm(
filename, context, CONTRACT_TYPE_PROMPT
)
is_contract = "YES"
contract_type = type_of_contract_answer
found_on_page = i + 1
break # Stop checking further pages once contract is found
# If not a contract, determine non-contract document type
if is_contract == "NO":
# Use the first page for non-contract classification
if "1" in text_dict:
context = text_dict["1"]
non_contract_type_answer = call_llm(
filename, context, NON_CONTRACT_TYPE_PROMPT
)
non_contract_type = non_contract_type_answer
found_on_page = 1
result = {
"is_contract": is_contract,
"contract_type": contract_type,
"non_contract_type": non_contract_type,
"found_on_page": found_on_page,
}
# Save result to JSON immediately
io_utils.save_result_to_json(filename, result, json_folder)
# Update progress counter
with progress_lock:
progress_counter["completed"] += 1
completed = progress_counter["completed"]
total = progress_counter["total"]
if (
completed % 5 == 0 or completed == total
): # Log every 5 files or at completion
logging.info(
f"Progress: {completed}/{total} files processed ({completed/total*100:.1f}%)"
)
return filename, result
except Exception as e:
logging.error(f"Error processing {filename}: {str(e)}")
error_result = {
"is_contract": "ERROR",
"contract_type": "ERROR",
"non_contract_type": "ERROR",
"found_on_page": "ERROR",
}
# Save error result to JSON
io_utils.save_result_to_json(filename, error_result, json_folder)
return filename, error_result
def main(input_dict, run_timestamp):
"""Run document type classification on a dictionary of files.
This function is designed to be called from investment/main.py as a pre-filter.
It classifies all files and returns a dict mapping filenames to classification results.
Args:
input_dict: Dictionary mapping filename -> file_text
run_timestamp: Timestamp string for organizing output files (format: run_YYYYMMDD_HH-MM_BATCHID)
Returns:
dict: Mapping of filename -> classification result dict with keys:
- is_contract: "YES" or "NO"
- contract_type: Type if contract, otherwise "N/A"
- non_contract_type: Type if non-contract, otherwise "N/A"
- found_on_page: Page number where classification was determined
Side effects:
- Saves results to S3 (if config.WRITE_TO_S3) or local filesystem
- Saves individual JSON files to config.DTC_JSON_OUTPUT_FOLDER
"""
logging.info(f"Starting DTC classification for {len(input_dict)} files...")
logging.info(f"Max pages to check: {config.DTC_MAX_PAGES_TO_CHECK}")
logging.info(f"Results will be saved to: {config.DTC_OUTPUT_FILE}")
# Initialize progress counter
progress_counter["total"] = len(input_dict)
progress_counter["completed"] = 0
# Create output folder
os.makedirs(config.DTC_JSON_OUTPUT_FOLDER, exist_ok=True)
logging.info(f"JSON output folder: {config.DTC_JSON_OUTPUT_FOLDER}")
# Record start time
start_time = datetime.now()
answer_dict = {}
# Process files concurrently
with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
future_to_file = {}
for filename, contract_text in input_dict.items():
future = executor.submit(
process_single_file,
filename,
contract_text,
config.DTC_JSON_OUTPUT_FOLDER,
)
future_to_file[future] = filename
# Collect results as they complete
for future in concurrent.futures.as_completed(future_to_file):
filename = future_to_file[future]
try:
result_filename, result_data = future.result()
answer_dict[result_filename] = result_data
except Exception as e:
logging.error(f"Exception occurred for {filename}: {str(e)}")
answer_dict[filename] = {
"is_contract": "ERROR",
"contract_type": "ERROR",
"non_contract_type": "ERROR",
"found_on_page": "ERROR",
}
# Record end time
end_time = datetime.now()
duration = (end_time - start_time).total_seconds()
logging.info(f"DTC classification complete in {duration:.2f} seconds")
logging.info(f"Average time per file: {duration/len(input_dict):.2f} seconds")
# Create DataFrame and save results
df = pd.DataFrame([{"FILE_NAME": k, **v} for k, v in answer_dict.items()])
# Log summary statistics
contracts_found = len(df[df["is_contract"] == "YES"])
non_contracts = len(df[df["is_contract"] == "NO"])
errors = len(df[df["is_contract"] == "ERROR"])
logging.info(
f"DTC Summary: {contracts_found} contracts, {non_contracts} non-contracts, {errors} errors"
)
# Create summary DataFrame with aggregate statistics
summary_data = {
"total_files": [len(df)],
"contracts": [contracts_found],
"non_contracts": [non_contracts],
"errors": [errors],
"contract_rate": [
f"{(contracts_found/len(df)*100):.1f}%" if len(df) > 0 else "0%"
],
"run_timestamp": [run_timestamp],
"batch_id": [config.BATCH_ID],
}
# Add contract type breakdown
if contracts_found > 0:
contract_types = df[df["is_contract"] == "YES"]["contract_type"].value_counts()
summary_data["most_common_contract_type"] = [
contract_types.index[0] if len(contract_types) > 0 else "N/A"
]
summary_data["contract_type_breakdown"] = [contract_types.to_dict()]
else:
summary_data["most_common_contract_type"] = ["N/A"]
summary_data["contract_type_breakdown"] = [{}]
# Add non-contract type breakdown
if non_contracts > 0:
non_contract_types = df[df["is_contract"] == "NO"][
"non_contract_type"
].value_counts()
summary_data["most_common_non_contract_type"] = [
non_contract_types.index[0] if len(non_contract_types) > 0 else "N/A"
]
summary_data["non_contract_type_breakdown"] = [non_contract_types.to_dict()]
else:
summary_data["most_common_non_contract_type"] = ["N/A"]
summary_data["non_contract_type_breakdown"] = [{}]
summary_df = pd.DataFrame(summary_data)
# Save detailed results using same pattern as investment pipeline
if config.WRITE_TO_S3:
io_utils.write_s3(df, "", run_timestamp, "dtc")
logging.info(
f"DTC results uploaded to S3: {config.BATCH_ID}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
)
# Save summary
io_utils.write_s3(summary_df, "", run_timestamp, "dtc_summary")
logging.info(
f"DTC summary uploaded to S3: {config.BATCH_ID}/{run_timestamp}/{config.BATCH_ID}-DTC-SUMMARY.csv"
)
else:
io_utils.write_local(df, "", run_timestamp, "dtc")
logging.info(
f"DTC results saved locally: {config.CONSOLIDATED_OUTPUT_DIRECTORY}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
)
# Save summary
io_utils.write_local(summary_df, "", run_timestamp, "dtc_summary")
logging.info(
f"DTC summary saved locally: {config.CONSOLIDATED_OUTPUT_DIRECTORY}/{run_timestamp}/{config.BATCH_ID}-DTC-SUMMARY.csv"
)
return answer_dict
# ── DTC Report: Keyword Search Helpers ──────────────────────────────────────
def _search_tin_in_page(page_text):
"""Check if TIN pattern is present in a single page of text.
Args:
page_text: Text content of a single page.
Returns:
bool: True if a TIN-like pattern was found.
"""
return bool(_TIN_COMPILED.search(page_text))
def _search_effective_date_in_page(page_text):
"""Check if any effective date keyword or date format is present in a single page.
Args:
page_text: Text content of a single page.
Returns:
bool: True if an effective date indicator was found.
"""
return any(p.search(page_text) for p in _EFFECTIVE_DATE_COMPILED)
def _search_lob_in_page(page_text):
"""Find all LOB keywords present in a single page of text.
Args:
page_text: Text content of a single page.
Returns:
list[tuple[str, str]]: List of (keyword, normalized_lob) tuples matched on this page.
"""
return [
(kw, normalized_lob)
for kw, normalized_lob, pattern in _LOB_COMPILED
if pattern.search(page_text)
]
# ── DTC Report: Per-file Processing ────────────────────────────────────────
def _process_file_for_report(filename, file_text, json_folder):
"""Process a single file for the DTC report: classify + keyword search.
Combines LLM-based document type classification with regex-based keyword
detection for TIN, effective date, and line of business across all pages.
Args:
filename: Name of the file being processed.
file_text: Raw text content (Textract output with page markers).
json_folder: Path to folder for saving individual JSON results.
Returns:
tuple: (filename, result_dict) where result_dict contains all report columns.
"""
try:
cleaned_text = clean_law_symbols(clean_newlines(file_text))
text_dict = split_text(cleaned_text)
# ── Step 1: Document Type Classification (LLM-based) ──
is_contract = "NO"
document_type = "N/A"
for i in range(config.DTC_MAX_PAGES_TO_CHECK):
page_key = str(i + 1)
if page_key not in text_dict:
break
context = text_dict[page_key]
is_contract_answer = call_llm(
filename, context, DOCUMENT_TYPE_CLASSIFICATION_PROMPT
)
if is_contract_answer.strip().upper() == "YES":
is_contract = "YES"
document_type = call_llm(filename, context, CONTRACT_TYPE_PROMPT)
break
if is_contract == "NO" and "1" in text_dict:
document_type = call_llm(filename, text_dict["1"], NON_CONTRACT_TYPE_PROMPT)
# ── Step 2: Keyword Search Across ALL Pages ──
tin_pages = []
eff_date_pages = []
lob_pages = []
lob_keywords_all = set()
lob_derived_all = set()
for page_key in sorted(text_dict.keys(), key=lambda x: int(x)):
page_text = text_dict[page_key]
if _search_tin_in_page(page_text):
tin_pages.append(page_key)
if _search_effective_date_in_page(page_text):
eff_date_pages.append(page_key)
matched_lob = _search_lob_in_page(page_text)
if matched_lob:
lob_pages.append(page_key)
for kw, normalized_lob in matched_lob:
lob_keywords_all.add(kw)
lob_derived_all.add(normalized_lob)
result = {
"IS_CONTRACT": is_contract,
"DOCUMENT_TYPE": document_type,
"TIN_FOUND": "YES" if tin_pages else "NO",
"TIN_PAGES": ", ".join(tin_pages) if tin_pages else "N/A",
"EFFECTIVE_DATE_FOUND": "YES" if eff_date_pages else "NO",
"EFFECTIVE_DATE_PAGES": (
", ".join(eff_date_pages) if eff_date_pages else "N/A"
),
"LOB_FOUND": "YES" if lob_pages else "NO",
"LOB_PAGES": ", ".join(lob_pages) if lob_pages else "N/A",
"LOB_KEYWORDS_MATCHED": (
", ".join(sorted(lob_keywords_all)) if lob_keywords_all else "N/A"
),
"LOB_DERIVED": (
", ".join(sorted(lob_derived_all)) if lob_derived_all else "N/A"
),
}
# Save individual result to JSON
io_utils.save_result_to_json(filename, result, json_folder)
# Update progress
with _report_progress_lock:
_report_progress["completed"] += 1
done = _report_progress["completed"]
total = _report_progress["total"]
if done % 5 == 0 or done == total:
logging.info(
f"DTC Report Progress: {done}/{total} files processed "
f"({done / total * 100:.1f}%)"
)
return filename, result
except Exception as e:
logging.error(f"Error processing {filename} for DTC report: {e}")
error_result = {
"IS_CONTRACT": "ERROR",
"DOCUMENT_TYPE": "ERROR",
"TIN_FOUND": "ERROR",
"TIN_PAGES": "ERROR",
"EFFECTIVE_DATE_FOUND": "ERROR",
"EFFECTIVE_DATE_PAGES": "ERROR",
"LOB_FOUND": "ERROR",
"LOB_PAGES": "ERROR",
"LOB_KEYWORDS_MATCHED": "ERROR",
"LOB_DERIVED": "ERROR",
}
io_utils.save_result_to_json(filename, error_result, json_folder)
return filename, error_result
# ── DTC Report: Main Entry Point ───────────────────────────────────────────
def generate_dtc_report(input_dict, run_timestamp):
"""Generate a DTC report combining document classification with keyword detection.
For each file, produces a row with:
- FILE_NAME: Original filename
- IS_CONTRACT: YES/NO from LLM classification
- DOCUMENT_TYPE: Contract type or non-contract type
- TIN_FOUND / TIN_PAGES: Whether TIN was detected and on which pages
- EFFECTIVE_DATE_FOUND / EFFECTIVE_DATE_PAGES: Whether effective date was found
- LOB_FOUND / LOB_PAGES / LOB_KEYWORDS_MATCHED: LOB detection results
Args:
input_dict: Dictionary mapping filename -> file_text.
run_timestamp: Timestamp string for organizing output files
(format: run_YYYYMMDD_HH-MM_BATCHID).
Returns:
pd.DataFrame: The complete DTC report DataFrame.
Side effects:
- Saves CSV results to S3 (if config.WRITE_TO_S3) or local filesystem.
- Saves individual JSON files to config.DTC_JSON_OUTPUT_FOLDER.
"""
logging.info(f"Starting DTC Report generation for {len(input_dict)} files...")
logging.info(
f"Max pages to check for classification: {config.DTC_MAX_PAGES_TO_CHECK}"
)
# Initialize progress
_report_progress["total"] = len(input_dict)
_report_progress["completed"] = 0
# Create output folder for individual JSONs
json_folder = config.DTC_JSON_OUTPUT_FOLDER + "_report"
os.makedirs(json_folder, exist_ok=True)
logging.info(f"DTC Report JSON output folder: {json_folder}")
start_time = datetime.now()
results = {}
# Process files concurrently
with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
future_to_file = {
executor.submit(_process_file_for_report, fname, ftext, json_folder): fname
for fname, ftext in input_dict.items()
}
for future in concurrent.futures.as_completed(future_to_file):
fname = future_to_file[future]
try:
result_fname, result_data = future.result()
results[result_fname] = result_data
except Exception as e:
logging.error(f"Exception for {fname} in DTC Report: {e}")
results[fname] = {
"IS_CONTRACT": "ERROR",
"DOCUMENT_TYPE": "ERROR",
"TIN_FOUND": "ERROR",
"TIN_PAGES": "ERROR",
"EFFECTIVE_DATE_FOUND": "ERROR",
"EFFECTIVE_DATE_PAGES": "ERROR",
"LOB_FOUND": "ERROR",
"LOB_PAGES": "ERROR",
"LOB_KEYWORDS_MATCHED": "ERROR",
"LOB_DERIVED": "ERROR",
}
duration = (datetime.now() - start_time).total_seconds()
logging.info(f"DTC Report generation complete in {duration:.2f} seconds")
if len(input_dict) > 0:
logging.info(f"Average time per file: {duration / len(input_dict):.2f} seconds")
# Build DataFrame with explicit column ordering
df = pd.DataFrame([{"FILE_NAME": k, **v} for k, v in results.items()])
col_order = [
"FILE_NAME",
"IS_CONTRACT",
"DOCUMENT_TYPE",
"TIN_FOUND",
"TIN_PAGES",
"EFFECTIVE_DATE_FOUND",
"EFFECTIVE_DATE_PAGES",
"LOB_FOUND",
"LOB_PAGES",
"LOB_KEYWORDS_MATCHED",
"LOB_DERIVED",
]
df = df[[c for c in col_order if c in df.columns]]
# Log summary statistics
total_files = len(df)
tin_count = (df["TIN_FOUND"] == "YES").sum() if "TIN_FOUND" in df.columns else 0
eff_date_count = (
(df["EFFECTIVE_DATE_FOUND"] == "YES").sum()
if "EFFECTIVE_DATE_FOUND" in df.columns
else 0
)
lob_count = (df["LOB_FOUND"] == "YES").sum() if "LOB_FOUND" in df.columns else 0
error_count = (
(df["IS_CONTRACT"] == "ERROR").sum() if "IS_CONTRACT" in df.columns else 0
)
logging.info(
f"DTC Report Summary: {total_files} files | "
f"TIN detected: {tin_count} | "
f"Effective Date detected: {eff_date_count} | "
f"LOB detected: {lob_count} | "
f"Errors: {error_count}"
)
# Save output using existing io_utils pattern
if config.WRITE_TO_S3:
io_utils.write_s3(df, "", run_timestamp, "dtc")
logging.info(
f"DTC Report uploaded to S3: "
f"{config.BATCH_ID}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
)
else:
io_utils.write_local(df, "", run_timestamp, "dtc")
logging.info(
f"DTC Report saved locally: "
f"{config.CONSOLIDATED_OUTPUT_DIRECTORY}/{run_timestamp}/{config.BATCH_ID}-DTC.csv"
)
return df
if __name__ == "__main__":
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
force=True,
)
# Suppress AWS SDK logging
logging.getLogger("botocore").setLevel(logging.WARNING)
logging.getLogger("boto3").setLevel(logging.WARNING)
# Generate timestamp for this run
run_timestamp = f"run_{datetime.now().strftime('%Y%m%d_%H-%M')}_{config.BATCH_ID}"
# Load input files
logging.info("Loading input files...")
input_dict = io_utils.read_input()
# Run DTC Report (classification + keyword search)
report_df = generate_dtc_report(input_dict, run_timestamp)
logging.info(f"DTC Report complete. Processed {len(report_df)} files.")