Files
OCR/docengine/app/repositories/template_repository.py

356 lines
13 KiB
Python

from __future__ import annotations
import uuid
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.models.template import (
DocumentCell,
DocumentFormat,
DocumentRegion,
ImageRegion,
TableColumn,
TableFormat,
TableRow,
TemplateFingerprint,
Watermark,
)
from app.repositories.base import BaseRepository
class TemplateRepository(BaseRepository[DocumentFormat]):
"""Repository for DocumentFormat (template) operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, DocumentFormat)
def get_active_templates(self, offset: int = 0, limit: int = 100) -> list[DocumentFormat]:
"""Get all active templates."""
query = (
select(DocumentFormat)
.where(DocumentFormat.is_active.is_(True))
.order_by(DocumentFormat.created_at.desc())
.offset(offset)
.limit(limit)
)
result = self.db.execute(query)
return list(result.scalars().all())
def count_active(self) -> int:
"""Count active templates."""
return self.count(filters={"is_active": True})
def get_by_name(self, name: str) -> DocumentFormat | None:
"""Get template by name."""
query = select(DocumentFormat).where(DocumentFormat.name == name)
result = self.db.execute(query)
return result.scalars().first()
def get_by_source_document(self, document_id: uuid.UUID) -> DocumentFormat | None:
"""Get template generated from a specific source document."""
query = select(DocumentFormat).where(
DocumentFormat.source_document_id == document_id,
DocumentFormat.is_active.is_(True),
)
result = self.db.execute(query)
return result.scalars().first()
def create_template(
self,
name: str,
page_width: float,
page_height: float,
page_count: int = 1,
description: str | None = None,
margin_top: float = 72.0,
margin_right: float = 72.0,
margin_bottom: float = 72.0,
margin_left: float = 72.0,
fingerprint: dict | None = None,
source_document_id: uuid.UUID | None = None,
created_by: uuid.UUID | None = None,
is_active: bool = True,
) -> DocumentFormat:
"""Create a new template."""
template = DocumentFormat(
name=name,
page_width=page_width,
page_height=page_height,
page_count=page_count,
description=description,
margin_top=margin_top,
margin_right=margin_right,
margin_bottom=margin_bottom,
margin_left=margin_left,
fingerprint=fingerprint,
source_document_id=source_document_id,
created_by=created_by,
is_active=is_active,
)
return self.create(template)
def deactivate_template(self, template_id: uuid.UUID) -> DocumentFormat | None:
"""Soft-delete a template by deactivating it."""
template = self.get_by_id(template_id)
if template:
template.is_active = False
self.db.flush()
self.db.refresh(template)
return template
def get_all_with_fingerprints(self) -> list[DocumentFormat]:
"""Get all active templates with their fingerprints."""
query = (
select(DocumentFormat)
.where(DocumentFormat.is_active.is_(True))
.order_by(DocumentFormat.created_at.desc())
)
result = self.db.execute(query)
return list(result.scalars().all())
class DocumentCellRepository(BaseRepository[DocumentCell]):
"""Repository for DocumentCell operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, DocumentCell)
def get_template_cells(self, format_id: uuid.UUID) -> list[DocumentCell]:
"""Get all cells for a template."""
query = (
select(DocumentCell)
.where(DocumentCell.format_id == format_id)
.order_by(DocumentCell.page_number, DocumentCell.sequence)
)
result = self.db.execute(query)
return list(result.scalars().all())
def get_page_cells(self, format_id: uuid.UUID, page_number: int) -> list[DocumentCell]:
"""Get cells for a specific page of a template."""
query = (
select(DocumentCell)
.where(
DocumentCell.format_id == format_id,
DocumentCell.page_number == page_number,
)
.order_by(DocumentCell.sequence)
)
result = self.db.execute(query)
return list(result.scalars().all())
def get_dynamic_cells(self, format_id: uuid.UUID) -> list[DocumentCell]:
"""Get all dynamic cells for a template."""
query = (
select(DocumentCell)
.where(
DocumentCell.format_id == format_id,
DocumentCell.is_dynamic.is_(True),
)
.order_by(DocumentCell.page_number, DocumentCell.sequence)
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_cell(self, format_id: uuid.UUID, **kwargs) -> DocumentCell: # noqa: ANN003
"""Create a new cell for a template."""
cell = DocumentCell(format_id=format_id, **kwargs)
return self.create(cell)
class DocumentRegionRepository(BaseRepository[DocumentRegion]):
"""Repository for DocumentRegion operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, DocumentRegion)
def get_template_regions(self, format_id: uuid.UUID) -> list[DocumentRegion]:
"""Get all regions for a template."""
query = (
select(DocumentRegion)
.where(DocumentRegion.format_id == format_id)
.order_by(DocumentRegion.page_number, DocumentRegion.sequence)
)
result = self.db.execute(query)
return list(result.scalars().all())
def get_regions_by_type(self, format_id: uuid.UUID, region_type: str) -> list[DocumentRegion]:
"""Get regions of a specific type."""
query = (
select(DocumentRegion)
.where(
DocumentRegion.format_id == format_id,
DocumentRegion.region_type == region_type,
)
.order_by(DocumentRegion.page_number, DocumentRegion.sequence)
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_region(self, format_id: uuid.UUID, **kwargs) -> DocumentRegion: # noqa: ANN003
"""Create a new region for a template."""
region = DocumentRegion(format_id=format_id, **kwargs)
return self.create(region)
class TableFormatRepository(BaseRepository[TableFormat]):
"""Repository for TableFormat operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, TableFormat)
def get_template_tables(self, format_id: uuid.UUID) -> list[TableFormat]:
"""Get all table formats for a template."""
query = (
select(TableFormat)
.where(TableFormat.format_id == format_id)
.order_by(TableFormat.page_number)
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_table_format(self, format_id: uuid.UUID, **kwargs) -> TableFormat: # noqa: ANN003
"""Create a new table format."""
table_format = TableFormat(format_id=format_id, **kwargs)
return self.create(table_format)
class TableColumnRepository(BaseRepository[TableColumn]):
"""Repository for TableColumn operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, TableColumn)
def get_table_columns(self, table_format_id: uuid.UUID) -> list[TableColumn]:
"""Get all columns for a table format."""
query = (
select(TableColumn)
.where(TableColumn.table_format_id == table_format_id)
.order_by(TableColumn.column_index)
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_column(self, table_format_id: uuid.UUID, **kwargs) -> TableColumn: # noqa: ANN003
"""Create a new table column."""
column = TableColumn(table_format_id=table_format_id, **kwargs)
return self.create(column)
class TableRowRepository(BaseRepository[TableRow]):
"""Repository for TableRow operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, TableRow)
def get_table_rows(self, table_format_id: uuid.UUID) -> list[TableRow]:
"""Get all rows for a table format."""
query = (
select(TableRow)
.where(TableRow.table_format_id == table_format_id)
.order_by(TableRow.row_index)
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_row(self, table_format_id: uuid.UUID, **kwargs) -> TableRow: # noqa: ANN003
"""Create a new table row."""
row = TableRow(table_format_id=table_format_id, **kwargs)
return self.create(row)
class WatermarkRepository(BaseRepository[Watermark]):
"""Repository for Watermark operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, Watermark)
def get_template_watermarks(self, format_id: uuid.UUID) -> list[Watermark]:
"""Get all watermarks for a template."""
query = select(Watermark).where(Watermark.format_id == format_id)
result = self.db.execute(query)
return list(result.scalars().all())
def create_watermark(self, format_id: uuid.UUID, **kwargs) -> Watermark: # noqa: ANN003
"""Create a new watermark."""
watermark = Watermark(format_id=format_id, **kwargs)
return self.create(watermark)
class ImageRegionRepository(BaseRepository[ImageRegion]):
"""Repository for ImageRegion operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, ImageRegion)
def get_template_images(self, format_id: uuid.UUID) -> list[ImageRegion]:
"""Get all image regions for a template."""
query = select(ImageRegion).where(ImageRegion.format_id == format_id)
result = self.db.execute(query)
return list(result.scalars().all())
def get_static_images(self, format_id: uuid.UUID) -> list[ImageRegion]:
"""Get static image regions."""
query = select(ImageRegion).where(
ImageRegion.format_id == format_id,
ImageRegion.is_static.is_(True),
)
result = self.db.execute(query)
return list(result.scalars().all())
def create_image_region(self, format_id: uuid.UUID, **kwargs) -> ImageRegion: # noqa: ANN003
"""Create a new image region."""
image_region = ImageRegion(format_id=format_id, **kwargs)
return self.create(image_region)
class TemplateFingerprintRepository(BaseRepository[TemplateFingerprint]):
"""Repository for TemplateFingerprint operations."""
def __init__(self, db: Session) -> None:
super().__init__(db, TemplateFingerprint)
def get_by_format_id(self, format_id: uuid.UUID) -> TemplateFingerprint | None:
"""Get fingerprint by template format ID."""
query = select(TemplateFingerprint).where(TemplateFingerprint.format_id == format_id)
result = self.db.execute(query)
return result.scalars().first()
def get_by_hash(self, fingerprint_hash: str) -> TemplateFingerprint | None:
"""Get fingerprint by hash."""
query = select(TemplateFingerprint).where(
TemplateFingerprint.fingerprint_hash == fingerprint_hash
)
result = self.db.execute(query)
return result.scalars().first()
def get_all_fingerprints(self) -> list[TemplateFingerprint]:
"""Get all fingerprints."""
query = select(TemplateFingerprint)
result = self.db.execute(query)
return list(result.scalars().all())
def create_fingerprint(
self,
format_id: uuid.UUID,
fingerprint_hash: str,
page_dimensions: dict | None = None,
logo_coordinates: dict | None = None,
header_coordinates: dict | None = None,
footer_coordinates: dict | None = None,
table_coordinates: dict | None = None,
cell_coordinates: dict | None = None,
) -> TemplateFingerprint:
"""Create a new template fingerprint."""
fp = TemplateFingerprint(
format_id=format_id,
fingerprint_hash=fingerprint_hash,
page_dimensions=page_dimensions,
logo_coordinates=logo_coordinates,
header_coordinates=header_coordinates,
footer_coordinates=footer_coordinates,
table_coordinates=table_coordinates,
cell_coordinates=cell_coordinates,
)
return self.create(fp)