Add comprehensive test suite for image processing and related services
- Implement tests for database generator to ensure proper session handling. - Create tests for EXIF extraction and conversion functions. - Add tests for image-related endpoints, ensuring proper data retrieval and isolation between clients. - Develop tests for OCR functionality, including language detection and text extraction. - Introduce tests for the image processing pipeline, covering success and failure scenarios. - Validate rate limiting functionality and ensure independent counters for different clients. - Implement scraper tests to verify HTML content fetching and error handling. - Add unit tests for various services, including storage and filename generation. - Establish worker entry point for ARQ to handle background image processing tasks.
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# Imago
|
||||
+121
@@ -0,0 +1,121 @@
|
||||
"""
|
||||
Configuration centralisée — chargée depuis .env
|
||||
"""
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
from pydantic_settings import BaseSettings
|
||||
from pydantic import field_validator
|
||||
import json
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
# Application
|
||||
APP_NAME: str = "Imago"
|
||||
APP_VERSION: str = "1.0.0"
|
||||
DEBUG: bool = False
|
||||
SECRET_KEY: str = "changez-moi"
|
||||
|
||||
# Serveur
|
||||
HOST: str = "0.0.0.0"
|
||||
PORT: int = 8000
|
||||
|
||||
# Base de données
|
||||
DATABASE_URL: str = "sqlite+aiosqlite:///./data/imago.db"
|
||||
|
||||
# Stockage
|
||||
UPLOAD_DIR: str = "./data/uploads"
|
||||
THUMBNAILS_DIR: str = "./data/thumbnails"
|
||||
MAX_UPLOAD_SIZE_MB: int = 50
|
||||
|
||||
# AI — Configuration
|
||||
AI_ENABLED: bool = True
|
||||
AI_PROVIDER: str = "openrouter"
|
||||
|
||||
# AI — Google Gemini
|
||||
GEMINI_API_KEY: str = ""
|
||||
GEMINI_MODEL: str = "gemini-3.1-pro-preview"
|
||||
GEMINI_MAX_TOKENS: int = 1024
|
||||
|
||||
# AI — OpenRouter
|
||||
OPENROUTER_API_KEY: str = ""
|
||||
OPENROUTER_MODEL: str = "qwen/qwen2.5-vl-72b-instruct"
|
||||
|
||||
# AI — Comportement
|
||||
AI_TAGS_MIN: int = 5
|
||||
AI_TAGS_MAX: int = 10
|
||||
AI_DESCRIPTION_LANGUAGE: str = "français"
|
||||
AI_CACHE_DAYS: int = 30
|
||||
|
||||
# OCR
|
||||
OCR_ENABLED: bool = True
|
||||
TESSERACT_CMD: str = "/usr/bin/tesseract"
|
||||
OCR_LANGUAGES: str = "fra+eng"
|
||||
|
||||
# CORS
|
||||
CORS_ORIGINS: List[str] = ["http://localhost:3000", "http://localhost:5173"]
|
||||
|
||||
# Authentification
|
||||
ADMIN_API_KEY: str = ""
|
||||
JWT_SECRET_KEY: str = "changez-moi-jwt-secret"
|
||||
JWT_ALGORITHM: str = "HS256"
|
||||
|
||||
# Rate limiting — global (legacy)
|
||||
RATE_LIMIT_UPLOAD: int = 10
|
||||
RATE_LIMIT_AI: int = 20
|
||||
|
||||
# Rate limiting — par plan (requêtes/heure)
|
||||
RATE_LIMIT_FREE_UPLOAD: int = 20
|
||||
RATE_LIMIT_FREE_AI: int = 50
|
||||
RATE_LIMIT_STANDARD_UPLOAD: int = 100
|
||||
RATE_LIMIT_STANDARD_AI: int = 200
|
||||
RATE_LIMIT_PREMIUM_UPLOAD: int = 500
|
||||
RATE_LIMIT_PREMIUM_AI: int = 1000
|
||||
|
||||
# Redis + ARQ Worker
|
||||
REDIS_URL: str = "redis://localhost:6379"
|
||||
WORKER_MAX_JOBS: int = 10
|
||||
WORKER_JOB_TIMEOUT: int = 180
|
||||
WORKER_MAX_TRIES: int = 3
|
||||
AI_STEP_TIMEOUT: int = 120
|
||||
OCR_STEP_TIMEOUT: int = 30
|
||||
|
||||
# Storage Backend
|
||||
STORAGE_BACKEND: str = "local" # "local" | "s3"
|
||||
S3_BUCKET: str = ""
|
||||
S3_REGION: str = "us-east-1"
|
||||
S3_ENDPOINT_URL: str = "" # vide = AWS, sinon MinIO/R2
|
||||
S3_ACCESS_KEY: str = ""
|
||||
S3_SECRET_KEY: str = ""
|
||||
S3_PREFIX: str = "imago"
|
||||
SIGNED_URL_SECRET: str = "changez-moi-signed-url"
|
||||
|
||||
@field_validator("CORS_ORIGINS", mode="before")
|
||||
@classmethod
|
||||
def parse_cors(cls, v):
|
||||
if isinstance(v, str):
|
||||
try:
|
||||
return json.loads(v)
|
||||
except Exception:
|
||||
return [v]
|
||||
return v
|
||||
|
||||
@property
|
||||
def upload_path(self) -> Path:
|
||||
p = Path(self.UPLOAD_DIR)
|
||||
p.mkdir(parents=True, exist_ok=True)
|
||||
return p
|
||||
|
||||
@property
|
||||
def thumbnails_path(self) -> Path:
|
||||
p = Path(self.THUMBNAILS_DIR)
|
||||
p.mkdir(parents=True, exist_ok=True)
|
||||
return p
|
||||
|
||||
@property
|
||||
def max_upload_bytes(self) -> int:
|
||||
return self.MAX_UPLOAD_SIZE_MB * 1024 * 1024
|
||||
|
||||
model_config = {"env_file": ".env", "case_sensitive": True, "extra": "ignore"}
|
||||
|
||||
|
||||
settings = Settings()
|
||||
@@ -0,0 +1,77 @@
|
||||
"""
|
||||
Configuration SQLAlchemy — session async
|
||||
"""
|
||||
import logging
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
engine = create_async_engine(
|
||||
settings.DATABASE_URL,
|
||||
echo=settings.DEBUG,
|
||||
future=True,
|
||||
)
|
||||
|
||||
AsyncSessionLocal = async_sessionmaker(
|
||||
bind=engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
autoflush=False,
|
||||
autocommit=False,
|
||||
)
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
async def get_db() -> AsyncSession:
|
||||
"""Dependency FastAPI — injecte une session DB dans chaque requête."""
|
||||
async with AsyncSessionLocal() as session:
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
finally:
|
||||
await session.close()
|
||||
|
||||
|
||||
async def init_db():
|
||||
"""Crée toutes les tables et initialise un client par défaut si nécessaire."""
|
||||
import secrets
|
||||
import hashlib
|
||||
from sqlalchemy import select
|
||||
from app.models.client import APIClient, ClientPlan
|
||||
|
||||
async with engine.begin() as conn:
|
||||
from app.models import image # noqa: F401
|
||||
from app.models import client # noqa: F401
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
# Vérifier s'il y a déjà des clients
|
||||
async with AsyncSessionLocal() as session:
|
||||
result = await session.execute(select(APIClient).limit(1))
|
||||
if result.scalar_one_or_none() is None:
|
||||
# Table vide -> Création du client bootstrap
|
||||
raw_key = secrets.token_urlsafe(32)
|
||||
key_hash = hashlib.sha256(raw_key.encode("utf-8")).hexdigest()
|
||||
|
||||
bootstrap_client = APIClient(
|
||||
name="Default Admin",
|
||||
api_key_hash=key_hash,
|
||||
scopes=["images:read", "images:write", "images:delete", "ai:use", "admin"],
|
||||
plan=ClientPlan.PREMIUM,
|
||||
)
|
||||
session.add(bootstrap_client)
|
||||
await session.commit()
|
||||
|
||||
logger.info("bootstrap.client_created", extra={
|
||||
"client_id": bootstrap_client.id,
|
||||
"api_key": raw_key,
|
||||
"warning": "Notez cette clé ! Elle ne sera plus affichée.",
|
||||
})
|
||||
@@ -0,0 +1,3 @@
|
||||
from app.dependencies.auth import get_current_client, require_scope, verify_api_key
|
||||
|
||||
__all__ = ["get_current_client", "require_scope", "verify_api_key"]
|
||||
@@ -0,0 +1,115 @@
|
||||
"""
|
||||
Dépendances FastAPI — authentification par API Key + vérification de scopes.
|
||||
|
||||
Usage dans les routers :
|
||||
client = Depends(get_current_client)
|
||||
_ = Depends(require_scope("images:read"))
|
||||
"""
|
||||
import hashlib
|
||||
import logging
|
||||
from typing import Callable
|
||||
|
||||
from fastapi import Depends, Header, HTTPException, Request, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.models.client import APIClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def hash_api_key(api_key: str) -> str:
|
||||
"""Hash SHA-256 d'une clé API — fonction utilitaire réutilisable."""
|
||||
return hashlib.sha256(api_key.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def verify_api_key(
|
||||
request: Request,
|
||||
authorization: str = Header(
|
||||
...,
|
||||
alias="Authorization",
|
||||
description="Clé API au format 'Bearer <key>'",
|
||||
),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> APIClient:
|
||||
"""
|
||||
Vérifie la clé API fournie dans le header Authorization.
|
||||
Injecte client_id et client_plan dans request.state pour le rate limiter.
|
||||
|
||||
Raises:
|
||||
HTTPException 401: clé absente, invalide ou client inactif.
|
||||
"""
|
||||
# ── Extraction du token ───────────────────────────────────
|
||||
if not authorization.startswith("Bearer "):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentification requise",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
raw_key = authorization[7:] # strip "Bearer "
|
||||
if not raw_key:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentification requise",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
# ── Lookup par hash ───────────────────────────────────────
|
||||
key_hash = hash_api_key(raw_key)
|
||||
result = await db.execute(
|
||||
select(APIClient).where(APIClient.api_key_hash == key_hash)
|
||||
)
|
||||
client = result.scalar_one_or_none()
|
||||
|
||||
if client is None:
|
||||
logger.warning("Tentative d'authentification avec une clé invalide")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentification requise",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
if not client.is_active:
|
||||
logger.warning("Tentative d'authentification avec un client inactif: %s", client.id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentification requise",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
# Injecter dans request.state pour le rate limiter
|
||||
request.state.client_id = client.id
|
||||
request.state.client_plan = client.plan.value if client.plan else "free"
|
||||
|
||||
return client
|
||||
|
||||
|
||||
# Alias pratique pour injection dans les routers
|
||||
get_current_client = verify_api_key
|
||||
|
||||
|
||||
def require_scope(scope: str) -> Callable:
|
||||
"""
|
||||
Factory qui retourne une dépendance FastAPI vérifiant qu'un scope est accordé.
|
||||
|
||||
Usage:
|
||||
@router.get("/...", dependencies=[Depends(require_scope("images:read"))])
|
||||
"""
|
||||
|
||||
async def _check_scope(
|
||||
client: APIClient = Depends(get_current_client),
|
||||
) -> APIClient:
|
||||
if not client.has_scope(scope):
|
||||
logger.warning(
|
||||
"Client %s (%s) a tenté d'accéder au scope '%s' sans autorisation",
|
||||
client.id, client.name, scope,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Permission insuffisante",
|
||||
)
|
||||
return client
|
||||
|
||||
return _check_scope
|
||||
@@ -0,0 +1,58 @@
|
||||
"""
|
||||
Configuration structlog — logging structuré JSON/Console.
|
||||
|
||||
En production (DEBUG=False) : JSON pour les agrégateurs (ELK, Datadog, etc.)
|
||||
En développement (DEBUG=True) : Console colorée lisible.
|
||||
"""
|
||||
import logging
|
||||
import sys
|
||||
|
||||
import structlog
|
||||
|
||||
|
||||
def configure_logging(debug: bool = False) -> None:
|
||||
"""Configure structlog + stdlib logging."""
|
||||
shared_processors = [
|
||||
structlog.contextvars.merge_contextvars,
|
||||
structlog.stdlib.add_log_level,
|
||||
structlog.stdlib.add_logger_name,
|
||||
structlog.processors.TimeStamper(fmt="iso"),
|
||||
structlog.processors.StackInfoRenderer(),
|
||||
structlog.processors.UnicodeDecoder(),
|
||||
]
|
||||
|
||||
if debug:
|
||||
# Console lisible en dev
|
||||
renderer = structlog.dev.ConsoleRenderer(colors=True)
|
||||
else:
|
||||
# JSON en production
|
||||
renderer = structlog.processors.JSONRenderer()
|
||||
|
||||
structlog.configure(
|
||||
processors=[
|
||||
*shared_processors,
|
||||
structlog.stdlib.ProcessorFormatter.wrap_for_formatter,
|
||||
],
|
||||
logger_factory=structlog.stdlib.LoggerFactory(),
|
||||
wrapper_class=structlog.stdlib.BoundLogger,
|
||||
cache_logger_on_first_use=True,
|
||||
)
|
||||
|
||||
formatter = structlog.stdlib.ProcessorFormatter(
|
||||
processor=renderer,
|
||||
foreign_pre_chain=shared_processors,
|
||||
)
|
||||
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setFormatter(formatter)
|
||||
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.handlers.clear()
|
||||
root_logger.addHandler(handler)
|
||||
root_logger.setLevel(logging.DEBUG if debug else logging.INFO)
|
||||
|
||||
# Réduire le bruit des librairies tierces
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
logging.getLogger("sqlalchemy.engine").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
+310
@@ -0,0 +1,310 @@
|
||||
"""
|
||||
Imago — Application principale FastAPI
|
||||
"""
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from slowapi import _rate_limit_exceeded_handler
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
from app.config import settings
|
||||
from app.database import init_db
|
||||
from app.logging_config import configure_logging
|
||||
from app.routers import images_router, ai_router, auth_router, files_router
|
||||
from app.middleware import limiter
|
||||
from app.middleware.logging_middleware import LoggingMiddleware
|
||||
from app.workers.redis_client import get_redis_pool, close_redis_pool
|
||||
|
||||
# Configure le logging structuré dès l'import
|
||||
configure_logging(debug=settings.DEBUG)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
from arq import create_pool
|
||||
from arq.connections import RedisSettings
|
||||
_arq_available = True
|
||||
except ImportError:
|
||||
_arq_available = False
|
||||
|
||||
try:
|
||||
from prometheus_fastapi_instrumentator import Instrumentator
|
||||
_prometheus_available = True
|
||||
except ImportError:
|
||||
_prometheus_available = False
|
||||
|
||||
|
||||
def _arq_redis_settings() -> "RedisSettings":
|
||||
"""Parse REDIS_URL en RedisSettings ARQ."""
|
||||
url = settings.REDIS_URL
|
||||
if url.startswith("redis://"):
|
||||
url = url[8:]
|
||||
elif url.startswith("rediss://"):
|
||||
url = url[9:]
|
||||
|
||||
password = None
|
||||
host = "localhost"
|
||||
port = 6379
|
||||
database = 0
|
||||
|
||||
if "@" in url:
|
||||
auth_part, url = url.rsplit("@", 1)
|
||||
if ":" in auth_part:
|
||||
password = auth_part.split(":", 1)[1]
|
||||
else:
|
||||
password = auth_part
|
||||
|
||||
if "/" in url:
|
||||
host_port, db_str = url.split("/", 1)
|
||||
if db_str:
|
||||
database = int(db_str)
|
||||
else:
|
||||
host_port = url
|
||||
|
||||
if ":" in host_port:
|
||||
host, port_str = host_port.rsplit(":", 1)
|
||||
if port_str:
|
||||
port = int(port_str)
|
||||
else:
|
||||
host = host_port
|
||||
|
||||
return RedisSettings(
|
||||
host=host or "localhost",
|
||||
port=port,
|
||||
password=password,
|
||||
database=database,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Lifespan — initialisation au démarrage
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
# Création des répertoires de données
|
||||
settings.upload_path
|
||||
settings.thumbnails_path
|
||||
|
||||
# Initialisation de la base de données (création des tables)
|
||||
await init_db()
|
||||
|
||||
active_model = settings.OPENROUTER_MODEL if settings.AI_PROVIDER == "openrouter" else settings.GEMINI_MODEL
|
||||
|
||||
logger.info("startup.db_initialized", extra={"upload_dir": settings.UPLOAD_DIR})
|
||||
logger.info("startup.ai_config", extra={
|
||||
"provider": settings.AI_PROVIDER,
|
||||
"model": active_model,
|
||||
"ocr_enabled": settings.OCR_ENABLED,
|
||||
})
|
||||
|
||||
# Initialisation Redis + ARQ pool
|
||||
try:
|
||||
app.state.redis = await get_redis_pool()
|
||||
logger.info("startup.redis_connected", extra={"url": settings.REDIS_URL})
|
||||
except Exception as e:
|
||||
app.state.redis = None
|
||||
logger.warning("startup.redis_unavailable", extra={"error": str(e)})
|
||||
|
||||
if _arq_available:
|
||||
try:
|
||||
app.state.arq_pool = await create_pool(_arq_redis_settings())
|
||||
logger.info("startup.arq_pool_created")
|
||||
except Exception as e:
|
||||
app.state.arq_pool = _FallbackArqPool()
|
||||
logger.warning("startup.arq_fallback", extra={"error": str(e)})
|
||||
else:
|
||||
app.state.arq_pool = _FallbackArqPool()
|
||||
logger.warning("startup.arq_not_installed")
|
||||
|
||||
yield
|
||||
|
||||
# Fermeture propre
|
||||
if hasattr(app.state, "arq_pool") and hasattr(app.state.arq_pool, "close"):
|
||||
await app.state.arq_pool.close()
|
||||
await close_redis_pool()
|
||||
logger.info("shutdown.complete")
|
||||
|
||||
|
||||
class _FallbackArqPool:
|
||||
"""Fallback quand Redis/ARQ n'est pas disponible."""
|
||||
|
||||
async def enqueue_job(self, *args, **kwargs):
|
||||
logger.warning("arq.fallback_enqueue", extra={"args": str(args)})
|
||||
return None
|
||||
|
||||
async def close(self):
|
||||
pass
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Application
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
app = FastAPI(
|
||||
title=settings.APP_NAME,
|
||||
version=settings.APP_VERSION,
|
||||
description="""
|
||||
## Imago
|
||||
|
||||
Backend de gestion d'images et fonctionnalités AI pour l'interface Shaarli.
|
||||
|
||||
### Fonctionnalités
|
||||
- 📸 **Upload et stockage d'images** avec génération de thumbnails
|
||||
- 🔍 **Extraction EXIF** automatique (appareil, GPS, paramètres de prise de vue)
|
||||
- 📝 **OCR** — extraction de texte depuis les images (Tesseract)
|
||||
- 🤖 **Vision AI** — description et classification par tags (Gemini)
|
||||
- 🔗 **Résumé d'URL** — scraping + résumé AI de pages web
|
||||
- ✅ **Rédaction de tâches** — génération structurée via AI
|
||||
- 📋 **File de tâches ARQ** — pipeline persistant avec retry automatique
|
||||
- 📊 **Métriques Prometheus** — /metrics endpoint
|
||||
|
||||
### Pipeline de traitement
|
||||
Chaque image uploadée est automatiquement traitée via ARQ (Redis) :
|
||||
`EXIF → OCR → Vision AI → stockage BDD`
|
||||
""",
|
||||
lifespan=lifespan,
|
||||
docs_url="/docs",
|
||||
redoc_url="/redoc",
|
||||
)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Middleware CORS
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=settings.CORS_ORIGINS,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Middleware Logging HTTP
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Rate Limiting (slowapi)
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Prometheus Metrics
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
if _prometheus_available:
|
||||
Instrumentator(
|
||||
should_group_status_codes=True,
|
||||
should_ignore_untemplated=True,
|
||||
excluded_handlers=["/health", "/health/detailed", "/metrics"],
|
||||
).instrument(app).expose(app, endpoint="/metrics", tags=["Observabilité"])
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Fichiers statiques / URLs signées
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
if settings.STORAGE_BACKEND == "local":
|
||||
app.include_router(files_router)
|
||||
app.mount("/static/uploads", StaticFiles(directory=str(settings.upload_path)), name="uploads")
|
||||
app.mount("/static/thumbnails", StaticFiles(directory=str(settings.thumbnails_path)), name="thumbnails")
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Routers
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
app.include_router(images_router)
|
||||
app.include_router(ai_router)
|
||||
app.include_router(auth_router)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Routes utilitaires
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@app.get("/", tags=["Santé"])
|
||||
async def root():
|
||||
return {
|
||||
"app": settings.APP_NAME,
|
||||
"version": settings.APP_VERSION,
|
||||
"status": "running",
|
||||
"docs": "/docs",
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health", tags=["Santé"])
|
||||
async def health():
|
||||
ai_configured = (
|
||||
(settings.AI_PROVIDER == "gemini" and bool(settings.GEMINI_API_KEY)) or
|
||||
(settings.AI_PROVIDER == "openrouter" and bool(settings.OPENROUTER_API_KEY))
|
||||
)
|
||||
active_model = settings.OPENROUTER_MODEL if settings.AI_PROVIDER == "openrouter" else settings.GEMINI_MODEL
|
||||
|
||||
return {
|
||||
"status": "healthy",
|
||||
"ai_enabled": settings.AI_ENABLED,
|
||||
"ai_provider": settings.AI_PROVIDER,
|
||||
"ai_configured": ai_configured,
|
||||
"ocr_enabled": settings.OCR_ENABLED,
|
||||
"model": active_model,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/health/detailed", tags=["Santé"])
|
||||
async def health_detailed(request: Request):
|
||||
"""Endpoint de santé détaillé pour monitoring avancé."""
|
||||
checks = {}
|
||||
|
||||
# DB check
|
||||
try:
|
||||
from app.database import AsyncSessionLocal
|
||||
from sqlalchemy import text
|
||||
async with AsyncSessionLocal() as session:
|
||||
await session.execute(text("SELECT 1"))
|
||||
checks["database"] = {"status": "ok"}
|
||||
except Exception as e:
|
||||
checks["database"] = {"status": "error", "error": str(e)}
|
||||
|
||||
# Redis check
|
||||
redis = getattr(request.app.state, "redis", None)
|
||||
if redis:
|
||||
try:
|
||||
await redis.ping()
|
||||
checks["redis"] = {"status": "ok"}
|
||||
except Exception as e:
|
||||
checks["redis"] = {"status": "error", "error": str(e)}
|
||||
else:
|
||||
checks["redis"] = {"status": "not_configured"}
|
||||
|
||||
# ARQ check
|
||||
arq_pool = getattr(request.app.state, "arq_pool", None)
|
||||
if arq_pool and not isinstance(arq_pool, _FallbackArqPool):
|
||||
checks["arq"] = {"status": "ok"}
|
||||
else:
|
||||
checks["arq"] = {"status": "fallback"}
|
||||
|
||||
# OCR check
|
||||
checks["ocr"] = {"status": "enabled" if settings.OCR_ENABLED else "disabled"}
|
||||
|
||||
# Storage check
|
||||
checks["storage"] = {
|
||||
"backend": settings.STORAGE_BACKEND,
|
||||
"status": "ok",
|
||||
}
|
||||
|
||||
overall = "healthy" if all(
|
||||
c.get("status") in ("ok", "enabled", "disabled", "not_configured", "fallback")
|
||||
for c in checks.values()
|
||||
) else "degraded"
|
||||
|
||||
return {
|
||||
"status": overall,
|
||||
"checks": checks,
|
||||
"version": settings.APP_VERSION,
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
"""
|
||||
Métriques Prometheus custom pour le hub d'images.
|
||||
|
||||
Exposed via /metrics par prometheus-fastapi-instrumentator.
|
||||
"""
|
||||
from prometheus_client import Counter, Histogram, Gauge
|
||||
|
||||
# ── Images ────────────────────────────────────────────────────
|
||||
hub_images_uploaded = Counter(
|
||||
"hub_images_uploaded_total",
|
||||
"Nombre total d'images uploadées",
|
||||
["client_plan"],
|
||||
)
|
||||
|
||||
hub_images_deleted = Counter(
|
||||
"hub_images_deleted_total",
|
||||
"Nombre total d'images supprimées",
|
||||
)
|
||||
|
||||
# ── Pipeline ──────────────────────────────────────────────────
|
||||
hub_pipeline_duration = Histogram(
|
||||
"hub_pipeline_duration_seconds",
|
||||
"Durée du pipeline de traitement complet",
|
||||
buckets=[1, 5, 10, 30, 60, 120, 300],
|
||||
)
|
||||
|
||||
hub_pipeline_step_duration = Histogram(
|
||||
"hub_pipeline_step_duration_seconds",
|
||||
"Durée de chaque étape du pipeline",
|
||||
["step"],
|
||||
buckets=[0.1, 0.5, 1, 5, 10, 30, 60],
|
||||
)
|
||||
|
||||
hub_pipeline_errors = Counter(
|
||||
"hub_pipeline_errors_total",
|
||||
"Nombre d'erreurs pipeline",
|
||||
["step"],
|
||||
)
|
||||
|
||||
# ── Storage ───────────────────────────────────────────────────
|
||||
hub_storage_used_bytes = Gauge(
|
||||
"hub_storage_used_bytes",
|
||||
"Espace de stockage utilisé par client",
|
||||
["client_id"],
|
||||
)
|
||||
|
||||
# ── ARQ ───────────────────────────────────────────────────────
|
||||
hub_arq_jobs_enqueued = Counter(
|
||||
"hub_arq_jobs_enqueued_total",
|
||||
"Nombre de jobs ARQ enfilés",
|
||||
["queue"],
|
||||
)
|
||||
|
||||
hub_arq_jobs_completed = Counter(
|
||||
"hub_arq_jobs_completed_total",
|
||||
"Nombre de jobs ARQ terminés",
|
||||
)
|
||||
|
||||
hub_arq_jobs_failed = Counter(
|
||||
"hub_arq_jobs_failed_total",
|
||||
"Nombre de jobs ARQ échoués",
|
||||
)
|
||||
@@ -0,0 +1,65 @@
|
||||
"""
|
||||
Middleware de rate limiting — par client et par plan.
|
||||
|
||||
Utilise slowapi (basé sur limits) avec identification par client_id
|
||||
plutôt que par IP, pour compter les requêtes par client API.
|
||||
|
||||
Limites par plan (par heure) :
|
||||
- free : 20 uploads, 50 AI
|
||||
- standard : 100 uploads, 200 AI
|
||||
- premium : 500 uploads, 1000 AI
|
||||
"""
|
||||
import logging
|
||||
from slowapi import Limiter
|
||||
from slowapi.util import get_remote_address
|
||||
from starlette.requests import Request
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_client_id_from_request(request: Request) -> str:
|
||||
"""
|
||||
Extrait le client_id depuis la state de la requête.
|
||||
Fallback vers l'IP si le client n'est pas encore authentifié.
|
||||
"""
|
||||
# Le client_id est injecté par le middleware ou la dépendance auth
|
||||
client_id = getattr(request.state, "client_id", None)
|
||||
if client_id:
|
||||
return str(client_id)
|
||||
return get_remote_address(request)
|
||||
|
||||
|
||||
# Instance globale du limiter
|
||||
limiter = Limiter(key_func=_get_client_id_from_request)
|
||||
|
||||
|
||||
def get_upload_rate_limit(plan: str) -> str:
|
||||
"""Retourne la limite de rate pour les uploads selon le plan."""
|
||||
limits = {
|
||||
"free": f"{settings.RATE_LIMIT_FREE_UPLOAD}/hour",
|
||||
"standard": f"{settings.RATE_LIMIT_STANDARD_UPLOAD}/hour",
|
||||
"premium": f"{settings.RATE_LIMIT_PREMIUM_UPLOAD}/hour",
|
||||
}
|
||||
return limits.get(plan, limits["free"])
|
||||
|
||||
|
||||
def get_ai_rate_limit(plan: str) -> str:
|
||||
"""Retourne la limite de rate pour les endpoints AI selon le plan."""
|
||||
limits = {
|
||||
"free": f"{settings.RATE_LIMIT_FREE_AI}/hour",
|
||||
"standard": f"{settings.RATE_LIMIT_STANDARD_AI}/hour",
|
||||
"premium": f"{settings.RATE_LIMIT_PREMIUM_AI}/hour",
|
||||
}
|
||||
return limits.get(plan, limits["free"])
|
||||
|
||||
|
||||
def upload_rate_limit_key(request: Request) -> str:
|
||||
"""Clé dynamique pour le rate limiting des uploads."""
|
||||
return _get_client_id_from_request(request)
|
||||
|
||||
|
||||
def ai_rate_limit_key(request: Request) -> str:
|
||||
"""Clé dynamique pour le rate limiting des endpoints AI."""
|
||||
return _get_client_id_from_request(request)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""
|
||||
Middleware de logging HTTP — enregistre chaque requête avec structlog.
|
||||
|
||||
Exclut les endpoints de santé (/health, /metrics) pour réduire le bruit.
|
||||
"""
|
||||
import time
|
||||
import logging
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
logger = logging.getLogger("http")
|
||||
|
||||
# Chemins exclus du logging
|
||||
_EXCLUDED_PATHS = {"/health", "/health/detailed", "/metrics", "/favicon.ico"}
|
||||
|
||||
|
||||
class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next) -> Response:
|
||||
if request.url.path in _EXCLUDED_PATHS:
|
||||
return await call_next(request)
|
||||
|
||||
start = time.monotonic()
|
||||
client_id = getattr(request.state, "client_id", "anonymous")
|
||||
|
||||
response = await call_next(request)
|
||||
|
||||
duration_ms = int((time.monotonic() - start) * 1000)
|
||||
|
||||
logger.info(
|
||||
"http.request",
|
||||
extra={
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"status": response.status_code,
|
||||
"duration_ms": duration_ms,
|
||||
"client_id": str(client_id),
|
||||
},
|
||||
)
|
||||
|
||||
return response
|
||||
@@ -0,0 +1,3 @@
|
||||
from app.middleware import limiter, get_upload_rate_limit, get_ai_rate_limit
|
||||
|
||||
__all__ = ["limiter", "get_upload_rate_limit", "get_ai_rate_limit"]
|
||||
@@ -0,0 +1,4 @@
|
||||
from app.models.image import Image, ProcessingStatus
|
||||
from app.models.client import APIClient, ClientPlan
|
||||
|
||||
__all__ = ["Image", "ProcessingStatus", "APIClient", "ClientPlan"]
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
Modèle SQLAlchemy — APIClient : clients authentifiés du hub
|
||||
"""
|
||||
import enum
|
||||
import uuid as uuid_lib
|
||||
from datetime import datetime, timezone
|
||||
from sqlalchemy import (
|
||||
Column, String, JSON, Boolean, DateTime, Enum as SAEnum,
|
||||
BigInteger, Integer,
|
||||
)
|
||||
from sqlalchemy.orm import relationship
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class ClientPlan(str, enum.Enum):
|
||||
FREE = "free"
|
||||
STANDARD = "standard"
|
||||
PREMIUM = "premium"
|
||||
|
||||
|
||||
class APIClient(Base):
|
||||
__tablename__ = "api_clients"
|
||||
|
||||
# ── Identité ──────────────────────────────────────────────
|
||||
id = Column(
|
||||
String(36),
|
||||
primary_key=True,
|
||||
default=lambda: str(uuid_lib.uuid4()),
|
||||
index=True,
|
||||
)
|
||||
name = Column(String(256), nullable=False)
|
||||
|
||||
# ── Authentification ──────────────────────────────────────
|
||||
api_key_hash = Column(String(64), nullable=False, unique=True, index=True)
|
||||
|
||||
# ── Permissions ───────────────────────────────────────────
|
||||
scopes = Column(JSON, nullable=False, default=list)
|
||||
plan = Column(
|
||||
SAEnum(ClientPlan),
|
||||
default=ClientPlan.FREE,
|
||||
nullable=False,
|
||||
)
|
||||
is_active = Column(Boolean, default=True, nullable=False)
|
||||
|
||||
# ── Quota tracking ─────────────────────────────────────────
|
||||
storage_used_bytes = Column(BigInteger, default=0, nullable=False)
|
||||
quota_storage_mb = Column(Integer, default=500, nullable=False)
|
||||
quota_images = Column(Integer, default=1000, nullable=False)
|
||||
|
||||
# ── Timestamps ────────────────────────────────────────────
|
||||
created_at = Column(DateTime, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(
|
||||
DateTime,
|
||||
default=lambda: datetime.now(timezone.utc),
|
||||
onupdate=lambda: datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
# ── Relations (ajoutée par Livrable 1.2) ──────────────────
|
||||
images = relationship(
|
||||
"Image",
|
||||
back_populates="client",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<APIClient id={self.id} name={self.name} plan={self.plan}>"
|
||||
|
||||
def has_scope(self, scope: str) -> bool:
|
||||
"""Vérifie si le client possède le scope demandé."""
|
||||
return scope in (self.scopes or [])
|
||||
@@ -0,0 +1,97 @@
|
||||
"""
|
||||
Modèle SQLAlchemy — Image et métadonnées associées
|
||||
"""
|
||||
import enum
|
||||
from datetime import datetime, timezone
|
||||
from sqlalchemy import (
|
||||
Column, Integer, String, Text, DateTime,
|
||||
JSON, Float, Enum as SAEnum, BigInteger, Boolean, ForeignKey
|
||||
)
|
||||
from sqlalchemy.orm import relationship
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class ProcessingStatus(str, enum.Enum):
|
||||
PENDING = "pending"
|
||||
PROCESSING = "processing"
|
||||
DONE = "done"
|
||||
ERROR = "error"
|
||||
|
||||
|
||||
class Image(Base):
|
||||
__tablename__ = "images"
|
||||
|
||||
# ── Identité ──────────────────────────────────────────────
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
uuid = Column(String(36), unique=True, index=True, nullable=False)
|
||||
|
||||
# ── Client (multi-tenant) ─────────────────────────────────
|
||||
client_id = Column(String(36), ForeignKey("api_clients.id"), nullable=False, index=True)
|
||||
client = relationship("APIClient", back_populates="images")
|
||||
|
||||
# ── Fichier ───────────────────────────────────────────────
|
||||
original_name = Column(String(512), nullable=False)
|
||||
filename = Column(String(512), nullable=False) # nom sur disque (uuid-based)
|
||||
file_path = Column(String(1024), nullable=False)
|
||||
thumbnail_path = Column(String(1024))
|
||||
mime_type = Column(String(128))
|
||||
file_size = Column(BigInteger) # bytes
|
||||
width = Column(Integer)
|
||||
height = Column(Integer)
|
||||
uploaded_at = Column(DateTime, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
# ── Statut du pipeline AI ─────────────────────────────────
|
||||
processing_status = Column(
|
||||
SAEnum(ProcessingStatus),
|
||||
default=ProcessingStatus.PENDING,
|
||||
nullable=False,
|
||||
index=True
|
||||
)
|
||||
processing_error = Column(Text)
|
||||
processing_started_at = Column(DateTime)
|
||||
processing_done_at = Column(DateTime)
|
||||
|
||||
# ── Métadonnées EXIF ──────────────────────────────────────
|
||||
exif_raw = Column(JSON) # dict complet brut
|
||||
exif_make = Column(String(256)) # Appareil — fabricant
|
||||
exif_model = Column(String(256)) # Appareil — modèle
|
||||
exif_lens = Column(String(256))
|
||||
exif_taken_at = Column(DateTime) # DateTimeOriginal EXIF
|
||||
exif_gps_lat = Column(Float)
|
||||
exif_gps_lon = Column(Float)
|
||||
exif_altitude = Column(Float)
|
||||
exif_iso = Column(Integer)
|
||||
exif_aperture = Column(String(32)) # ex: "f/2.8"
|
||||
exif_shutter = Column(String(32)) # ex: "1/250"
|
||||
exif_focal = Column(String(32)) # ex: "50mm"
|
||||
exif_flash = Column(Boolean)
|
||||
exif_orientation = Column(Integer)
|
||||
exif_software = Column(String(256))
|
||||
|
||||
# ── OCR ───────────────────────────────────────────────────
|
||||
ocr_text = Column(Text)
|
||||
ocr_language = Column(String(64))
|
||||
ocr_confidence = Column(Float) # 0.0 – 1.0
|
||||
ocr_has_text = Column(Boolean, default=False)
|
||||
|
||||
# ── AI Vision ─────────────────────────────────────────────
|
||||
ai_description = Column(Text)
|
||||
ai_tags = Column(JSON) # ["nature", "paysage", ...]
|
||||
ai_confidence = Column(Float) # score de confiance global
|
||||
ai_model_used = Column(String(128))
|
||||
ai_processed_at = Column(DateTime)
|
||||
ai_prompt_tokens = Column(Integer)
|
||||
ai_output_tokens = Column(Integer)
|
||||
|
||||
def __repr__(self):
|
||||
return f"<Image id={self.id} name={self.original_name} status={self.processing_status}>"
|
||||
|
||||
@property
|
||||
def has_gps(self) -> bool:
|
||||
return self.exif_gps_lat is not None and self.exif_gps_lon is not None
|
||||
|
||||
@property
|
||||
def dimensions(self) -> str | None:
|
||||
if self.width and self.height:
|
||||
return f"{self.width}x{self.height}"
|
||||
return None
|
||||
@@ -0,0 +1,6 @@
|
||||
from app.routers.images import router as images_router
|
||||
from app.routers.ai import router as ai_router
|
||||
from app.routers.auth import router as auth_router
|
||||
from app.routers.files import router as files_router
|
||||
|
||||
__all__ = ["images_router", "ai_router", "auth_router", "files_router"]
|
||||
@@ -0,0 +1,102 @@
|
||||
"""
|
||||
Router — AI : résumé d'URL, rédaction de tâches
|
||||
Sécurisé : authentification par API Key + scope ai:use + rate limiting.
|
||||
"""
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
|
||||
from app.dependencies.auth import get_current_client, require_scope
|
||||
from app.models.client import APIClient
|
||||
from app.schemas import (
|
||||
SummarizeRequest, SummarizeResponse,
|
||||
DraftTaskRequest, DraftTaskResponse,
|
||||
)
|
||||
from app.services.scraper import fetch_page_content
|
||||
from app.services.ai_vision import summarize_url, draft_task
|
||||
from app.config import settings
|
||||
from app.middleware import limiter
|
||||
|
||||
|
||||
router = APIRouter(prefix="/ai", tags=["Intelligence Artificielle"])
|
||||
|
||||
|
||||
@router.post(
|
||||
"/summarize",
|
||||
response_model=SummarizeResponse,
|
||||
summary="Résumé AI d'une URL",
|
||||
description=(
|
||||
"Scrappe le contenu d'une URL et génère un résumé structuré + tags via AI. "
|
||||
"Utile pour enrichir les bookmarks Shaarli."
|
||||
),
|
||||
dependencies=[Depends(require_scope("ai:use"))],
|
||||
)
|
||||
@limiter.limit("1000/hour")
|
||||
async def summarize_link(
|
||||
request: Request,
|
||||
body: SummarizeRequest,
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
if not settings.AI_ENABLED:
|
||||
raise HTTPException(status_code=503, detail="AI désactivée")
|
||||
|
||||
# Scraping
|
||||
page = await fetch_page_content(body.url)
|
||||
if page.get("error"):
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"Impossible de récupérer la page : {page['error']}",
|
||||
)
|
||||
|
||||
content = " ".join(filter(None, [
|
||||
page.get("title"),
|
||||
page.get("description"),
|
||||
page.get("text"),
|
||||
]))
|
||||
|
||||
if not content.strip():
|
||||
raise HTTPException(status_code=422, detail="Aucun contenu texte trouvé sur cette page")
|
||||
|
||||
# Résumé AI
|
||||
result = await summarize_url(
|
||||
url=body.url,
|
||||
content=content,
|
||||
language=body.language,
|
||||
)
|
||||
|
||||
return SummarizeResponse(
|
||||
url=body.url,
|
||||
title=page.get("title"),
|
||||
summary=result["summary"],
|
||||
tags=result["tags"],
|
||||
model=result["model"],
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/draft-task",
|
||||
response_model=DraftTaskResponse,
|
||||
summary="Rédaction AI d'une tâche",
|
||||
description=(
|
||||
"Génère une tâche structurée (titre, description, étapes, priorité) "
|
||||
"à partir d'une description libre."
|
||||
),
|
||||
dependencies=[Depends(require_scope("ai:use"))],
|
||||
)
|
||||
@limiter.limit("1000/hour")
|
||||
async def generate_task(
|
||||
request: Request,
|
||||
body: DraftTaskRequest,
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
if not settings.AI_ENABLED:
|
||||
raise HTTPException(status_code=503, detail="AI désactivée")
|
||||
|
||||
result = await draft_task(
|
||||
description=body.description,
|
||||
context=body.context,
|
||||
language=body.language,
|
||||
)
|
||||
|
||||
if not result.get("title"):
|
||||
raise HTTPException(status_code=500, detail="Échec de la génération de la tâche")
|
||||
|
||||
return DraftTaskResponse(**result)
|
||||
@@ -0,0 +1,198 @@
|
||||
"""
|
||||
Router — Auth : gestion des clients API (CRUD + rotation de clé)
|
||||
"""
|
||||
import logging
|
||||
import secrets
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.dependencies.auth import get_current_client, hash_api_key, require_scope
|
||||
from app.models.client import APIClient
|
||||
from app.schemas.auth import (
|
||||
ClientCreate,
|
||||
ClientCreateResponse,
|
||||
ClientResponse,
|
||||
ClientUpdate,
|
||||
KeyRotateResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["Authentification"])
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# CRÉER UN CLIENT
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.post(
|
||||
"/clients",
|
||||
response_model=ClientCreateResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Créer un nouveau client API",
|
||||
description=(
|
||||
"Crée un client et retourne la clé API **en clair une seule fois**. "
|
||||
"Stockez-la immédiatement — elle ne sera plus jamais affichée."
|
||||
),
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def create_client(
|
||||
body: ClientCreate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> ClientCreateResponse:
|
||||
# Génération de la clé API
|
||||
raw_key = secrets.token_urlsafe(32)
|
||||
key_hash = hash_api_key(raw_key)
|
||||
|
||||
client = APIClient(
|
||||
name=body.name,
|
||||
api_key_hash=key_hash,
|
||||
scopes=body.scopes,
|
||||
plan=body.plan,
|
||||
)
|
||||
db.add(client)
|
||||
await db.flush()
|
||||
await db.refresh(client)
|
||||
|
||||
logger.info("Client créé : %s (%s)", client.name, client.id)
|
||||
|
||||
return ClientCreateResponse(
|
||||
id=client.id,
|
||||
name=client.name,
|
||||
scopes=client.scopes,
|
||||
plan=client.plan,
|
||||
is_active=client.is_active,
|
||||
created_at=client.created_at,
|
||||
updated_at=client.updated_at,
|
||||
api_key=raw_key,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# LISTER LES CLIENTS
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/clients",
|
||||
response_model=list[ClientResponse],
|
||||
summary="Lister tous les clients API",
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def list_clients(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> list[ClientResponse]:
|
||||
result = await db.execute(select(APIClient).order_by(APIClient.created_at.desc()))
|
||||
clients = result.scalars().all()
|
||||
return [ClientResponse.model_validate(c) for c in clients]
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DÉTAIL D'UN CLIENT
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/clients/{client_id}",
|
||||
response_model=ClientResponse,
|
||||
summary="Détail d'un client API",
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def get_client(
|
||||
client_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> ClientResponse:
|
||||
result = await db.execute(select(APIClient).where(APIClient.id == client_id))
|
||||
client = result.scalar_one_or_none()
|
||||
if not client:
|
||||
raise HTTPException(status_code=404, detail="Client introuvable")
|
||||
return ClientResponse.model_validate(client)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# MODIFIER UN CLIENT
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.patch(
|
||||
"/clients/{client_id}",
|
||||
response_model=ClientResponse,
|
||||
summary="Modifier un client API",
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def update_client(
|
||||
client_id: str,
|
||||
body: ClientUpdate,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> ClientResponse:
|
||||
result = await db.execute(select(APIClient).where(APIClient.id == client_id))
|
||||
client = result.scalar_one_or_none()
|
||||
if not client:
|
||||
raise HTTPException(status_code=404, detail="Client introuvable")
|
||||
|
||||
update_data = body.model_dump(exclude_unset=True)
|
||||
for field, value in update_data.items():
|
||||
setattr(client, field, value)
|
||||
|
||||
await db.flush()
|
||||
await db.refresh(client)
|
||||
|
||||
logger.info("Client mis à jour : %s (%s)", client.name, client.id)
|
||||
return ClientResponse.model_validate(client)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# ROTATION DE CLÉ
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.post(
|
||||
"/clients/{client_id}/rotate-key",
|
||||
response_model=KeyRotateResponse,
|
||||
summary="Régénérer la clé API d'un client",
|
||||
description="Invalide l'ancienne clé et en génère une nouvelle.",
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def rotate_key(
|
||||
client_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> KeyRotateResponse:
|
||||
result = await db.execute(select(APIClient).where(APIClient.id == client_id))
|
||||
client = result.scalar_one_or_none()
|
||||
if not client:
|
||||
raise HTTPException(status_code=404, detail="Client introuvable")
|
||||
|
||||
raw_key = secrets.token_urlsafe(32)
|
||||
client.api_key_hash = hash_api_key(raw_key)
|
||||
await db.flush()
|
||||
|
||||
logger.info("Clé API rotée pour client : %s (%s)", client.name, client.id)
|
||||
|
||||
return KeyRotateResponse(id=client.id, api_key=raw_key)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DÉSACTIVER UN CLIENT (soft delete)
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.delete(
|
||||
"/clients/{client_id}",
|
||||
response_model=ClientResponse,
|
||||
summary="Désactiver un client API",
|
||||
description="Soft delete — marque le client comme inactif sans supprimer les données.",
|
||||
dependencies=[Depends(require_scope("admin"))],
|
||||
)
|
||||
async def delete_client(
|
||||
client_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
) -> ClientResponse:
|
||||
result = await db.execute(select(APIClient).where(APIClient.id == client_id))
|
||||
client = result.scalar_one_or_none()
|
||||
if not client:
|
||||
raise HTTPException(status_code=404, detail="Client introuvable")
|
||||
|
||||
client.is_active = False
|
||||
await db.flush()
|
||||
await db.refresh(client)
|
||||
|
||||
logger.info("Client désactivé : %s (%s)", client.name, client.id)
|
||||
return ClientResponse.model_validate(client)
|
||||
@@ -0,0 +1,68 @@
|
||||
"""
|
||||
Router — Files : sert les fichiers locaux via URLs signées HMAC.
|
||||
|
||||
Monté uniquement quand STORAGE_BACKEND == "local".
|
||||
"""
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
from fastapi.responses import FileResponse
|
||||
|
||||
from app.config import settings
|
||||
from app.services.storage_backend import get_storage_backend, LocalStorage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/files", tags=["Fichiers"])
|
||||
|
||||
|
||||
@router.get(
|
||||
"/signed/{token}",
|
||||
summary="Télécharger un fichier via URL signée",
|
||||
description="Valide le token HMAC et retourne le fichier correspondant.",
|
||||
)
|
||||
async def serve_signed_file(token: str):
|
||||
"""Sert un fichier local via un token HMAC signé."""
|
||||
backend = get_storage_backend()
|
||||
|
||||
if not isinstance(backend, LocalStorage):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="Endpoint non disponible avec le backend de stockage actuel",
|
||||
)
|
||||
|
||||
# Valider le token
|
||||
path = backend.validate_token(token)
|
||||
if path is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Token invalide ou expiré",
|
||||
)
|
||||
|
||||
# Vérifier que le fichier existe
|
||||
abs_path = backend.get_absolute_path(path)
|
||||
if not abs_path.exists():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_410_GONE,
|
||||
detail="Le fichier n'existe plus",
|
||||
)
|
||||
|
||||
# Détecter le content type
|
||||
suffix = abs_path.suffix.lower()
|
||||
mime_map = {
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".png": "image/png",
|
||||
".gif": "image/gif",
|
||||
".webp": "image/webp",
|
||||
".bmp": "image/bmp",
|
||||
".tiff": "image/tiff",
|
||||
}
|
||||
media_type = mime_map.get(suffix, "application/octet-stream")
|
||||
|
||||
return FileResponse(
|
||||
path=str(abs_path),
|
||||
media_type=media_type,
|
||||
filename=abs_path.name,
|
||||
)
|
||||
@@ -0,0 +1,500 @@
|
||||
"""
|
||||
Router — Images : upload, lecture, suppression, retraitement
|
||||
Sécurisé : authentification par API Key + isolation par client_id.
|
||||
"""
|
||||
import logging
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import (
|
||||
APIRouter, Depends, HTTPException, UploadFile, File,
|
||||
Request, status, Query,
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select, func, or_
|
||||
|
||||
from app.database import get_db
|
||||
from app.dependencies.auth import get_current_client, require_scope
|
||||
from app.models.client import APIClient
|
||||
from app.models.image import Image, ProcessingStatus
|
||||
from app.schemas import (
|
||||
UploadResponse, ImageDetail, ImageSummary,
|
||||
StatusResponse, PaginatedImages, DeleteResponse,
|
||||
TagsResponse, ReprocessResponse,
|
||||
)
|
||||
from app.services import storage
|
||||
from app.middleware import limiter, get_upload_rate_limit
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/images", tags=["Images"])
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# UTILITAIRE : récupérer une image avec isolation client
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
async def get_image_or_404(
|
||||
image_id: int, client_id: str, db: AsyncSession
|
||||
) -> Image:
|
||||
"""
|
||||
Récupère une image par ID en vérifiant qu'elle appartient au client.
|
||||
Lève HTTP 404 si introuvable ou si elle n'appartient pas au client.
|
||||
"""
|
||||
result = await db.execute(
|
||||
select(Image).where(
|
||||
Image.id == image_id,
|
||||
Image.client_id == client_id,
|
||||
)
|
||||
)
|
||||
image = result.scalar_one_or_none()
|
||||
if not image:
|
||||
raise HTTPException(status_code=404, detail="Image introuvable")
|
||||
return image
|
||||
|
||||
|
||||
def _dynamic_upload_limit(key: str) -> str:
|
||||
"""Retourne la limite dynamique basée sur le plan du client."""
|
||||
# On parse le plan depuis la clé ou le state — fallback free
|
||||
return get_upload_rate_limit("free")
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# UPLOAD
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.post(
|
||||
"/upload",
|
||||
response_model=UploadResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Uploader une image",
|
||||
description="Upload une image, lance automatiquement le pipeline AI (EXIF + OCR + Vision).",
|
||||
dependencies=[Depends(require_scope("images:write"))],
|
||||
)
|
||||
@limiter.limit("500/hour")
|
||||
async def upload_image(
|
||||
request: Request,
|
||||
file: UploadFile = File(...),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
# Vérification quota avant upload
|
||||
quota_mb = client.quota_storage_mb or 500
|
||||
used_bytes = client.storage_used_bytes or 0
|
||||
if used_bytes >= quota_mb * 1024 * 1024:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"Quota de stockage dépassé ({quota_mb} MB)",
|
||||
)
|
||||
|
||||
# Sauvegarde fichier + thumbnail (isolé par client_id)
|
||||
file_data = await storage.save_upload(file, client_id=client.id)
|
||||
|
||||
# Création de l'enregistrement BDD
|
||||
image = Image(**file_data)
|
||||
db.add(image)
|
||||
|
||||
# Mise à jour du quota
|
||||
file_size = file_data.get("file_size", 0)
|
||||
client.storage_used_bytes = (client.storage_used_bytes or 0) + file_size
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(image)
|
||||
|
||||
# Enqueue dans ARQ (persistant, avec retry)
|
||||
arq_pool = request.app.state.arq_pool
|
||||
queue_name = "premium" if client.plan and client.plan.value == "premium" else "standard"
|
||||
await arq_pool.enqueue_job(
|
||||
"process_image_task",
|
||||
image.id,
|
||||
str(client.id),
|
||||
_queue_name=queue_name,
|
||||
)
|
||||
|
||||
return UploadResponse(
|
||||
id=image.id,
|
||||
uuid=image.uuid,
|
||||
original_name=image.original_name,
|
||||
status=image.processing_status,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# LISTE
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"",
|
||||
response_model=PaginatedImages,
|
||||
summary="Lister les images",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def list_images(
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=100),
|
||||
tag: Optional[str] = Query(None, description="Filtrer par tag AI"),
|
||||
status_filter: Optional[ProcessingStatus] = Query(None, alias="status"),
|
||||
search: Optional[str] = Query(None, description="Recherche dans description et OCR"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
# Filtre d'isolation par client
|
||||
query = select(Image).where(Image.client_id == client.id)
|
||||
|
||||
if status_filter:
|
||||
query = query.where(Image.processing_status == status_filter)
|
||||
|
||||
if tag:
|
||||
query = query.where(Image.ai_tags.contains([tag]))
|
||||
|
||||
if search:
|
||||
query = query.where(
|
||||
or_(
|
||||
Image.ai_description.ilike(f"%{search}%"),
|
||||
Image.ocr_text.ilike(f"%{search}%"),
|
||||
Image.original_name.ilike(f"%{search}%"),
|
||||
)
|
||||
)
|
||||
|
||||
# Count total
|
||||
count_query = select(func.count()).select_from(query.subquery())
|
||||
total_result = await db.execute(count_query)
|
||||
total = total_result.scalar_one()
|
||||
|
||||
# Pagination
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(Image.uploaded_at.desc()).offset(offset).limit(page_size)
|
||||
result = await db.execute(query)
|
||||
images = result.scalars().all()
|
||||
|
||||
# Quota info
|
||||
used_mb = round((client.storage_used_bytes or 0) / (1024 * 1024), 2)
|
||||
quota_mb = client.quota_storage_mb or 500
|
||||
pct = round(used_mb / quota_mb * 100, 1) if quota_mb > 0 else 0.0
|
||||
|
||||
return PaginatedImages(
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
pages=math.ceil(total / page_size) if total else 0,
|
||||
storage_used_mb=used_mb,
|
||||
storage_quota_mb=quota_mb,
|
||||
quota_pct=pct,
|
||||
items=[
|
||||
ImageSummary(
|
||||
id=img.id,
|
||||
uuid=img.uuid,
|
||||
original_name=img.original_name,
|
||||
mime_type=img.mime_type,
|
||||
file_size=img.file_size,
|
||||
width=img.width,
|
||||
height=img.height,
|
||||
uploaded_at=img.uploaded_at,
|
||||
processing_status=img.processing_status,
|
||||
ai_tags=img.ai_tags,
|
||||
ai_description=img.ai_description,
|
||||
thumbnail_path=img.thumbnail_path,
|
||||
)
|
||||
for img in images
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DÉTAIL COMPLET
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}",
|
||||
response_model=ImageDetail,
|
||||
summary="Détail complet d'une image",
|
||||
description="Retourne toutes les données : fichier, EXIF, OCR et résultats AI.",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_image(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
return ImageDetail.from_orm_full(image)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# STATUT DU PIPELINE
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/status",
|
||||
response_model=StatusResponse,
|
||||
summary="Statut du traitement AI",
|
||||
description="Permet de poller l'avancement du pipeline (pending → processing → done/error).",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_status(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
return StatusResponse(
|
||||
id=image.id,
|
||||
uuid=image.uuid,
|
||||
status=image.processing_status,
|
||||
error=image.processing_error,
|
||||
started_at=image.processing_started_at,
|
||||
done_at=image.processing_done_at,
|
||||
)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DONNÉES EXIF
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/exif",
|
||||
summary="Métadonnées EXIF de l'image",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_exif(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
return {
|
||||
"id": image.id,
|
||||
"camera": {
|
||||
"make": image.exif_make,
|
||||
"model": image.exif_model,
|
||||
"lens": image.exif_lens,
|
||||
"iso": image.exif_iso,
|
||||
"aperture": image.exif_aperture,
|
||||
"shutter_speed": image.exif_shutter,
|
||||
"focal_length": image.exif_focal,
|
||||
"flash": image.exif_flash,
|
||||
"orientation": image.exif_orientation,
|
||||
"software": image.exif_software,
|
||||
"taken_at": image.exif_taken_at,
|
||||
},
|
||||
"gps": {
|
||||
"latitude": image.exif_gps_lat,
|
||||
"longitude": image.exif_gps_lon,
|
||||
"altitude": image.exif_altitude,
|
||||
"has_gps": image.has_gps,
|
||||
"maps_url": (
|
||||
f"https://maps.google.com/?q={image.exif_gps_lat},{image.exif_gps_lon}"
|
||||
if image.has_gps else None
|
||||
),
|
||||
},
|
||||
"raw": image.exif_raw,
|
||||
}
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DONNÉES OCR
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/ocr",
|
||||
summary="Texte extrait de l'image (OCR)",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_ocr(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
return {
|
||||
"id": image.id,
|
||||
"has_text": image.ocr_has_text,
|
||||
"text": image.ocr_text,
|
||||
"language": image.ocr_language,
|
||||
"confidence": image.ocr_confidence,
|
||||
}
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# DONNÉES AI
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/ai",
|
||||
summary="Résultats AI (description + tags)",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_ai(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
return {
|
||||
"id": image.id,
|
||||
"description": image.ai_description,
|
||||
"tags": image.ai_tags,
|
||||
"confidence": image.ai_confidence,
|
||||
"model_used": image.ai_model_used,
|
||||
"processed_at": image.ai_processed_at,
|
||||
"tokens": {
|
||||
"prompt": image.ai_prompt_tokens,
|
||||
"output": image.ai_output_tokens,
|
||||
"total": (image.ai_prompt_tokens or 0) + (image.ai_output_tokens or 0),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# TAGS — Vue globale (filtrée par client)
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/tags/all",
|
||||
response_model=TagsResponse,
|
||||
summary="Tous les tags utilisés",
|
||||
description="Liste dédupliquée de tous les tags AI générés sur les images du client.",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_all_tags(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Image.ai_tags).where(
|
||||
Image.client_id == client.id,
|
||||
Image.ai_tags.isnot(None),
|
||||
)
|
||||
)
|
||||
all_tag_lists = result.scalars().all()
|
||||
|
||||
unique_tags = sorted(set(
|
||||
tag
|
||||
for tag_list in all_tag_lists
|
||||
for tag in (tag_list or [])
|
||||
))
|
||||
|
||||
return TagsResponse(tags=unique_tags, total=len(unique_tags))
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# RETRAITEMENT AI
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.post(
|
||||
"/{image_id}/reprocess",
|
||||
response_model=ReprocessResponse,
|
||||
summary="Relancer le pipeline AI",
|
||||
description="Reprocess une image existante (utile après changement de modèle AI).",
|
||||
dependencies=[Depends(require_scope("images:write"))],
|
||||
)
|
||||
@limiter.limit("500/hour")
|
||||
async def reprocess_image(
|
||||
request: Request,
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
# Reset du statut
|
||||
image.processing_status = ProcessingStatus.PENDING
|
||||
image.processing_error = None
|
||||
image.processing_started_at = None
|
||||
image.processing_done_at = None
|
||||
await db.commit()
|
||||
|
||||
# Enqueue dans ARQ
|
||||
arq_pool = request.app.state.arq_pool
|
||||
queue_name = "premium" if client.plan and client.plan.value == "premium" else "standard"
|
||||
await arq_pool.enqueue_job(
|
||||
"process_image_task",
|
||||
image_id,
|
||||
str(client.id),
|
||||
_queue_name=queue_name,
|
||||
)
|
||||
|
||||
return ReprocessResponse(id=image_id)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# SUPPRESSION
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.delete(
|
||||
"/{image_id}",
|
||||
response_model=DeleteResponse,
|
||||
summary="Supprimer une image",
|
||||
dependencies=[Depends(require_scope("images:delete"))],
|
||||
)
|
||||
async def delete_image(
|
||||
image_id: int,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
|
||||
# Décrémentation du quota
|
||||
file_size = image.file_size or 0
|
||||
client.storage_used_bytes = max(0, (client.storage_used_bytes or 0) - file_size)
|
||||
|
||||
# Suppression des fichiers sur disque
|
||||
storage.delete_files(image.file_path, image.thumbnail_path)
|
||||
|
||||
await db.delete(image)
|
||||
await db.commit()
|
||||
|
||||
return DeleteResponse(deleted_id=image_id)
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# URLs SIGNÉES
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/download-url",
|
||||
summary="URL signée de téléchargement",
|
||||
description="Retourne une URL signée temporaire pour télécharger l'image originale.",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_download_url(
|
||||
image_id: int,
|
||||
expires_in: int = Query(900, ge=60, le=86400, description="Durée de validité en secondes"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
from app.services.storage_backend import get_storage_backend
|
||||
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
backend = get_storage_backend()
|
||||
|
||||
url = await backend.get_signed_url(image.file_path, expires_in=expires_in)
|
||||
return {"url": url, "expires_in": expires_in}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{image_id}/thumbnail-url",
|
||||
summary="URL signée du thumbnail",
|
||||
description="Retourne une URL signée temporaire pour le thumbnail de l'image.",
|
||||
dependencies=[Depends(require_scope("images:read"))],
|
||||
)
|
||||
async def get_thumbnail_url(
|
||||
image_id: int,
|
||||
expires_in: int = Query(900, ge=60, le=86400, description="Durée de validité en secondes"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
client: APIClient = Depends(get_current_client),
|
||||
):
|
||||
from app.services.storage_backend import get_storage_backend
|
||||
|
||||
image = await get_image_or_404(image_id, client.id, db)
|
||||
if not image.thumbnail_path:
|
||||
raise HTTPException(status_code=404, detail="Thumbnail non disponible")
|
||||
|
||||
backend = get_storage_backend()
|
||||
url = await backend.get_signed_url(image.thumbnail_path, expires_in=expires_in)
|
||||
return {"url": url, "expires_in": expires_in}
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
"""
|
||||
Schémas Pydantic — validation et sérialisation des réponses API
|
||||
"""
|
||||
from datetime import datetime
|
||||
from typing import Any, List, Optional
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from app.models.image import ProcessingStatus
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Sous-schémas imbriqués
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
class ExifGPS(BaseModel):
|
||||
latitude: Optional[float] = None
|
||||
longitude: Optional[float] = None
|
||||
altitude: Optional[float] = None
|
||||
has_gps: bool = False
|
||||
|
||||
|
||||
class ExifCamera(BaseModel):
|
||||
make: Optional[str] = None
|
||||
model: Optional[str] = None
|
||||
lens: Optional[str] = None
|
||||
iso: Optional[int] = None
|
||||
aperture: Optional[str] = None
|
||||
shutter_speed: Optional[str] = None
|
||||
focal_length: Optional[str] = None
|
||||
flash: Optional[bool] = None
|
||||
orientation: Optional[int] = None
|
||||
software: Optional[str] = None
|
||||
taken_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class ExifData(BaseModel):
|
||||
camera: ExifCamera
|
||||
gps: ExifGPS
|
||||
raw: Optional[dict[str, Any]] = None
|
||||
|
||||
|
||||
class OcrData(BaseModel):
|
||||
text: Optional[str] = None
|
||||
language: Optional[str] = None
|
||||
confidence: Optional[float] = None
|
||||
has_text: bool = False
|
||||
|
||||
|
||||
class AiData(BaseModel):
|
||||
description: Optional[str] = None
|
||||
tags: Optional[List[str]] = None
|
||||
confidence: Optional[float] = None
|
||||
model_used: Optional[str] = None
|
||||
processed_at: Optional[datetime] = None
|
||||
prompt_tokens: Optional[int] = None
|
||||
output_tokens: Optional[int] = None
|
||||
|
||||
|
||||
class ProcessingInfo(BaseModel):
|
||||
status: ProcessingStatus
|
||||
error: Optional[str] = None
|
||||
started_at: Optional[datetime] = None
|
||||
done_at: Optional[datetime] = None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Réponses principales
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
class ImageBase(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: int
|
||||
uuid: str
|
||||
original_name: str
|
||||
mime_type: Optional[str] = None
|
||||
file_size: Optional[int] = None
|
||||
width: Optional[int] = None
|
||||
height: Optional[int] = None
|
||||
uploaded_at: Optional[datetime] = None
|
||||
processing_status: ProcessingStatus
|
||||
|
||||
|
||||
class ImageSummary(ImageBase):
|
||||
"""Version allégée pour les listes."""
|
||||
ai_tags: Optional[List[str]] = None
|
||||
ai_description: Optional[str] = None
|
||||
thumbnail_path: Optional[str] = None
|
||||
|
||||
|
||||
class ImageDetail(ImageBase):
|
||||
"""Version complète avec toutes les données collectées."""
|
||||
exif: ExifData
|
||||
ocr: OcrData
|
||||
ai: AiData
|
||||
processing: ProcessingInfo
|
||||
|
||||
@classmethod
|
||||
def from_orm_full(cls, img) -> "ImageDetail":
|
||||
return cls(
|
||||
id=img.id,
|
||||
uuid=img.uuid,
|
||||
original_name=img.original_name,
|
||||
mime_type=img.mime_type,
|
||||
file_size=img.file_size,
|
||||
width=img.width,
|
||||
height=img.height,
|
||||
uploaded_at=img.uploaded_at,
|
||||
processing_status=img.processing_status,
|
||||
thumbnail_path=img.thumbnail_path,
|
||||
exif=ExifData(
|
||||
camera=ExifCamera(
|
||||
make=img.exif_make,
|
||||
model=img.exif_model,
|
||||
lens=img.exif_lens,
|
||||
iso=img.exif_iso,
|
||||
aperture=img.exif_aperture,
|
||||
shutter_speed=img.exif_shutter,
|
||||
focal_length=img.exif_focal,
|
||||
flash=img.exif_flash,
|
||||
orientation=img.exif_orientation,
|
||||
software=img.exif_software,
|
||||
taken_at=img.exif_taken_at,
|
||||
),
|
||||
gps=ExifGPS(
|
||||
latitude=img.exif_gps_lat,
|
||||
longitude=img.exif_gps_lon,
|
||||
altitude=img.exif_altitude,
|
||||
has_gps=img.has_gps,
|
||||
),
|
||||
raw=img.exif_raw,
|
||||
),
|
||||
ocr=OcrData(
|
||||
text=img.ocr_text,
|
||||
language=img.ocr_language,
|
||||
confidence=img.ocr_confidence,
|
||||
has_text=img.ocr_has_text or False,
|
||||
),
|
||||
ai=AiData(
|
||||
description=img.ai_description,
|
||||
tags=img.ai_tags,
|
||||
confidence=img.ai_confidence,
|
||||
model_used=img.ai_model_used,
|
||||
processed_at=img.ai_processed_at,
|
||||
prompt_tokens=img.ai_prompt_tokens,
|
||||
output_tokens=img.ai_output_tokens,
|
||||
),
|
||||
processing=ProcessingInfo(
|
||||
status=img.processing_status,
|
||||
error=img.processing_error,
|
||||
started_at=img.processing_started_at,
|
||||
done_at=img.processing_done_at,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class UploadResponse(BaseModel):
|
||||
id: int
|
||||
uuid: str
|
||||
original_name: str
|
||||
status: ProcessingStatus
|
||||
message: str = "Image uploadée — traitement AI en cours"
|
||||
|
||||
|
||||
class StatusResponse(BaseModel):
|
||||
id: int
|
||||
uuid: str
|
||||
status: ProcessingStatus
|
||||
error: Optional[str] = None
|
||||
started_at: Optional[datetime] = None
|
||||
done_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class PaginatedImages(BaseModel):
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
pages: int
|
||||
items: List[ImageSummary]
|
||||
# Quota tracking
|
||||
storage_used_mb: Optional[float] = None
|
||||
storage_quota_mb: Optional[int] = None
|
||||
quota_pct: Optional[float] = None
|
||||
|
||||
|
||||
class DeleteResponse(BaseModel):
|
||||
deleted_id: int
|
||||
message: str = "Image supprimée avec succès"
|
||||
|
||||
|
||||
class TagsResponse(BaseModel):
|
||||
tags: List[str]
|
||||
total: int
|
||||
|
||||
|
||||
class ReprocessResponse(BaseModel):
|
||||
id: int
|
||||
message: str = "Traitement AI relancé"
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# AI — Endpoints externes (résumé URL, rédaction)
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
class SummarizeRequest(BaseModel):
|
||||
url: str
|
||||
language: str = "français"
|
||||
|
||||
|
||||
class SummarizeResponse(BaseModel):
|
||||
url: str
|
||||
title: Optional[str] = None
|
||||
summary: str
|
||||
tags: List[str]
|
||||
model: str
|
||||
|
||||
|
||||
class DraftTaskRequest(BaseModel):
|
||||
description: str
|
||||
context: Optional[str] = None
|
||||
language: str = "français"
|
||||
|
||||
|
||||
class DraftTaskResponse(BaseModel):
|
||||
title: str
|
||||
description: str
|
||||
steps: List[str]
|
||||
estimated_time: Optional[str] = None
|
||||
priority: Optional[str] = None
|
||||
@@ -0,0 +1,67 @@
|
||||
"""
|
||||
Schémas Pydantic — authentification et gestion des clients API
|
||||
"""
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from app.models.client import ClientPlan
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Requêtes
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
class ClientCreate(BaseModel):
|
||||
"""Créer un nouveau client API."""
|
||||
name: str = Field(..., min_length=1, max_length=256, description="Nom de l'application cliente")
|
||||
scopes: List[str] = Field(
|
||||
default=["images:read", "images:write"],
|
||||
description="Permissions accordées",
|
||||
)
|
||||
plan: ClientPlan = Field(default=ClientPlan.FREE, description="Plan tarifaire")
|
||||
|
||||
|
||||
class ClientUpdate(BaseModel):
|
||||
"""Modifier un client API existant."""
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=256)
|
||||
scopes: Optional[List[str]] = None
|
||||
plan: Optional[ClientPlan] = None
|
||||
is_active: Optional[bool] = None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
# Réponses
|
||||
# ─────────────────────────────────────────────────────────────
|
||||
|
||||
class ClientResponse(BaseModel):
|
||||
"""Réponse de base pour un client API."""
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
name: str
|
||||
scopes: List[str]
|
||||
plan: ClientPlan
|
||||
is_active: bool
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class ClientCreateResponse(ClientResponse):
|
||||
"""
|
||||
Réponse après création d'un client.
|
||||
La clé API est retournée EN CLAIR une seule fois.
|
||||
"""
|
||||
api_key: str = Field(
|
||||
...,
|
||||
description="Clé API en clair — stockez-la, elle ne sera plus jamais affichée",
|
||||
)
|
||||
|
||||
|
||||
class KeyRotateResponse(BaseModel):
|
||||
"""Réponse après rotation de la clé API."""
|
||||
id: str
|
||||
api_key: str = Field(
|
||||
...,
|
||||
description="Nouvelle clé API en clair — stockez-la, elle ne sera plus jamais affichée",
|
||||
)
|
||||
message: str = "Clé API régénérée avec succès — l'ancienne clé est désormais invalide"
|
||||
@@ -0,0 +1,3 @@
|
||||
from app.services import storage, exif_service, ocr_service, ai_vision, scraper, pipeline
|
||||
|
||||
__all__ = ["storage", "exif_service", "ocr_service", "ai_vision", "scraper", "pipeline"]
|
||||
@@ -0,0 +1,417 @@
|
||||
"""
|
||||
Service AI Vision — description, classification et tags via Google Gemini ou OpenRouter
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import base64
|
||||
import httpx
|
||||
from pathlib import Path
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_client: Optional[genai.Client] = None
|
||||
|
||||
|
||||
def _get_client() -> genai.Client:
|
||||
global _client
|
||||
if _client is None:
|
||||
_client = genai.Client(api_key=settings.GEMINI_API_KEY)
|
||||
return _client
|
||||
|
||||
|
||||
def _read_image(file_path: str) -> tuple[bytes, str]:
|
||||
"""Lit l'image en bytes et détecte le media_type."""
|
||||
path = Path(file_path)
|
||||
suffix = path.suffix.lower()
|
||||
|
||||
mime_map = {
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".png": "image/png",
|
||||
".gif": "image/gif",
|
||||
".webp": "image/webp",
|
||||
}
|
||||
media_type = mime_map.get(suffix, "image/jpeg")
|
||||
|
||||
with open(path, "rb") as f:
|
||||
data = f.read()
|
||||
|
||||
return data, media_type
|
||||
|
||||
|
||||
def _extract_json(text: str) -> Optional[dict]:
|
||||
cleaned = re.sub(r"```json\s*|```\s*", "", (text or "")).strip()
|
||||
json_match = re.search(r"\{.*\}", cleaned, re.DOTALL)
|
||||
if not json_match:
|
||||
return None
|
||||
try:
|
||||
return json.loads(json_match.group())
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
|
||||
def _usage_tokens_gemini(response) -> tuple[Optional[int], Optional[int]]:
|
||||
usage = getattr(response, "usage_metadata", None)
|
||||
if not usage:
|
||||
return None, None
|
||||
prompt_tokens = getattr(usage, "prompt_token_count", None)
|
||||
output_tokens = getattr(usage, "candidates_token_count", None)
|
||||
return prompt_tokens, output_tokens
|
||||
|
||||
|
||||
async def _generate_gemini(
|
||||
prompt: str,
|
||||
image_bytes: Optional[bytes] = None,
|
||||
media_type: Optional[str] = None,
|
||||
max_tokens: int = 1024
|
||||
) -> dict:
|
||||
"""Appel à Google Gemini via SDK."""
|
||||
if not settings.GEMINI_API_KEY:
|
||||
logger.warning("ai.gemini.no_key")
|
||||
return {"text": None, "usage": (None, None)}
|
||||
|
||||
client = _get_client()
|
||||
contents = []
|
||||
if image_bytes and media_type:
|
||||
contents.append(types.Part.from_bytes(data=image_bytes, mime_type=media_type))
|
||||
contents.append(prompt)
|
||||
|
||||
try:
|
||||
# Le SDK est sync, on le run dans un thread
|
||||
response = await asyncio.to_thread(
|
||||
client.models.generate_content,
|
||||
model=settings.GEMINI_MODEL,
|
||||
contents=contents,
|
||||
config=types.GenerateContentConfig(
|
||||
max_output_tokens=max_tokens,
|
||||
response_mime_type="application/json",
|
||||
),
|
||||
)
|
||||
usage = _usage_tokens_gemini(response)
|
||||
return {"text": getattr(response, "text", ""), "usage": usage}
|
||||
except Exception as e:
|
||||
logger.error("ai.gemini.error", extra={"error": str(e)})
|
||||
return {"text": None, "usage": (None, None), "error": str(e)}
|
||||
|
||||
|
||||
async def _generate_openrouter(
|
||||
prompt: str,
|
||||
image_bytes: Optional[bytes] = None,
|
||||
media_type: Optional[str] = None,
|
||||
max_tokens: int = 1024
|
||||
) -> dict:
|
||||
"""Appel à OpenRouter via HTTP."""
|
||||
if not settings.OPENROUTER_API_KEY:
|
||||
logger.warning("ai.openrouter.no_key")
|
||||
return {"text": None, "usage": (None, None)}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {settings.OPENROUTER_API_KEY}",
|
||||
"Content-Type": "application/json",
|
||||
"HTTP-Referer": settings.HOST,
|
||||
"X-Title": settings.APP_NAME,
|
||||
}
|
||||
|
||||
messages = []
|
||||
content_payload = []
|
||||
|
||||
content_payload.append({"type": "text", "text": prompt})
|
||||
|
||||
if image_bytes and media_type:
|
||||
b64_img = base64.b64encode(image_bytes).decode("utf-8")
|
||||
content_payload.append({
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{media_type};base64,{b64_img}"
|
||||
}
|
||||
})
|
||||
|
||||
messages.append({"role": "user", "content": content_payload})
|
||||
|
||||
payload = {
|
||||
"model": settings.OPENROUTER_MODEL,
|
||||
"messages": messages,
|
||||
"max_tokens": max_tokens,
|
||||
# OpenRouter/OpenAI support response_format={"type": "json_object"} pour certains modèles
|
||||
# On tente le coup si le modèle est compatible, sinon le prompt engineering fait le travail
|
||||
"response_format": {"type": "json_object"}
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
response = await client.post(
|
||||
"https://openrouter.ai/api/v1/chat/completions",
|
||||
json=payload,
|
||||
headers=headers,
|
||||
timeout=60.0
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
text = ""
|
||||
if "choices" in data and len(data["choices"]) > 0:
|
||||
text = data["choices"][0]["message"]["content"]
|
||||
|
||||
usage_data = data.get("usage", {})
|
||||
prompt_tokens = usage_data.get("prompt_tokens")
|
||||
output_tokens = usage_data.get("completion_tokens")
|
||||
|
||||
return {"text": text, "usage": (prompt_tokens, output_tokens)}
|
||||
|
||||
except Exception as e:
|
||||
logger.error("ai.openrouter.error", extra={"error": str(e)})
|
||||
return {"text": None, "usage": (None, None), "error": str(e)}
|
||||
|
||||
|
||||
async def _generate(
|
||||
prompt: str,
|
||||
image_bytes: Optional[bytes] = None,
|
||||
media_type: Optional[str] = None,
|
||||
max_tokens: int = 1024
|
||||
) -> dict:
|
||||
"""Dispatcher vers le bon provider."""
|
||||
provider = settings.AI_PROVIDER.lower()
|
||||
logger.info("ai.generate", extra={"provider": provider})
|
||||
|
||||
if provider == "openrouter":
|
||||
return await _generate_openrouter(prompt, image_bytes, media_type, max_tokens)
|
||||
else:
|
||||
# Default to Gemini
|
||||
return await _generate_gemini(prompt, image_bytes, media_type, max_tokens)
|
||||
|
||||
|
||||
def _build_prompt(ocr_hint: Optional[str], language: str) -> str:
|
||||
ocr_section = ""
|
||||
if ocr_hint and len(ocr_hint.strip()) > 5:
|
||||
ocr_section = f"""
|
||||
Texte détecté dans l'image par OCR (utilise-le pour enrichir ta réponse) :
|
||||
\"\"\"
|
||||
{ocr_hint[:500]}
|
||||
\"\"\"
|
||||
"""
|
||||
|
||||
return f"""Analyse cette image avec précision et retourne UNIQUEMENT un objet JSON valide avec ces champs :
|
||||
|
||||
{{
|
||||
"description": "Description complète et détaillée en {language}, 2-4 phrases. Décris le sujet principal, le contexte, les couleurs, l'ambiance.",
|
||||
"tags": ["tag1", "tag2", "tag3"],
|
||||
"confidence": 0.95
|
||||
}}
|
||||
|
||||
Règles pour les tags :
|
||||
- Entre {settings.AI_TAGS_MIN} et {settings.AI_TAGS_MAX} tags
|
||||
- En minuscules, sans espaces (utiliser des tirets si nécessaire)
|
||||
- Couvrir : sujet principal, type d'image, couleurs dominantes, style, contexte
|
||||
- Exemples : portrait, paysage, architecture, nature, nourriture, texte, document, animal, sport, technologie, intérieur, extérieur
|
||||
{ocr_section}
|
||||
Réponds UNIQUEMENT avec le JSON, sans texte avant ou après, sans balises markdown."""
|
||||
|
||||
|
||||
async def analyze_image(
|
||||
file_path: str,
|
||||
ocr_hint: Optional[str] = None,
|
||||
language: str = "français",
|
||||
) -> dict:
|
||||
"""
|
||||
Envoie l'image à l'AI pour analyse (Description + Tags).
|
||||
"""
|
||||
if not settings.AI_ENABLED:
|
||||
return {}
|
||||
|
||||
result = {
|
||||
"description": None,
|
||||
"tags": [],
|
||||
"confidence": None,
|
||||
"model": settings.OPENROUTER_MODEL if settings.AI_PROVIDER == "openrouter" else settings.GEMINI_MODEL,
|
||||
"prompt_tokens": None,
|
||||
"output_tokens": None,
|
||||
}
|
||||
|
||||
try:
|
||||
image_bytes, media_type = _read_image(file_path)
|
||||
prompt = _build_prompt(ocr_hint, language)
|
||||
|
||||
response = await _generate(
|
||||
prompt=prompt,
|
||||
image_bytes=image_bytes,
|
||||
media_type=media_type,
|
||||
max_tokens=settings.GEMINI_MAX_TOKENS # Ou une config unifiée
|
||||
)
|
||||
|
||||
text = response.get("text")
|
||||
result["prompt_tokens"], result["output_tokens"] = response.get("usage")
|
||||
|
||||
if text:
|
||||
parsed = _extract_json(text)
|
||||
if parsed:
|
||||
result["description"] = parsed.get("description")
|
||||
result["tags"] = parsed.get("tags", [])
|
||||
result["confidence"] = parsed.get("confidence")
|
||||
else:
|
||||
logger.warning("ai.vision.json_parse_failed", extra={"raw": text[:100]})
|
||||
|
||||
if response.get("error"):
|
||||
logger.error("ai.vision.provider_error", extra={"error": response['error']})
|
||||
|
||||
except Exception as e:
|
||||
logger.error("ai.vision.unexpected_error", extra={"error": str(e)})
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def extract_text_with_ai(file_path: str) -> dict:
|
||||
"""
|
||||
Utilise l'AI comme fallback OCR.
|
||||
"""
|
||||
result = {
|
||||
"text": None,
|
||||
"has_text": False,
|
||||
"language": "unknown",
|
||||
"confidence": 0.0,
|
||||
"method": f"ai-{settings.AI_PROVIDER}"
|
||||
}
|
||||
|
||||
if not settings.AI_ENABLED:
|
||||
return result
|
||||
|
||||
logger.info("ai.ocr.fallback_start", extra={"file": Path(file_path).name})
|
||||
|
||||
try:
|
||||
image_bytes, media_type = _read_image(file_path)
|
||||
prompt = """Agis comme un moteur OCR avancé.
|
||||
Extrais TOUT le texte visible dans cette image.
|
||||
Retourne UNIQUEMENT un objet JSON :
|
||||
{
|
||||
"text": "Le texte complet extrait ici...",
|
||||
"language": "fr" (code langue ISO 2 lettres, ex: fr, en, es),
|
||||
"confidence": 0.9 (estimation confiance 0.0 à 1.0)
|
||||
}
|
||||
Si aucun texte n'est visible, retourne : {"text": "", "has_text": false}
|
||||
"""
|
||||
|
||||
response = await _generate(
|
||||
prompt=prompt,
|
||||
image_bytes=image_bytes,
|
||||
media_type=media_type,
|
||||
max_tokens=1024
|
||||
)
|
||||
|
||||
text = response.get("text")
|
||||
if text:
|
||||
parsed = _extract_json(text)
|
||||
if parsed:
|
||||
extracted = parsed.get("text", "").strip()
|
||||
result["text"] = extracted
|
||||
result["has_text"] = bool(extracted) or parsed.get("has_text", False)
|
||||
result["language"] = parsed.get("language", "unknown")
|
||||
result["confidence"] = parsed.get("confidence", 0.0)
|
||||
logger.info("ai.ocr.success", extra={"chars": len(extracted)})
|
||||
else:
|
||||
logger.warning("ai.ocr.json_parse_failed")
|
||||
else:
|
||||
logger.info("ai.ocr.empty_response")
|
||||
|
||||
except Exception as e:
|
||||
logger.error("ai.ocr.error", extra={"error": str(e)})
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def summarize_url(url: str, content: str, language: str = "français") -> dict:
|
||||
"""Génère un résumé et des tags pour un contenu web."""
|
||||
result = {
|
||||
"summary": "",
|
||||
"tags": [],
|
||||
"model": settings.AI_PROVIDER,
|
||||
}
|
||||
|
||||
if not settings.AI_ENABLED:
|
||||
return result
|
||||
|
||||
prompt = f"""Tu reçois le contenu d'une page web. Génère un résumé et des tags en {language}.
|
||||
|
||||
URL : {url}
|
||||
|
||||
Contenu :
|
||||
\"\"\"
|
||||
{content[:3000]}
|
||||
\"\"\"
|
||||
|
||||
Retourne UNIQUEMENT ce JSON :
|
||||
{{
|
||||
"summary": "Résumé clair en 3-5 phrases en {language}",
|
||||
"tags": ["tag1", "tag2", "tag3"]
|
||||
}}"""
|
||||
|
||||
try:
|
||||
response = await _generate(
|
||||
prompt=prompt,
|
||||
max_tokens=settings.GEMINI_MAX_TOKENS
|
||||
)
|
||||
|
||||
text = response.get("text")
|
||||
if text:
|
||||
parsed = _extract_json(text)
|
||||
if parsed:
|
||||
result["summary"] = parsed.get("summary", "")
|
||||
result["tags"] = parsed.get("tags", [])
|
||||
|
||||
except Exception as e:
|
||||
logger.error("ai.summarize_url.error", extra={"error": str(e)})
|
||||
|
||||
return result
|
||||
|
||||
|
||||
async def draft_task(description: str, context: Optional[str], language: str = "français") -> dict:
|
||||
"""Génère une tâche structurée à partir d'une description."""
|
||||
result = {
|
||||
"title": "",
|
||||
"description": "",
|
||||
"steps": [],
|
||||
"estimated_time": None,
|
||||
"priority": None,
|
||||
}
|
||||
|
||||
if not settings.AI_ENABLED:
|
||||
return result
|
||||
|
||||
ctx_section = f"\nContexte : {context}" if context else ""
|
||||
prompt = f"""Tu es un assistant de gestion de tâches. Génère une tâche structurée en {language}.
|
||||
|
||||
Description : {description}{ctx_section}
|
||||
|
||||
Retourne UNIQUEMENT ce JSON :
|
||||
{{
|
||||
"title": "Titre court et actionnable",
|
||||
"description": "Description complète de la tâche",
|
||||
"steps": ["Étape 1", "Étape 2", "Étape 3"],
|
||||
"estimated_time": "30 minutes",
|
||||
"priority": "haute|moyenne|basse"
|
||||
}}"""
|
||||
|
||||
try:
|
||||
response = await _generate(
|
||||
prompt=prompt,
|
||||
max_tokens=settings.GEMINI_MAX_TOKENS
|
||||
)
|
||||
|
||||
text = response.get("text")
|
||||
if text:
|
||||
parsed = _extract_json(text)
|
||||
if parsed:
|
||||
result.update(parsed)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("ai.draft_task.error", extra={"error": str(e)})
|
||||
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
"""
|
||||
Service d'extraction EXIF — Pillow + piexif
|
||||
"""
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
import piexif
|
||||
from PIL import Image as PILImage
|
||||
from PIL.ExifTags import TAGS, GPSTAGS
|
||||
|
||||
|
||||
def _dms_to_decimal(dms: tuple, ref: str) -> float | None:
|
||||
"""Convertit les coordonnées GPS DMS (degrés/minutes/secondes) en décimal."""
|
||||
try:
|
||||
degrees = dms[0][0] / dms[0][1]
|
||||
minutes = dms[1][0] / dms[1][1]
|
||||
seconds = dms[2][0] / dms[2][1]
|
||||
decimal = degrees + minutes / 60 + seconds / 3600
|
||||
if ref in ("S", "W"):
|
||||
decimal = -decimal
|
||||
return round(decimal, 7)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _parse_rational(value) -> str | None:
|
||||
"""Convertit un rationnel EXIF en chaîne lisible."""
|
||||
try:
|
||||
if isinstance(value, tuple) and len(value) == 2:
|
||||
num, den = value
|
||||
if den == 0:
|
||||
return None
|
||||
return f"{num}/{den}"
|
||||
return str(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _safe_str(value: Any) -> str | None:
|
||||
"""Décode les bytes en string si nécessaire."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8", errors="ignore").strip("\x00")
|
||||
return str(value)
|
||||
|
||||
|
||||
def extract_exif(file_path: str) -> dict:
|
||||
"""
|
||||
Extrait toutes les métadonnées EXIF d'une image.
|
||||
Retourne un dict structuré avec les données parsées.
|
||||
"""
|
||||
result = {
|
||||
"raw": {},
|
||||
"make": None,
|
||||
"model": None,
|
||||
"lens": None,
|
||||
"taken_at": None,
|
||||
"gps_lat": None,
|
||||
"gps_lon": None,
|
||||
"altitude": None,
|
||||
"iso": None,
|
||||
"aperture": None,
|
||||
"shutter": None,
|
||||
"focal": None,
|
||||
"flash": None,
|
||||
"orientation": None,
|
||||
"software": None,
|
||||
}
|
||||
|
||||
try:
|
||||
path = Path(file_path)
|
||||
if not path.exists():
|
||||
return result
|
||||
|
||||
# ── Lecture EXIF brute via piexif ─────────────────────
|
||||
try:
|
||||
exif_data = piexif.load(str(path))
|
||||
except Exception:
|
||||
# JPEG sans EXIF, PNG, etc.
|
||||
return result
|
||||
|
||||
raw_dict = {}
|
||||
|
||||
# ── IFD 0 (Image principale) ──────────────────────────
|
||||
ifd0 = exif_data.get("0th", {})
|
||||
result["make"] = _safe_str(ifd0.get(piexif.ImageIFD.Make))
|
||||
result["model"] = _safe_str(ifd0.get(piexif.ImageIFD.Model))
|
||||
result["software"] = _safe_str(ifd0.get(piexif.ImageIFD.Software))
|
||||
result["orientation"] = ifd0.get(piexif.ImageIFD.Orientation)
|
||||
|
||||
# ── EXIF IFD ──────────────────────────────────────────
|
||||
exif = exif_data.get("Exif", {})
|
||||
|
||||
# Date de prise de vue
|
||||
taken_raw = _safe_str(exif.get(piexif.ExifIFD.DateTimeOriginal))
|
||||
if taken_raw:
|
||||
try:
|
||||
result["taken_at"] = datetime.strptime(taken_raw, "%Y:%m:%d %H:%M:%S")
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# Paramètres de prise de vue
|
||||
iso_val = exif.get(piexif.ExifIFD.ISOSpeedRatings)
|
||||
result["iso"] = int(iso_val) if iso_val else None
|
||||
|
||||
aperture_val = exif.get(piexif.ExifIFD.FNumber)
|
||||
if aperture_val:
|
||||
try:
|
||||
f = aperture_val[0] / aperture_val[1]
|
||||
result["aperture"] = f"f/{f:.1f}"
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
shutter_val = exif.get(piexif.ExifIFD.ExposureTime)
|
||||
if shutter_val:
|
||||
result["shutter"] = _parse_rational(shutter_val)
|
||||
|
||||
focal_val = exif.get(piexif.ExifIFD.FocalLength)
|
||||
if focal_val:
|
||||
try:
|
||||
f = focal_val[0] / focal_val[1]
|
||||
result["focal"] = f"{f:.0f}mm"
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
flash_val = exif.get(piexif.ExifIFD.Flash)
|
||||
result["flash"] = bool(flash_val & 1) if flash_val is not None else None
|
||||
|
||||
lens_val = _safe_str(exif.get(piexif.ExifIFD.LensModel))
|
||||
result["lens"] = lens_val
|
||||
|
||||
# ── GPS IFD ───────────────────────────────────────────
|
||||
gps = exif_data.get("GPS", {})
|
||||
if gps:
|
||||
lat_val = gps.get(piexif.GPSIFD.GPSLatitude)
|
||||
lat_ref = _safe_str(gps.get(piexif.GPSIFD.GPSLatitudeRef))
|
||||
lon_val = gps.get(piexif.GPSIFD.GPSLongitude)
|
||||
lon_ref = _safe_str(gps.get(piexif.GPSIFD.GPSLongitudeRef))
|
||||
|
||||
if lat_val and lat_ref:
|
||||
result["gps_lat"] = _dms_to_decimal(lat_val, lat_ref)
|
||||
if lon_val and lon_ref:
|
||||
result["gps_lon"] = _dms_to_decimal(lon_val, lon_ref)
|
||||
|
||||
alt_val = gps.get(piexif.GPSIFD.GPSAltitude)
|
||||
if alt_val:
|
||||
try:
|
||||
result["altitude"] = round(alt_val[0] / alt_val[1], 2)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ── Dict brut lisible (TAGS humains) ──────────────────
|
||||
with PILImage.open(path) as img:
|
||||
raw_exif = img._getexif()
|
||||
if raw_exif:
|
||||
for tag_id, val in raw_exif.items():
|
||||
tag = TAGS.get(tag_id, str(tag_id))
|
||||
if isinstance(val, bytes):
|
||||
val = val.decode("utf-8", errors="ignore")
|
||||
elif isinstance(val, tuple):
|
||||
val = list(val)
|
||||
raw_dict[tag] = val
|
||||
result["raw"] = raw_dict
|
||||
|
||||
except Exception as e:
|
||||
logger.error("exif.extraction_error", extra={"file": file_path, "error": str(e)})
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,107 @@
|
||||
"""
|
||||
Service OCR — extraction de texte via Tesseract
|
||||
"""
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from PIL import Image as PILImage
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
import pytesseract
|
||||
_ocr_import_error: Exception | None = None
|
||||
except Exception as e:
|
||||
pytesseract = None
|
||||
_ocr_import_error = e
|
||||
|
||||
|
||||
def _detect_language(text: str) -> str:
|
||||
"""Détection grossière de la langue à partir du texte extrait."""
|
||||
if not text:
|
||||
return "unknown"
|
||||
|
||||
# Mots communs français
|
||||
fr_words = {"le", "la", "les", "de", "du", "des", "un", "une", "et", "en", "est", "que"}
|
||||
# Mots communs anglais
|
||||
en_words = {"the", "is", "are", "and", "or", "of", "to", "in", "a", "an", "for", "with"}
|
||||
|
||||
words = set(text.lower().split())
|
||||
fr_score = len(words & fr_words)
|
||||
en_score = len(words & en_words)
|
||||
|
||||
if fr_score == 0 and en_score == 0:
|
||||
return "unknown"
|
||||
return "fr" if fr_score >= en_score else "en"
|
||||
|
||||
|
||||
def extract_text(file_path: str) -> dict:
|
||||
"""
|
||||
Extrait le texte d'une image via Tesseract OCR.
|
||||
Retourne un dict avec le texte, la langue et le score de confiance.
|
||||
"""
|
||||
result = {
|
||||
"text": None,
|
||||
"language": None,
|
||||
"confidence": None,
|
||||
"has_text": False,
|
||||
}
|
||||
|
||||
if not settings.OCR_ENABLED:
|
||||
return result
|
||||
|
||||
if pytesseract is None:
|
||||
logger.warning("ocr.unavailable", extra={"error": str(_ocr_import_error)})
|
||||
return result
|
||||
|
||||
path = Path(file_path)
|
||||
if not path.exists():
|
||||
return result
|
||||
|
||||
try:
|
||||
# Configuration Tesseract
|
||||
if settings.TESSERACT_CMD:
|
||||
pytesseract.pytesseract.tesseract_cmd = settings.TESSERACT_CMD
|
||||
|
||||
with PILImage.open(path) as img:
|
||||
# Convertit en RGB si nécessaire
|
||||
if img.mode not in ("RGB", "L"):
|
||||
img = img.convert("RGB")
|
||||
|
||||
# Extraction avec données de confiance
|
||||
data = pytesseract.image_to_data(
|
||||
img,
|
||||
lang=settings.OCR_LANGUAGES,
|
||||
output_type=pytesseract.Output.DICT,
|
||||
)
|
||||
|
||||
# Calcul de la confiance moyenne (on ignore les -1)
|
||||
confidences = [
|
||||
int(c) for c in data["conf"]
|
||||
if str(c).strip() not in ("-1", "")
|
||||
]
|
||||
avg_confidence = (
|
||||
round(sum(confidences) / len(confidences) / 100, 3)
|
||||
if confidences else 0.0
|
||||
)
|
||||
|
||||
# Texte nettoyé
|
||||
raw_text = pytesseract.image_to_string(
|
||||
img,
|
||||
lang=settings.OCR_LANGUAGES,
|
||||
).strip()
|
||||
|
||||
if raw_text and len(raw_text) > 3:
|
||||
result["text"] = raw_text
|
||||
result["has_text"] = True
|
||||
result["confidence"] = avg_confidence
|
||||
result["language"] = _detect_language(raw_text)
|
||||
else:
|
||||
result["has_text"] = False
|
||||
|
||||
except pytesseract.TesseractNotFoundError:
|
||||
logger.warning("ocr.tesseract_not_found")
|
||||
except Exception as e:
|
||||
logger.error("ocr.extraction_error", extra={"file": file_path, "error": str(e)})
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,207 @@
|
||||
"""
|
||||
Pipeline de traitement AI — orchestration des 3 étapes
|
||||
Chaque étape est indépendante : un échec partiel n'arrête pas le pipeline.
|
||||
|
||||
Publie des événements Redis (si disponible) pour le suivi en temps réel.
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models.image import Image, ProcessingStatus
|
||||
from app.services.exif_service import extract_exif
|
||||
from app.services.ocr_service import extract_text
|
||||
from app.services.ai_vision import analyze_image, extract_text_with_ai
|
||||
import asyncio
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _run_sync_in_thread(func: Any, *args: Any) -> Any:
|
||||
"""Exécute une fonction synchrone dans un thread pour ne pas bloquer l'event loop."""
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(None, func, *args)
|
||||
|
||||
|
||||
async def _publish_event(
|
||||
redis: Any, image_id: int, event: str, data: dict | None = None
|
||||
) -> None:
|
||||
"""Publie un événement sur le channel Redis pipeline:{image_id}."""
|
||||
if redis is None:
|
||||
return
|
||||
try:
|
||||
payload = {"event": event, "image_id": image_id, "timestamp": time.time()}
|
||||
if data:
|
||||
payload["data"] = data
|
||||
await redis.publish(f"pipeline:{image_id}", json.dumps(payload))
|
||||
except Exception:
|
||||
pass # Pub/Sub non critique — ne doit pas bloquer le pipeline
|
||||
|
||||
|
||||
async def process_image_pipeline(
|
||||
image_id: int, db: AsyncSession, redis: Any = None
|
||||
) -> None:
|
||||
"""
|
||||
Pipeline complet de traitement d'une image :
|
||||
1. Extraction EXIF (sync → thread)
|
||||
2. OCR — extraction texte (sync → thread)
|
||||
3. Vision AI — description + tags (async)
|
||||
4. Sauvegarde finale en BDD
|
||||
|
||||
Le statut est mis à jour à chaque étape pour permettre le polling.
|
||||
Publie des événements Redis sur le channel pipeline:{image_id}.
|
||||
"""
|
||||
# ── Chargement de l'image ─────────────────────────────────
|
||||
result = await db.execute(select(Image).where(Image.id == image_id))
|
||||
image = result.scalar_one_or_none()
|
||||
|
||||
if not image:
|
||||
logger.warning("pipeline.image_not_found", extra={"image_id": image_id})
|
||||
return
|
||||
|
||||
# ── Démarrage ─────────────────────────────────────────────
|
||||
image.processing_status = ProcessingStatus.PROCESSING
|
||||
image.processing_started_at = datetime.now(timezone.utc)
|
||||
await db.commit()
|
||||
await db.refresh(image)
|
||||
|
||||
await _publish_event(redis, image_id, "pipeline.started")
|
||||
|
||||
errors: list[str] = []
|
||||
file_path = image.file_path
|
||||
|
||||
# ════════════════════════════════════════════════════════════
|
||||
# ÉTAPE 1 — Extraction EXIF
|
||||
# ════════════════════════════════════════════════════════════
|
||||
try:
|
||||
logger.info("pipeline.step.start", extra={"image_id": image_id, "step": "exif", "step_num": "1/3"})
|
||||
t0 = time.time()
|
||||
exif = await _run_sync_in_thread(extract_exif, file_path)
|
||||
|
||||
image.exif_raw = exif.get("raw")
|
||||
image.exif_make = exif.get("make")
|
||||
image.exif_model = exif.get("model")
|
||||
image.exif_lens = exif.get("lens")
|
||||
image.exif_taken_at = exif.get("taken_at")
|
||||
image.exif_gps_lat = exif.get("gps_lat")
|
||||
image.exif_gps_lon = exif.get("gps_lon")
|
||||
image.exif_altitude = exif.get("altitude")
|
||||
image.exif_iso = exif.get("iso")
|
||||
image.exif_aperture = exif.get("aperture")
|
||||
image.exif_shutter = exif.get("shutter")
|
||||
image.exif_focal = exif.get("focal")
|
||||
image.exif_flash = exif.get("flash")
|
||||
image.exif_orientation = exif.get("orientation")
|
||||
image.exif_software = exif.get("software")
|
||||
|
||||
await db.commit()
|
||||
elapsed = int((time.time() - t0) * 1000)
|
||||
logger.info("pipeline.step.done", extra={"image_id": image_id, "step": "exif", "duration_ms": elapsed, "camera": image.exif_make})
|
||||
await _publish_event(redis, image_id, "step.completed", {
|
||||
"step": "exif", "duration_ms": elapsed, "camera": image.exif_make,
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
msg = f"EXIF : {str(e)}"
|
||||
errors.append(msg)
|
||||
logger.error("pipeline.step.error", extra={"image_id": image_id, "step": "exif", "error": str(e)})
|
||||
|
||||
# ════════════════════════════════════════════════════════════
|
||||
# ÉTAPE 2 — OCR
|
||||
# ════════════════════════════════════════════════════════════
|
||||
try:
|
||||
logger.info("pipeline.step.start", extra={"image_id": image_id, "step": "ocr", "step_num": "2/3"})
|
||||
t0 = time.time()
|
||||
ocr = await _run_sync_in_thread(extract_text, file_path)
|
||||
|
||||
# Fallback AI si OCR classique échoue ou ne trouve rien
|
||||
if not ocr.get("has_text", False):
|
||||
logger.info("pipeline.ocr.fallback", extra={"image_id": image_id, "reason": "tesseract_empty"})
|
||||
ai_ocr = await extract_text_with_ai(file_path)
|
||||
|
||||
if ai_ocr.get("has_text"):
|
||||
ocr = ai_ocr
|
||||
logger.info("pipeline.ocr.fallback_success", extra={"image_id": image_id, "chars": len(ocr.get("text", ""))})
|
||||
else:
|
||||
logger.info("pipeline.ocr.fallback_empty", extra={"image_id": image_id})
|
||||
|
||||
image.ocr_text = ocr.get("text")
|
||||
image.ocr_language = ocr.get("language")
|
||||
image.ocr_confidence = ocr.get("confidence")
|
||||
image.ocr_has_text = ocr.get("has_text", False)
|
||||
|
||||
await db.commit()
|
||||
elapsed = int((time.time() - t0) * 1000)
|
||||
logger.info("pipeline.step.done", extra={"image_id": image_id, "step": "ocr", "duration_ms": elapsed, "has_text": image.ocr_has_text})
|
||||
await _publish_event(redis, image_id, "step.completed", {
|
||||
"step": "ocr", "duration_ms": elapsed, "has_text": image.ocr_has_text,
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
msg = f"OCR : {str(e)}"
|
||||
errors.append(msg)
|
||||
logger.error("pipeline.step.error", extra={"image_id": image_id, "step": "ocr", "error": str(e)})
|
||||
|
||||
# ════════════════════════════════════════════════════════════
|
||||
# ÉTAPE 3 — Vision AI (description + tags)
|
||||
# ════════════════════════════════════════════════════════════
|
||||
try:
|
||||
logger.info("pipeline.step.start", extra={"image_id": image_id, "step": "ai", "step_num": "3/3"})
|
||||
t0 = time.time()
|
||||
ai = await analyze_image(
|
||||
file_path=file_path,
|
||||
ocr_hint=image.ocr_text,
|
||||
)
|
||||
|
||||
image.ai_description = ai.get("description")
|
||||
image.ai_tags = ai.get("tags", [])
|
||||
image.ai_confidence = ai.get("confidence")
|
||||
image.ai_model_used = ai.get("model")
|
||||
image.ai_processed_at = datetime.now(timezone.utc)
|
||||
image.ai_prompt_tokens = ai.get("prompt_tokens")
|
||||
image.ai_output_tokens = ai.get("output_tokens")
|
||||
|
||||
await db.commit()
|
||||
elapsed = int((time.time() - t0) * 1000)
|
||||
logger.info("pipeline.step.done", extra={"image_id": image_id, "step": "ai", "duration_ms": elapsed, "tags_count": len(image.ai_tags or [])})
|
||||
await _publish_event(redis, image_id, "step.completed", {
|
||||
"step": "ai", "duration_ms": elapsed, "tags_count": len(image.ai_tags or []),
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
msg = f"AI Vision : {str(e)}"
|
||||
errors.append(msg)
|
||||
logger.error("pipeline.step.error", extra={"image_id": image_id, "step": "ai", "error": str(e)})
|
||||
|
||||
# ════════════════════════════════════════════════════════════
|
||||
# FINALISATION
|
||||
# ════════════════════════════════════════════════════════════
|
||||
image.processing_done_at = datetime.now(timezone.utc)
|
||||
|
||||
if errors:
|
||||
if image.ai_description:
|
||||
image.processing_status = ProcessingStatus.DONE
|
||||
image.processing_error = f"Avertissements : {'; '.join(errors)}"
|
||||
else:
|
||||
image.processing_status = ProcessingStatus.ERROR
|
||||
image.processing_error = "; ".join(errors)
|
||||
else:
|
||||
image.processing_status = ProcessingStatus.DONE
|
||||
image.processing_error = None
|
||||
|
||||
await db.commit()
|
||||
logger.info("pipeline.completed", extra={
|
||||
"image_id": image_id,
|
||||
"status": image.processing_status.value,
|
||||
"errors": len(errors),
|
||||
})
|
||||
|
||||
if errors:
|
||||
await _publish_event(redis, image_id, "pipeline.error", {"errors": errors})
|
||||
else:
|
||||
await _publish_event(redis, image_id, "pipeline.done")
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
Service de scraping — extraction de contenu web pour résumés AI
|
||||
"""
|
||||
from typing import Optional
|
||||
import httpx
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
|
||||
HEADERS = {
|
||||
"User-Agent": (
|
||||
"Mozilla/5.0 (compatible; ShaarliBot/1.0)"
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
async def fetch_page_content(url: str) -> dict:
|
||||
"""
|
||||
Récupère le contenu d'une URL et extrait :
|
||||
- Titre de la page
|
||||
- Méta description
|
||||
- Texte principal
|
||||
"""
|
||||
result = {
|
||||
"url": url,
|
||||
"title": None,
|
||||
"description": None,
|
||||
"text": None,
|
||||
"error": None,
|
||||
}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
headers=HEADERS,
|
||||
timeout=15.0,
|
||||
follow_redirects=True,
|
||||
) as client:
|
||||
response = await client.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
soup = BeautifulSoup(response.text, "html.parser")
|
||||
|
||||
# Titre
|
||||
title_tag = soup.find("title")
|
||||
result["title"] = title_tag.get_text(strip=True) if title_tag else None
|
||||
|
||||
# Meta description
|
||||
meta_desc = soup.find("meta", attrs={"name": "description"})
|
||||
if not meta_desc:
|
||||
meta_desc = soup.find("meta", attrs={"property": "og:description"})
|
||||
result["description"] = meta_desc.get("content", "") if meta_desc else None
|
||||
|
||||
# Texte principal — on retire scripts, styles, nav
|
||||
for tag in soup(["script", "style", "nav", "footer", "header", "aside"]):
|
||||
tag.decompose()
|
||||
|
||||
# Priorité aux balises sémantiques
|
||||
main = soup.find("article") or soup.find("main") or soup.find("body")
|
||||
if main:
|
||||
paragraphs = main.find_all("p")
|
||||
text = " ".join(p.get_text(strip=True) for p in paragraphs if len(p.get_text(strip=True)) > 30)
|
||||
result["text"] = text[:5000] if text else None
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
result["error"] = f"HTTP {e.response.status_code}"
|
||||
except httpx.RequestError as e:
|
||||
result["error"] = f"Connexion impossible : {str(e)}"
|
||||
except Exception as e:
|
||||
result["error"] = str(e)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,123 @@
|
||||
"""
|
||||
Service de stockage — sauvegarde fichiers, génération thumbnails
|
||||
Multi-tenant : les fichiers sont isolés par client_id.
|
||||
"""
|
||||
import uuid
|
||||
import logging
|
||||
import aiofiles
|
||||
from pathlib import Path
|
||||
from datetime import datetime, timezone
|
||||
from PIL import Image as PILImage
|
||||
from fastapi import UploadFile, HTTPException, status
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
ALLOWED_MIME_TYPES = {
|
||||
"image/jpeg", "image/png", "image/gif",
|
||||
"image/webp", "image/bmp", "image/tiff",
|
||||
}
|
||||
|
||||
THUMBNAIL_SIZE = (320, 320)
|
||||
|
||||
|
||||
def _generate_filename(original: str) -> tuple[str, str]:
|
||||
"""Retourne (uuid_filename, extension)."""
|
||||
suffix = Path(original).suffix.lower() or ".jpg"
|
||||
uid = str(uuid.uuid4())
|
||||
return f"{uid}{suffix}", uid
|
||||
|
||||
|
||||
def _get_client_upload_path(client_id: str) -> Path:
|
||||
"""Retourne le répertoire d'upload pour un client donné."""
|
||||
p = settings.upload_path / client_id
|
||||
p.mkdir(parents=True, exist_ok=True)
|
||||
return p
|
||||
|
||||
|
||||
def _get_client_thumbnails_path(client_id: str) -> Path:
|
||||
"""Retourne le répertoire de thumbnails pour un client donné."""
|
||||
p = settings.thumbnails_path / client_id
|
||||
p.mkdir(parents=True, exist_ok=True)
|
||||
return p
|
||||
|
||||
|
||||
async def save_upload(file: UploadFile, client_id: str) -> dict:
|
||||
"""
|
||||
Valide, sauvegarde le fichier uploadé et génère un thumbnail.
|
||||
Les fichiers sont stockés dans uploads/{client_id}/ pour l'isolation.
|
||||
Retourne un dict avec toutes les métadonnées fichier.
|
||||
"""
|
||||
# ── Validation MIME ───────────────────────────────────────
|
||||
if file.content_type not in ALLOWED_MIME_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
|
||||
detail=f"Type non supporté : {file.content_type}. "
|
||||
f"Acceptés : {', '.join(ALLOWED_MIME_TYPES)}",
|
||||
)
|
||||
|
||||
# ── Lecture du contenu ────────────────────────────────────
|
||||
content = await file.read()
|
||||
|
||||
if len(content) > settings.max_upload_bytes:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"Fichier trop volumineux. Max : {settings.MAX_UPLOAD_SIZE_MB} MB",
|
||||
)
|
||||
|
||||
# ── Nommage et chemins ────────────────────────────────────
|
||||
filename, file_uuid = _generate_filename(file.filename or "image")
|
||||
upload_dir = _get_client_upload_path(client_id)
|
||||
thumb_dir = _get_client_thumbnails_path(client_id)
|
||||
|
||||
file_path = upload_dir / filename
|
||||
thumb_filename = f"thumb_{filename}"
|
||||
thumb_path = thumb_dir / thumb_filename
|
||||
|
||||
# ── Sauvegarde fichier original ───────────────────────────
|
||||
async with aiofiles.open(file_path, "wb") as f:
|
||||
await f.write(content)
|
||||
|
||||
# ── Dimensions + thumbnail ────────────────────────────────
|
||||
width, height = None, None
|
||||
try:
|
||||
with PILImage.open(file_path) as img:
|
||||
width, height = img.size
|
||||
img.thumbnail(THUMBNAIL_SIZE, PILImage.LANCZOS)
|
||||
# Convertit en RGB si nécessaire (ex: PNG RGBA)
|
||||
if img.mode in ("RGBA", "P"):
|
||||
img = img.convert("RGB")
|
||||
img.save(thumb_path, "JPEG", quality=85)
|
||||
except Exception as e:
|
||||
# Thumbnail non bloquant
|
||||
thumb_path = None
|
||||
logger.warning("Erreur génération thumbnail : %s", e)
|
||||
|
||||
return {
|
||||
"uuid": file_uuid,
|
||||
"original_name": file.filename,
|
||||
"filename": filename,
|
||||
"file_path": str(file_path),
|
||||
"thumbnail_path": str(thumb_path) if thumb_path else None,
|
||||
"mime_type": file.content_type,
|
||||
"file_size": len(content),
|
||||
"width": width,
|
||||
"height": height,
|
||||
"uploaded_at": datetime.now(timezone.utc),
|
||||
"client_id": client_id,
|
||||
}
|
||||
|
||||
|
||||
def delete_files(file_path: str, thumbnail_path: str | None = None) -> None:
|
||||
"""Supprime le fichier original et son thumbnail du disque."""
|
||||
for path_str in [file_path, thumbnail_path]:
|
||||
if path_str:
|
||||
p = Path(path_str)
|
||||
if p.exists():
|
||||
p.unlink()
|
||||
|
||||
|
||||
def get_image_url(filename: str, client_id: str, thumb: bool = False) -> str:
|
||||
"""Construit l'URL publique d'une image."""
|
||||
prefix = "thumbnails" if thumb else "uploads"
|
||||
return f"/static/{prefix}/{client_id}/{filename}"
|
||||
@@ -0,0 +1,227 @@
|
||||
"""
|
||||
Abstraction StorageBackend — interface commune pour le stockage de fichiers.
|
||||
|
||||
Deux implémentations :
|
||||
- LocalStorage : fichiers sur disque local + URLs signées HMAC
|
||||
- S3Storage : AWS S3 / MinIO / Cloudflare R2 via aioboto3
|
||||
|
||||
Le reste du code utilise exclusivement get_storage_backend() et l'interface
|
||||
StorageBackend — jamais les classes concrètes directement.
|
||||
"""
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
|
||||
import aiofiles
|
||||
from itsdangerous import URLSafeTimedSerializer, BadSignature, SignatureExpired
|
||||
|
||||
from app.config import settings
|
||||
|
||||
# Singleton backend
|
||||
_backend: "StorageBackend | None" = None
|
||||
|
||||
|
||||
class StorageBackend(ABC):
|
||||
"""Interface abstraite pour le stockage de fichiers."""
|
||||
|
||||
@abstractmethod
|
||||
async def save(self, content: bytes, path: str, content_type: str) -> str:
|
||||
"""Sauvegarde un fichier. Retourne le chemin stocké."""
|
||||
|
||||
@abstractmethod
|
||||
async def delete(self, path: str) -> None:
|
||||
"""Supprime un fichier."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_signed_url(self, path: str, expires_in: int = 900) -> str:
|
||||
"""Retourne une URL d'accès temporaire signée."""
|
||||
|
||||
@abstractmethod
|
||||
async def exists(self, path: str) -> bool:
|
||||
"""Vérifie qu'un fichier existe."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_size(self, path: str) -> int:
|
||||
"""Retourne la taille en bytes."""
|
||||
|
||||
|
||||
class LocalStorage(StorageBackend):
|
||||
"""Stockage sur disque local avec URLs signées HMAC."""
|
||||
|
||||
def __init__(self, base_dir: str, secret: str) -> None:
|
||||
self._base_dir = Path(base_dir)
|
||||
self._serializer = URLSafeTimedSerializer(secret)
|
||||
|
||||
def _full_path(self, path: str) -> Path:
|
||||
return self._base_dir / path
|
||||
|
||||
async def save(self, content: bytes, path: str, content_type: str) -> str:
|
||||
"""Sauvegarde un fichier sur disque."""
|
||||
full = self._full_path(path)
|
||||
full.parent.mkdir(parents=True, exist_ok=True)
|
||||
async with aiofiles.open(full, "wb") as f:
|
||||
await f.write(content)
|
||||
return path
|
||||
|
||||
async def delete(self, path: str) -> None:
|
||||
"""Supprime un fichier du disque."""
|
||||
full = self._full_path(path)
|
||||
if full.exists():
|
||||
full.unlink()
|
||||
|
||||
async def get_signed_url(self, path: str, expires_in: int = 900) -> str:
|
||||
"""Génère un token HMAC signé pour accéder au fichier."""
|
||||
token = self._serializer.dumps({"path": path, "max_age": expires_in})
|
||||
return f"/files/signed/{token}"
|
||||
|
||||
def validate_token(self, token: str) -> str | None:
|
||||
"""Valide un token HMAC et retourne le path, None si invalide/expiré."""
|
||||
try:
|
||||
data = self._serializer.loads(token, max_age=3600)
|
||||
return data.get("path")
|
||||
except (BadSignature, SignatureExpired):
|
||||
return None
|
||||
|
||||
def validate_token_with_max_age(self, token: str, max_age: int = 3600) -> str | None:
|
||||
"""Valide un token HMAC avec un max_age spécifique."""
|
||||
try:
|
||||
data = self._serializer.loads(token, max_age=max_age)
|
||||
return data.get("path")
|
||||
except (BadSignature, SignatureExpired):
|
||||
return None
|
||||
|
||||
async def exists(self, path: str) -> bool:
|
||||
"""Vérifie l'existence du fichier sur disque."""
|
||||
return self._full_path(path).exists()
|
||||
|
||||
async def get_size(self, path: str) -> int:
|
||||
"""Retourne la taille du fichier."""
|
||||
full = self._full_path(path)
|
||||
if full.exists():
|
||||
return full.stat().st_size
|
||||
return 0
|
||||
|
||||
def get_absolute_path(self, path: str) -> Path:
|
||||
"""Retourne le chemin absolu d'un fichier (pour FileResponse)."""
|
||||
return self._full_path(path)
|
||||
|
||||
|
||||
class S3Storage(StorageBackend):
|
||||
"""Stockage S3/MinIO via aioboto3."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bucket: str,
|
||||
prefix: str = "",
|
||||
region: str = "us-east-1",
|
||||
endpoint_url: str = "",
|
||||
access_key: str = "",
|
||||
secret_key: str = "",
|
||||
) -> None:
|
||||
self._bucket = bucket
|
||||
self._prefix = prefix.rstrip("/")
|
||||
self._region = region
|
||||
self._endpoint_url = endpoint_url or None
|
||||
self._access_key = access_key
|
||||
self._secret_key = secret_key
|
||||
|
||||
def _s3_key(self, path: str) -> str:
|
||||
if self._prefix:
|
||||
return f"{self._prefix}/{path}"
|
||||
return path
|
||||
|
||||
def _get_session(self):
|
||||
import aioboto3
|
||||
return aioboto3.Session(
|
||||
aws_access_key_id=self._access_key,
|
||||
aws_secret_access_key=self._secret_key,
|
||||
region_name=self._region,
|
||||
)
|
||||
|
||||
async def save(self, content: bytes, path: str, content_type: str) -> str:
|
||||
"""Upload vers S3/MinIO."""
|
||||
session = self._get_session()
|
||||
async with session.client("s3", endpoint_url=self._endpoint_url) as client:
|
||||
await client.put_object(
|
||||
Bucket=self._bucket,
|
||||
Key=self._s3_key(path),
|
||||
Body=content,
|
||||
ContentType=content_type,
|
||||
)
|
||||
return path
|
||||
|
||||
async def delete(self, path: str) -> None:
|
||||
"""Supprime un objet S3."""
|
||||
session = self._get_session()
|
||||
async with session.client("s3", endpoint_url=self._endpoint_url) as client:
|
||||
await client.delete_object(
|
||||
Bucket=self._bucket,
|
||||
Key=self._s3_key(path),
|
||||
)
|
||||
|
||||
async def get_signed_url(self, path: str, expires_in: int = 900) -> str:
|
||||
"""Génère une URL présignée S3."""
|
||||
session = self._get_session()
|
||||
async with session.client("s3", endpoint_url=self._endpoint_url) as client:
|
||||
url = await client.generate_presigned_url(
|
||||
"get_object",
|
||||
Params={"Bucket": self._bucket, "Key": self._s3_key(path)},
|
||||
ExpiresIn=expires_in,
|
||||
)
|
||||
return url
|
||||
|
||||
async def exists(self, path: str) -> bool:
|
||||
"""Vérifie l'existence via head_object."""
|
||||
session = self._get_session()
|
||||
async with session.client("s3", endpoint_url=self._endpoint_url) as client:
|
||||
try:
|
||||
await client.head_object(
|
||||
Bucket=self._bucket,
|
||||
Key=self._s3_key(path),
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
async def get_size(self, path: str) -> int:
|
||||
"""Retourne la taille via head_object."""
|
||||
session = self._get_session()
|
||||
async with session.client("s3", endpoint_url=self._endpoint_url) as client:
|
||||
try:
|
||||
resp = await client.head_object(
|
||||
Bucket=self._bucket,
|
||||
Key=self._s3_key(path),
|
||||
)
|
||||
return resp.get("ContentLength", 0)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def get_storage_backend() -> StorageBackend:
|
||||
"""Factory : retourne le backend de stockage configuré (singleton)."""
|
||||
global _backend
|
||||
if _backend is not None:
|
||||
return _backend
|
||||
|
||||
if settings.STORAGE_BACKEND == "s3":
|
||||
_backend = S3Storage(
|
||||
bucket=settings.S3_BUCKET,
|
||||
prefix=settings.S3_PREFIX,
|
||||
region=settings.S3_REGION,
|
||||
endpoint_url=settings.S3_ENDPOINT_URL,
|
||||
access_key=settings.S3_ACCESS_KEY,
|
||||
secret_key=settings.S3_SECRET_KEY,
|
||||
)
|
||||
else:
|
||||
_backend = LocalStorage(
|
||||
base_dir=str(settings.upload_path.parent),
|
||||
secret=settings.SIGNED_URL_SECRET,
|
||||
)
|
||||
|
||||
return _backend
|
||||
|
||||
|
||||
def reset_storage_backend() -> None:
|
||||
"""Reset le singleton (utile pour les tests)."""
|
||||
global _backend
|
||||
_backend = None
|
||||
@@ -0,0 +1 @@
|
||||
# Workers package — ARQ task queue
|
||||
@@ -0,0 +1,191 @@
|
||||
"""
|
||||
Worker ARQ — traitement asynchrone des images via Redis.
|
||||
|
||||
Lance avec : python worker.py
|
||||
|
||||
Fonctionnalités :
|
||||
- File persistante Redis (survit aux redémarrages)
|
||||
- Retry automatique avec backoff exponentiel
|
||||
- Queues prioritaires (premium / standard)
|
||||
- Dead-letter : marquage error après max_tries
|
||||
"""
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from arq import cron, func
|
||||
from arq.connections import RedisSettings
|
||||
|
||||
from app.config import settings
|
||||
from app.database import AsyncSessionLocal
|
||||
from app.models.image import Image, ProcessingStatus
|
||||
from app.services.pipeline import process_image_pipeline
|
||||
from sqlalchemy import select
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Backoff exponentiel : délais entre tentatives (en secondes)
|
||||
RETRY_DELAYS = [1, 4, 16]
|
||||
|
||||
|
||||
async def process_image_task(ctx: dict, image_id: int, client_id: str) -> str:
|
||||
"""
|
||||
Tâche ARQ : traite une image via le pipeline EXIF → OCR → AI.
|
||||
|
||||
Args:
|
||||
ctx: Contexte ARQ (contient job_try, redis, etc.)
|
||||
image_id: ID de l'image à traiter
|
||||
client_id: ID du client propriétaire
|
||||
"""
|
||||
job_try = ctx.get("job_try", 1)
|
||||
redis = ctx.get("redis")
|
||||
|
||||
logger.info(
|
||||
"worker.job.started",
|
||||
extra={"image_id": image_id, "client_id": client_id, "job_try": job_try},
|
||||
)
|
||||
|
||||
async with AsyncSessionLocal() as db:
|
||||
try:
|
||||
await process_image_pipeline(image_id, db, redis=redis)
|
||||
logger.info(
|
||||
"worker.job.completed",
|
||||
extra={"image_id": image_id, "client_id": client_id},
|
||||
)
|
||||
return f"OK image_id={image_id}"
|
||||
|
||||
except Exception as e:
|
||||
max_tries = settings.WORKER_MAX_TRIES
|
||||
logger.error(
|
||||
"worker.job.failed",
|
||||
extra={
|
||||
"image_id": image_id,
|
||||
"client_id": client_id,
|
||||
"job_try": job_try,
|
||||
"max_tries": max_tries,
|
||||
"error": str(e),
|
||||
},
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
if job_try >= max_tries:
|
||||
# Dead-letter : marquer l'image en erreur définitive
|
||||
await _mark_image_error(db, image_id, str(e), job_try)
|
||||
logger.error(
|
||||
"worker.job.dead_letter",
|
||||
extra={
|
||||
"image_id": image_id,
|
||||
"client_id": client_id,
|
||||
"total_tries": job_try,
|
||||
},
|
||||
)
|
||||
return f"DEAD_LETTER image_id={image_id} after {job_try} tries"
|
||||
|
||||
# Retry avec backoff
|
||||
delay_idx = min(job_try - 1, len(RETRY_DELAYS) - 1)
|
||||
retry_delay = RETRY_DELAYS[delay_idx]
|
||||
logger.warning(
|
||||
"worker.job.retry_scheduled",
|
||||
extra={
|
||||
"image_id": image_id,
|
||||
"retry_in_seconds": retry_delay,
|
||||
"next_try": job_try + 1,
|
||||
},
|
||||
)
|
||||
raise # ARQ replanifie automatiquement
|
||||
|
||||
|
||||
async def _mark_image_error(
|
||||
db, image_id: int, error_msg: str, total_tries: int
|
||||
) -> None:
|
||||
"""Marque une image en erreur définitive après épuisement des retries."""
|
||||
result = await db.execute(select(Image).where(Image.id == image_id))
|
||||
image = result.scalar_one_or_none()
|
||||
if image:
|
||||
image.processing_status = ProcessingStatus.ERROR
|
||||
image.processing_error = f"Échec après {total_tries} tentatives : {error_msg}"
|
||||
image.processing_done_at = datetime.now(timezone.utc)
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def on_startup(ctx: dict) -> None:
|
||||
"""Hook ARQ : appelé au démarrage du worker."""
|
||||
logger.info("worker.startup", extra={"max_jobs": settings.WORKER_MAX_JOBS})
|
||||
|
||||
|
||||
async def on_shutdown(ctx: dict) -> None:
|
||||
"""Hook ARQ : appelé à l'arrêt du worker."""
|
||||
logger.info("worker.shutdown")
|
||||
|
||||
|
||||
async def on_job_start(ctx: dict) -> None:
|
||||
"""Hook ARQ : appelé au début de chaque job."""
|
||||
pass # Le logging est fait dans process_image_task
|
||||
|
||||
|
||||
async def on_job_end(ctx: dict) -> None:
|
||||
"""Hook ARQ : appelé à la fin de chaque job."""
|
||||
pass # Le logging est fait dans process_image_task
|
||||
|
||||
|
||||
def _parse_redis_settings() -> RedisSettings:
|
||||
"""Parse REDIS_URL en RedisSettings ARQ."""
|
||||
url = settings.REDIS_URL
|
||||
# redis://[:password@]host[:port][/db]
|
||||
if url.startswith("redis://"):
|
||||
url = url[8:]
|
||||
elif url.startswith("rediss://"):
|
||||
url = url[9:]
|
||||
|
||||
password = None
|
||||
host = "localhost"
|
||||
port = 6379
|
||||
database = 0
|
||||
|
||||
# Parse password
|
||||
if "@" in url:
|
||||
auth_part, url = url.rsplit("@", 1)
|
||||
if ":" in auth_part:
|
||||
password = auth_part.split(":", 1)[1]
|
||||
else:
|
||||
password = auth_part
|
||||
|
||||
# Parse host:port/db
|
||||
if "/" in url:
|
||||
host_port, db_str = url.split("/", 1)
|
||||
if db_str:
|
||||
database = int(db_str)
|
||||
else:
|
||||
host_port = url
|
||||
|
||||
if ":" in host_port:
|
||||
host, port_str = host_port.rsplit(":", 1)
|
||||
if port_str:
|
||||
port = int(port_str)
|
||||
else:
|
||||
host = host_port
|
||||
|
||||
return RedisSettings(
|
||||
host=host or "localhost",
|
||||
port=port,
|
||||
password=password,
|
||||
database=database,
|
||||
)
|
||||
|
||||
|
||||
class WorkerSettings:
|
||||
"""Configuration du worker ARQ."""
|
||||
|
||||
functions = [func(process_image_task, name="process_image_task")]
|
||||
redis_settings = _parse_redis_settings()
|
||||
max_jobs = settings.WORKER_MAX_JOBS
|
||||
job_timeout = settings.WORKER_JOB_TIMEOUT
|
||||
retry_jobs = True
|
||||
max_tries = settings.WORKER_MAX_TRIES
|
||||
queue_name = "standard" # Queue par défaut
|
||||
on_startup = on_startup
|
||||
on_shutdown = on_shutdown
|
||||
on_job_start = on_job_start
|
||||
on_job_end = on_job_end
|
||||
|
||||
# Le worker écoute les deux queues
|
||||
queues = ["standard", "premium"]
|
||||
@@ -0,0 +1,28 @@
|
||||
"""
|
||||
Client Redis partagé — pool de connexions async pour ARQ et Pub/Sub.
|
||||
"""
|
||||
from redis.asyncio import ConnectionPool, Redis
|
||||
|
||||
from app.config import settings
|
||||
|
||||
_pool: ConnectionPool | None = None
|
||||
|
||||
|
||||
async def get_redis_pool() -> Redis:
|
||||
"""Retourne un client Redis avec pool de connexions partagé."""
|
||||
global _pool
|
||||
if _pool is None:
|
||||
_pool = ConnectionPool.from_url(
|
||||
settings.REDIS_URL,
|
||||
max_connections=20,
|
||||
decode_responses=True,
|
||||
)
|
||||
return Redis(connection_pool=_pool)
|
||||
|
||||
|
||||
async def close_redis_pool() -> None:
|
||||
"""Ferme proprement le pool de connexions Redis."""
|
||||
global _pool
|
||||
if _pool is not None:
|
||||
await _pool.disconnect()
|
||||
_pool = None
|
||||
Reference in New Issue
Block a user