VQA / database /models.py
shivam-2211's picture
Deploy VQA Backend - production files only
82fea5c verified
Raw
History Blame Contribute Delete
10.9 kB
"""
VQA + Image Search Pipeline - Database Models
==============================================
Pydantic models for MongoDB documents and API schemas.
Following the database-schema-designer skill guidelines.
"""
from pydantic import BaseModel, Field, ConfigDict, model_validator
from typing import List, Optional, Dict, Any
from datetime import datetime, timezone
from uuid import uuid4
from enum import Enum
# =============================================================================
# ENUMS
# =============================================================================
class ImageStatus(str, Enum):
"""Status of image processing."""
PENDING = "pending"
PROCESSING = "processing"
COMPLETED = "completed"
FAILED = "failed"
# =============================================================================
# EMBEDDED MODELS
# =============================================================================
class BoundingBox(BaseModel):
"""Bounding box coordinates (normalized 0-1)."""
x_min: float = Field(..., ge=0, le=1, description="Left edge")
y_min: float = Field(..., ge=0, le=1, description="Top edge")
x_max: float = Field(..., ge=0, le=1, description="Right edge")
y_max: float = Field(..., ge=0, le=1, description="Bottom edge")
@model_validator(mode="after")
def validate_bounds(self) -> "BoundingBox":
"""Ensure min values are less than max values."""
if self.x_min >= self.x_max:
raise ValueError(f"x_min ({self.x_min}) must be less than x_max ({self.x_max})")
if self.y_min >= self.y_max:
raise ValueError(f"y_min ({self.y_min}) must be less than y_max ({self.y_max})")
return self
class DetectedObject(BaseModel):
"""An object detected in the image."""
label: str = Field(..., min_length=1, description="Object label/class")
confidence: float = Field(..., ge=0, le=1, description="Detection confidence")
bounding_box: Optional[BoundingBox] = Field(
default=None,
description="Object location in image"
)
class SceneTag(BaseModel):
"""A scene/environment tag for the image."""
label: str = Field(..., min_length=1, description="Scene label")
confidence: float = Field(..., ge=0, le=1, description="Classification confidence")
class ConfidenceScores(BaseModel):
"""Per-field confidence scores for quality tracking."""
objects: Optional[float] = Field(default=None, ge=0, le=1)
scene_tags: Optional[float] = Field(default=None, ge=0, le=1)
caption: Optional[float] = Field(default=None, ge=0, le=1)
ocr_text: Optional[float] = Field(default=None, ge=0, le=1)
embedding: Optional[float] = Field(default=None, ge=0, le=1)
# =============================================================================
# MAIN DOCUMENT MODEL (MongoDB)
# =============================================================================
class ImageDocument(BaseModel):
"""
Image document stored in MongoDB.
This represents the complete perception data extracted from an image.
The embedding vector is stored separately in Qdrant for efficient similarity search.
"""
# Primary identifiers
image_id: str = Field(
default_factory=lambda: str(uuid4()),
description="Unique image identifier (UUID)"
)
source_uri: str = Field(..., description="Original image location (URL or path)")
user_id: Optional[str] = Field(
default=None,
description="Owner user ID (None for legacy/pre-auth images)"
)
# Perception data
objects: List[DetectedObject] = Field(
default_factory=list,
description="Detected objects with bounding boxes"
)
scene_tags: List[SceneTag] = Field(
default_factory=list,
description="Scene/environment classification tags"
)
caption: Optional[str] = Field(
default=None,
description="AI-generated image description"
)
ocr_text: Optional[str] = Field(
default=None,
description="Extracted text from image (OCR)"
)
# Quality and versioning
model_version: str = Field(
...,
description="Perception pipeline version used"
)
confidence_scores: ConfidenceScores = Field(
default_factory=ConfidenceScores,
description="Per-field confidence metrics"
)
status: ImageStatus = Field(
default=ImageStatus.PENDING,
description="Processing status"
)
error_message: Optional[str] = Field(
default=None,
description="Error message if processing failed"
)
# Timestamps (using timezone-aware UTC)
created_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc),
description="Document creation time (UTC)"
)
updated_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc),
description="Last update time (UTC)"
)
# Additional metadata
file_size_bytes: Optional[int] = Field(
default=None,
ge=0,
description="Original file size"
)
image_width: Optional[int] = Field(default=None, ge=1, description="Image width in pixels")
image_height: Optional[int] = Field(default=None, ge=1, description="Image height in pixels")
mime_type: Optional[str] = Field(default=None, description="Image MIME type")
model_config = ConfigDict(
json_encoders={datetime: lambda v: v.isoformat()}
)
# =============================================================================
# API REQUEST/RESPONSE MODELS
# =============================================================================
class ImageIngestRequest(BaseModel):
"""Request to ingest a new image."""
source_uri: Optional[str] = Field(
default=None,
description="URL to fetch image from (alternative to file upload)"
)
class ImageIngestResponse(BaseModel):
"""Response after image ingestion."""
image_id: str
status: ImageStatus
message: str
class TextSearchRequest(BaseModel):
"""Request for text-based image search."""
query: str = Field(..., min_length=1, description="Search query text")
limit: int = Field(default=10, ge=1, le=100, description="Max results per page")
page: int = Field(default=1, ge=1, description="Page number (1-indexed)")
min_confidence: Optional[float] = Field(
default=None,
ge=0, le=1,
description="Minimum confidence threshold"
)
filters: Optional[Dict[str, Any]] = Field(
default=None,
description="Additional filters (e.g., scene_tags, objects)"
)
class ImageSearchRequest(BaseModel):
"""Request for image similarity search."""
limit: int = Field(default=10, ge=1, le=100, description="Max results")
min_similarity: Optional[float] = Field(
default=None,
ge=0, le=1,
description="Minimum similarity score"
)
class SearchResult(BaseModel):
"""Single search result."""
image_id: str
source_uri: str
score: float = Field(..., ge=0, le=1.01, description="Relevance/similarity score (0-1)")
caption: Optional[str] = None
objects: List[str] = Field(default_factory=list, description="Object labels")
scene_tags: List[str] = Field(default_factory=list, description="Scene labels")
class PaginationInfo(BaseModel):
"""Pagination metadata (shared between list and search endpoints)."""
page: int
limit: int
total: int
total_pages: int
has_next: bool
has_prev: bool
class SearchResponse(BaseModel):
"""Search results response."""
query: str
results: List[SearchResult]
total_count: int
search_time_ms: float
pagination: Optional[PaginationInfo] = None
class VQARequest(BaseModel):
"""Request to ask a question about an image."""
question: str = Field(..., min_length=1, description="Question about the image")
use_stored_context: bool = Field(
default=True,
description="Include stored perception data as context"
)
class VQAResponse(BaseModel):
"""VQA response."""
image_id: str
question: str
answer: str
confidence: Optional[float] = Field(
default=None,
description="Answer confidence score"
)
processing_time_ms: float
class QueryRequest(BaseModel):
"""Unified query request (routed to search or VQA)."""
query: str = Field(..., min_length=1, description="Natural language query")
image_id: Optional[str] = Field(
default=None,
description="Specific image ID for VQA"
)
class QueryResponse(BaseModel):
"""Unified query response."""
query_type: str = Field(..., description="'search' or 'vqa'")
search_results: Optional[List[SearchResult]] = None
vqa_answer: Optional[str] = None
processing_time_ms: float
# =============================================================================
# USER / AUTH MODELS
# =============================================================================
class UserDocument(BaseModel):
"""User document stored in MongoDB."""
user_id: str = Field(
default_factory=lambda: str(uuid4()),
description="Unique user identifier (UUID)"
)
email: str = Field(..., description="User email (unique)")
hashed_password: str = Field(..., description="bcrypt hashed password")
is_admin: bool = Field(default=False, description="Admin flag")
created_at: datetime = Field(
default_factory=lambda: datetime.now(timezone.utc),
)
model_config = ConfigDict(
json_encoders={datetime: lambda v: v.isoformat()}
)
class RegisterRequest(BaseModel):
"""User registration request."""
email: str = Field(..., min_length=3, description="Email address")
password: str = Field(..., min_length=6, description="Password (min 6 chars)")
class LoginRequest(BaseModel):
"""User login request."""
email: str = Field(..., description="Email address")
password: str = Field(..., description="Password")
class RefreshRequest(BaseModel):
"""Token refresh request."""
refresh_token: str = Field(..., description="Refresh token")
class TokenResponse(BaseModel):
"""JWT token pair response."""
access_token: str
refresh_token: str
token_type: str = "bearer"
expires_in: int = Field(description="Access token lifetime in seconds")
class UserResponse(BaseModel):
"""Public user profile."""
user_id: str
email: str
is_admin: bool
created_at: datetime