354 lines
13 KiB
Python
354 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,
|
|
) -> 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,
|
|
)
|
|
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)
|