tonigi's picture
Add complete Gradio interface
3c2cd23 verified
Raw History Blame Contribute Delete
6.91 kB
from __future__ import annotations
from enum import StrEnum
from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator
Pixel = Annotated[float, Field(ge=0)]
DEFAULT_SNAP_ANGLES = (0.0, 45.0, 90.0, -45.0, -90.0)
class CleanupMode(StrEnum):
OPENCV = "opencv"
LAMA = "lama"
class Device(StrEnum):
AUTO = "auto"
CPU = "cpu"
CUDA = "cuda"
MPS = "mps"
class BitmapExtractor(StrEnum):
OPENCV = "opencv"
SAM2 = "sam2"
class Point(BaseModel):
x: Pixel
y: Pixel
class Box(BaseModel):
x: Pixel
y: Pixel
width: Annotated[float, Field(gt=0)]
height: Annotated[float, Field(gt=0)]
class TextStyle(BaseModel):
font_family: str = "Lato"
font_size: Annotated[float, Field(ge=4, le=1000)] = 16
font_weight: Literal[400, 700] = 400
italic: bool = False
fill: str = "#111111"
opacity: Annotated[float, Field(ge=0, le=1)] = 1
stroke: str = "none"
stroke_width: Annotated[float, Field(ge=0, le=50)] = 0
letter_spacing: Annotated[float, Field(ge=-20, le=100)] = 0
anchor: Literal["start", "middle", "end"] = "start"
@field_validator("fill", "stroke")
@classmethod
def validate_color(cls, value: str) -> str:
if value == "none" or (
len(value) in {4, 7, 9} and value.startswith("#")
):
return value
raise ValueError("expected a CSS hexadecimal color or 'none'")
class TextRegion(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
text: str
confidence: Annotated[float, Field(ge=0, le=1)]
source_quad: Annotated[list[Point], Field(min_length=4, max_length=4)]
overlay_box: Box
rotation: Annotated[float, Field(ge=-180, le=180)] = 0
remove_text: bool = True
include_overlay: bool = True
mask_padding: Annotated[float, Field(ge=0, le=100)] = 3
style: TextStyle = Field(default_factory=TextStyle)
z_index: int = 1000
class SolidPaint(BaseModel):
model_config = ConfigDict(extra="forbid")
type: Literal["solid"] = "solid"
color: str
@field_validator("color")
@classmethod
def validate_color(cls, value: str) -> str:
if len(value) in {4, 7, 9} and value.startswith("#"):
return value
raise ValueError("expected a CSS hexadecimal color")
class GradientStop(BaseModel):
offset: Annotated[float, Field(ge=0, le=1)]
color: str
opacity: Annotated[float, Field(ge=0, le=1)] = 1
@field_validator("color")
@classmethod
def validate_color(cls, value: str) -> str:
if len(value) in {4, 7, 9} and value.startswith("#"):
return value
raise ValueError("expected a CSS hexadecimal color")
class LinearGradientPaint(BaseModel):
model_config = ConfigDict(extra="forbid")
type: Literal["linear-gradient"] = "linear-gradient"
angle: Annotated[float, Field(ge=-180, le=180)] = 0
stops: Annotated[list[GradientStop], Field(min_length=2, max_length=8)]
Paint = SolidPaint | LinearGradientPaint
class ShapeStyle(BaseModel):
fill: Paint = Field(default_factory=lambda: SolidPaint(color="#000000"))
fill_opacity: Annotated[float, Field(ge=0, le=1)] = 1
stroke: str = "none"
stroke_width: Annotated[float, Field(ge=0, le=100)] = 0
stroke_opacity: Annotated[float, Field(ge=0, le=1)] = 1
linecap: Literal["butt", "round", "square"] = "round"
linejoin: Literal["miter", "round", "bevel"] = "round"
@field_validator("stroke")
@classmethod
def validate_stroke(cls, value: str) -> str:
if value == "none" or (len(value) in {4, 7, 9} and value.startswith("#")):
return value
raise ValueError("expected a CSS hexadecimal color or 'none'")
class RectGeometry(BaseModel):
type: Literal["rect"] = "rect"
box: Box
rx: Annotated[float, Field(ge=0)] = 0
ry: Annotated[float, Field(ge=0)] = 0
rotation: Annotated[float, Field(ge=-180, le=180)] = 0
class EllipseGeometry(BaseModel):
type: Literal["ellipse"] = "ellipse"
cx: Pixel
cy: Pixel
rx: Annotated[float, Field(gt=0)]
ry: Annotated[float, Field(gt=0)]
rotation: Annotated[float, Field(ge=-180, le=180)] = 0
class PolygonGeometry(BaseModel):
type: Literal["polygon"] = "polygon"
points: Annotated[list[Point], Field(min_length=3)]
class PathGeometry(BaseModel):
type: Literal["path"] = "path"
points: Annotated[list[Point], Field(min_length=2)]
closed: bool = False
smooth: bool = False
arrow_start: bool = False
arrow_end: bool = False
ShapeGeometry = RectGeometry | EllipseGeometry | PolygonGeometry | PathGeometry
class ShapeRegion(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
confidence: Annotated[float, Field(ge=0, le=1)]
geometry: ShapeGeometry
source_contours: list[list[Point]] = Field(default_factory=list)
style: ShapeStyle = Field(default_factory=ShapeStyle)
remove_source: bool = True
include_overlay: bool = True
mask_padding: Annotated[float, Field(ge=0, le=100)] = 1
z_index: int = 100
is_container: bool = False
semantic_role: Literal["shape", "arrow"] = "shape"
class BitmapRegion(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
asset_name: str
box: Box
opacity: Annotated[float, Field(ge=0, le=1)] = 1
z_index: int = 500
class MaskStroke(BaseModel):
operation: Literal["add", "erase"]
radius: Annotated[float, Field(gt=0, le=500)]
points: Annotated[list[Point], Field(min_length=1)]
class Canvas(BaseModel):
width: Annotated[int, Field(gt=0)]
height: Annotated[int, Field(gt=0)]
class DocumentSpec(BaseModel):
model_config = ConfigDict(extra="forbid")
version: Literal[1, 2] = 1
source_name: str
canvas: Canvas
regions: list[TextRegion]
shapes: list[ShapeRegion] = Field(default_factory=list)
bitmap_regions: list[BitmapRegion] = Field(default_factory=list)
mask_strokes: list[MaskStroke] = Field(default_factory=list)
class ProcessOptions(BaseModel):
confidence: Annotated[float, Field(ge=0, le=1)] = 0.5
device: Device = Device.AUTO
cleanup: CleanupMode = CleanupMode.OPENCV
max_pixels: Annotated[int, Field(gt=0)] = 40_000_000
snap_angles: Annotated[
list[Annotated[float, Field(ge=-180, le=180)]],
Field(min_length=1),
] = Field(default_factory=lambda: list(DEFAULT_SNAP_ANGLES))
snap_tolerance: Annotated[float, Field(ge=0, le=180)] = 6.0
snap_font_sizes: bool = False
vectorize_shapes: bool = False
bitmap_extractor: BitmapExtractor = BitmapExtractor.OPENCV
class JobResponse(BaseModel):
id: str
document: DocumentSpec
original_url: str
background_url: str
class Capabilities(BaseModel):
ocr_available: bool
lama_installed: bool
devices: list[str]
fonts: list[str]