SantosPatazca
Add IAM, Feedback, and Multimodal modules with updated backend services
a162c90
Raw
History Blame Contribute Delete
6.58 kB
# src/expon/presentation/interfaces/rest/controllers/analysis_controller.py
from fastapi import APIRouter, HTTPException, Depends
from pydantic import BaseModel
from typing import Optional, Dict, Any, Union
import os
from uuid import UUID
from sqlalchemy.orm import Session
from src.expon.shared.infrastructure.dependencies import get_db
from src.expon.iam.infrastructure.authorization.sfs.auth_bearer import get_current_user
from src.expon.presentation.infrastructure.persistence.jpa.repositories.presentation_repository import PresentationRepository
from src.expon.presentation.infrastructure.services.storage.local_storage_service import LocalStorageService
# Servicios de dominio
from src.expon.presentation.domain.services.transcription_service import TranscriptionService
from src.expon.presentation.domain.services.multimodal_service import MultimodalService
router = APIRouter()
asr_service = TranscriptionService()
multimodal = MultimodalService() # ← usamos el modelo MULTIMODAL
HF_MODEL_ID = os.getenv("HF_MODEL_ID", "alexander1010/expon-emotions")
HF_REVISION = os.getenv("HF_REVISION", "multimodal")
def _env_float(name: str, default: float) -> float:
try:
v = os.getenv(name, "")
return float(v) if v not in (None, "", "None") else default
except Exception:
return default
MM_ALPHA = _env_float("MM_ALPHA", 0.10)
MM_TEMP_TEXT = _env_float("MM_TEMP_TEXT", 1.8)
# ========= Schemas =========
class AnalyzeRequest(BaseModel):
presentation_id: str
text: Optional[str] = None # opcional; si viene vacío, usamos ASR
class AnalyzeResponse(BaseModel):
presentation_id: str
transcript: Optional[str]
confidence: float
dominant_emotion: str
emotion_probabilities: Dict[str, float]
alpha: float
temp_text: float
provider: str
feedback: Dict[str, Any] # feedback va por endpoints dedicados
# ========= Helpers =========
def _extract_transcript(asr_result: Union[str, Dict[str, Any], None]) -> Optional[str]:
if asr_result is None:
return None
if isinstance(asr_result, str):
return asr_result
if isinstance(asr_result, dict):
for k in ("text", "transcript", "combined_text", "combined_transcript"):
val = asr_result.get(k)
if isinstance(val, str) and val.strip():
return val
result = asr_result.get("result")
if isinstance(result, dict):
for k in ("text", "transcript"):
val = result.get(k)
if isinstance(val, str) and val.strip():
return val
return None
def _resolve_audio_path(db: Session, user_id: str, presentation_id: str) -> str:
"""
Resuelve la ruta física del audio usando el MISMO repo y storage de /upload.
- En BD guardas 'filename'
- En storage (LocalStorageService.base_path) está el archivo real
"""
# Asegura UUID como en tus otros endpoints
try:
pid = UUID(presentation_id)
except Exception:
raise HTTPException(status_code=400, detail="presentation_id no es un UUID válido")
repo = PresentationRepository(db)
pres = repo.get_by_id_and_user(pid, user_id)
if pres is None:
raise HTTPException(status_code=404, detail="Presentación no encontrada")
filename = getattr(pres, "filename", None)
if not filename:
raise HTTPException(status_code=500, detail="La presentación no tiene filename almacenado")
storage = LocalStorageService() # usa el mismo base_path por defecto que en /upload (/tmp/storage/audio)
audio_path = storage.get_path(filename)
if not os.path.exists(audio_path):
print(f"[analyze] No existe el archivo en storage: {audio_path}")
raise HTTPException(status_code=500, detail="No se encontró el archivo de audio en el almacenamiento")
return audio_path
# ========= Endpoint =========
@router.post("/analyze", response_model=AnalyzeResponse)
async def analyze_presentation(
payload: AnalyzeRequest,
db: Session = Depends(get_db),
user=Depends(get_current_user),
):
"""
JSON:
{
"presentation_id": "<uuid>",
"text": "opcional (si viene, salta ASR)"
}
"""
presentation_id = payload.presentation_id
text = (payload.text or "").strip()
if not text or text.lower() == "string":
text = None
# 1) Ruta del audio desde storage/BD
audio_path = _resolve_audio_path(db, user.id, presentation_id)
# 2) Texto: si no llega, transcribimos el audio (ASR)
if text:
transcript = text
confidence_asr = 1.0
else:
asr_result = asr_service.transcribe(audio_path) # puede devolver str o dict
transcript = _extract_transcript(asr_result)
if not transcript:
raise HTTPException(status_code=500, detail="No se pudo obtener la transcripción del audio.")
confidence_asr = 1.0
if isinstance(asr_result, dict) and "confidence" in asr_result:
try:
confidence_asr = float(asr_result["confidence"])
except Exception:
pass
# 3) Emociones con modelo MULTIMODAL (usa audio + texto opcional)
# Firma asumida: predict(audio_path, text_or_none)
emo = multimodal.predict(audio_path, transcript if transcript else None)
# La salida esperada del servicio:
dominant = emo["dominant_emotion"]
probs = emo["emotion_probabilities"]
alpha = float(emo.get("alpha", MM_ALPHA))
temp_text = float(emo.get("temp_text", MM_TEMP_TEXT))
provider = emo.get("provider", f"hf:{HF_MODEL_ID}@{HF_REVISION}")
# Confianza devuelta: score de la emoción dominante
confidence_dom = float(probs.get(dominant, 0.0))
# (Opcional) Persistir análisis en BD si quieres consultarlo luego sin recalcular:
# try:
# repo = PresentationRepository(db)
# repo.update_analysis(
# presentation_id=UUID(presentation_id),
# transcript=transcript,
# dominant_emotion=dominant,
# emotion_probabilities=probs,
# confidence=confidence_dom,
# )
# except Exception:
# pass # no romper respuesta si falla la persistencia
return AnalyzeResponse(
presentation_id=presentation_id,
transcript=transcript,
confidence=confidence_dom,
dominant_emotion=dominant,
emotion_probabilities=probs,
alpha=alpha,
temp_text=temp_text,
provider=provider,
feedback={}, # el feedback se genera con /feedback
)