Spaces:
Sleeping
Sleeping
| """ | |
| 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") | |
| 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 | |