from __future__ import annotations import datetime import time from dataclasses import dataclass, field from pathlib import Path from typing import Any from ocr_pipeline.citations import extract_citations from ocr_pipeline.config import settings from ocr_pipeline.detectors import detect_figures, detect_tables from ocr_pipeline.ocr import ImagePreprocessor, OCREngine, compute_image_hash from ocr_pipeline.ocr.parallel import ParallelProcessor from ocr_pipeline.output import MarkdownWriter from ocr_pipeline.postprocess import chunk_text, clean_text, extract_entities from ocr_pipeline.utils.db import Database, get_db from ocr_pipeline.utils.logging import get_logger from ocr_pipeline.watch import Watcher logger = get_logger(__name__) @dataclass class PipelineResult: total_images: int = 0 successful: int = 0 failed: int = 0 skipped: int = 0 chunks_created: int = 0 output_files: list[Path] = field(default_factory=list) errors: list[dict[str, str]] = field(default_factory=list) class OCRPipeline: def __init__(self) -> None: self.engine = OCREngine() self.preprocessor = ImagePreprocessor() self.writer = MarkdownWriter() self.db: Database = get_db() self._run_id: int | None = None def _process_single(self, image_path: Path) -> list[Path]: image_hash = compute_image_hash(image_path) if settings.processing.skip_existing and self.db.is_processed(image_hash, image_path): logger.info("already_processed", path=str(image_path)) return [] original_image = self.preprocessor.load_image(image_path) # PaddleOCR reads the (size-capped) original; Tesseract the binarized derivative. ocr_base = self.preprocessor.resize_if_needed(original_image) preprocessed = self.preprocessor.preprocess_image(original_image) ocr_result = self.engine.process(ocr_base, preprocessed) cleaned = clean_text(ocr_result.text) chunks = chunk_text(cleaned.text) entities = extract_entities(cleaned.text) # Detectors receive the original image: binarization/line removal damages their inputs. figures = detect_figures(original_image) tables = detect_tables(original_image) citations = extract_citations(cleaned.text) timestamp = datetime.datetime.now(datetime.UTC).isoformat() output_files: list[Path] = [] for chunk in chunks: metadata = { "source_path": str(image_path), "source_hash": image_hash, "timestamp": timestamp, "ocr_engine": ocr_result.engine, "ocr_confidence_mean": ocr_result.confidence, "language": ocr_result.language, "entity_extraction_backend": entities.backend, "detected_entities": [entity.text for entity in entities.entities], "has_figures": bool(figures), "has_tables": bool(tables), "citations_found": [ f"{citation.type}:{citation.identifier}" for citation in citations ], "chunk_index": chunk.chunk_index, "total_chunks": len(chunks), } output_files.append( self.writer.write_chunk( content=chunk.content, metadata=metadata, figures=figures, tables=tables, entities=entities, citations=citations, ) ) self.db.mark_processed( file_hash=image_hash, path=str(image_path), timestamp=time.time(), engine=ocr_result.engine, confidence=ocr_result.confidence, chunks=len(chunks), run_id=self._run_id, ) return output_files def process_single(self, image_path: Path) -> list[Path]: """Public single-image API used by notebook migration and watch mode.""" return self._process_single(image_path) def process_all(self, image_paths: list[Path] | None = None) -> PipelineResult: if image_paths is None: image_paths = ParallelProcessor().find_images() else: image_paths = [path for path in image_paths if path.exists()] result = PipelineResult(total_images=len(image_paths)) if not image_paths: return result run_timestamp = datetime.datetime.now(datetime.UTC).isoformat() self._run_id, run_number = self.db.start_run() self.writer.start_run(run_timestamp, run_number) try: # Output/database writes are deliberately centralized to avoid collisions and SQLite contention. for path in image_paths: try: outputs = self._process_single(path) if not outputs and settings.processing.skip_existing: result.skipped += 1 else: result.successful += 1 result.chunks_created += len(outputs) result.output_files.extend(outputs) except Exception as exc: result.failed += 1 result.errors.append({"path": str(path), "error": str(exc)}) logger.exception("pipeline_failed", path=str(path)) if settings.output.write_consolidated: consolidated = self.writer.write_consolidated() if consolidated is not None: result.output_files.append(consolidated) finally: self.db.complete_run( self._run_id, { "total": result.total_images, "successful": result.successful, "failed": result.failed, "chunks": result.chunks_created, }, ) self._run_id = None return result def watch(self) -> None: def process_callback(path: Path) -> None: self.process_single(path) watcher = Watcher(process_callback) watcher.start() try: while True: time.sleep(1) except KeyboardInterrupt: watcher.stop() def get_stats(self, days: int = 30) -> dict[str, Any]: return self.db.get_stats(days)