CoolFace
Apppublic

Mery3391/YAMNET

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
main.py251 linesDownload Raw Back to root
1from fastapi import FastAPI, UploadFile, File, HTTPException, Depends, Body2from fastapi.responses import FileResponse3from fastapi.staticfiles import StaticFiles4from fastapi.middleware.cors import CORSMiddleware  # Ajoute CORS5from sqlalchemy.orm import Session6import tempfile7import os8from predict import predict_emotion9from database import SessionLocal, User, AudioAnalysis10from datetime import datetime11import hashlib12 13app = FastAPI(title="Speech Emotion API - Darija")14 15# ✅ Ajoute CORS pour permettre les requêtes du frontend16app.add_middleware(17    CORSMiddleware,18    allow_origins=["*"],  # En développement, en production spécifie les origines19    allow_credentials=True,20    allow_methods=["*"],21    allow_headers=["*"],22)23 24ALLOWED_EXT = {".wav", ".mp3", ".ogg", ".flac"}25 26# ── Database dependency ────────────────────────────────────────27def get_db():28    db = SessionLocal()29    try:30        yield db31    finally:32        db.close()33 34# ── Helper functions ───────────────────────────────────────────35def hash_password(password: str) -> str:36    return hashlib.sha256(password.encode()).hexdigest()37 38def verify_password(password: str, hashed: str) -> bool:39    return hash_password(password) == hashed40 41# ── Serve frontend static files ────────────────────────────────42app.mount("/css", StaticFiles(directory="frontend/css"), name="css")43app.mount("/js", StaticFiles(directory="frontend/js"), name="js")44 45# ── HTML pages ────────────────────────────────────────────────46@app.get("/")47async def root():48    return FileResponse("frontend/login.html")49 50@app.get("/login.html")51async def login_page():52    return FileResponse("frontend/login.html")53 54@app.get("/signup.html")55async def signup_page():56    return FileResponse("frontend/signup.html")57 58@app.get("/dashboard")59async def dashboard_page():60    return FileResponse("frontend/index.html")61 62# ── Static assets (images) ────────────────────────────────────63@app.get("/logo.jpg")64def logo():65    return FileResponse("logo.jpg", media_type="image/jpeg")66 67@app.get("/maroc.jpg")68def maroc():69    return FileResponse("maroc.jpg", media_type="image/jpeg")70 71# ── Authentication endpoints (corrigés pour JSON) ────────────────────72@app.post("/auth/signup")73async def signup(74    data: dict = Body(...),75    db: Session = Depends(get_db)76):77    try:78        email = data.get("email")79        password = data.get("password")80        name = data.get("name")81        82        print(f"Signup attempt - Email: {email}, Name: {name}")83        84        # Vérifier si l'utilisateur existe déjà85        existing_user = db.query(User).filter(User.email == email).first()86        if existing_user:87            raise HTTPException(status_code=400, detail="Email already registered")88        89        # Créer le nouvel utilisateur90        hashed_password = hash_password(password)91        new_user = User(92            email=email,93            password_hash=hashed_password,94            name=name95        )96        db.add(new_user)97        db.commit()98        db.refresh(new_user)99        100        print(f"Signup successful for: {email}")101        102        return {103            "status": "success",104            "user": {105                "email": new_user.email,106                "name": new_user.name,107                "id": new_user.id108            }109        }110    except HTTPException:111        raise112    except Exception as e:113        print(f"Signup error: {e}")114        raise HTTPException(status_code=500, detail=str(e))115 116@app.post("/auth/login")117async def login(118    data: dict = Body(...),119    db: Session = Depends(get_db)120):121    try:122        email = data.get("email")123        password = data.get("password")124        125        print(f"Login attempt - Email: {email}")126        127        user = db.query(User).filter(User.email == email).first()128        if not user or not verify_password(password, user.password_hash):129            raise HTTPException(status_code=401, detail="Invalid credentials")130        131        print(f"Login successful for: {email}")132        133        return {134            "status": "success",135            "user": {136                "email": user.email,137                "name": user.name,138                "id": user.id139            }140        }141    except HTTPException:142        raise143    except Exception as e:144        print(f"Login error: {e}")145        raise HTTPException(status_code=500, detail=str(e))146# ── Save analysis endpoint ────────────────────────────────────147@app.post("/save-analysis")148async def save_analysis(149    data: dict = Body(...),150    db: Session = Depends(get_db)151):152    user = db.query(User).filter(User.email == data.get("user_email")).first()153    if not user:154        raise HTTPException(status_code=404, detail="User not found")155    156    analysis = AudioAnalysis(157        user_id=user.id,158        filename=data.get("filename"),159        emotion=data.get("emotion"),160        confidence=data.get("confidence"),161        angry_score=data.get("angry_score", 0),162        happy_score=data.get("happy_score", 0),163        neutral_score=data.get("neutral_score", 0),164        sad_score=data.get("sad_score", 0)165    )166    db.add(analysis)167    db.commit()168    169    return {"status": "success", "analysis_id": analysis.id}170 171@app.get("/user-history/{user_email}")172async def get_history(user_email: str, db: Session = Depends(get_db)):173    user = db.query(User).filter(User.email == user_email).first()174    if not user:175        return {"history": []}176    177    analyses = db.query(AudioAnalysis).filter(178        AudioAnalysis.user_id == user.id179    ).order_by(AudioAnalysis.created_at.desc()).all()180    181    return {182        "history": [183            {184                "id": a.id,185                "filename": a.filename,186                "emotion": a.emotion,187                "confidence": a.confidence,188                "date": a.created_at.isoformat()189            }190            for a in analyses191        ]192    }193 194@app.delete("/delete-analysis/{analysis_id}")195async def delete_analysis(analysis_id: int, user_email: str, db: Session = Depends(get_db)):196    user = db.query(User).filter(User.email == user_email).first()197    if not user:198        raise HTTPException(status_code=404, detail="User not found")199    200    analysis = db.query(AudioAnalysis).filter(201        AudioAnalysis.id == analysis_id,202        AudioAnalysis.user_id == user.id203    ).first()204    205    if not analysis:206        raise HTTPException(status_code=404, detail="Analysis not found")207    208    db.delete(analysis)209    db.commit()210    211    return {"status": "success"}212 213# ── Prediction ────────────────────────────────────────────────214@app.post("/predict")215async def predict(file: UploadFile = File(...)):216    ext = os.path.splitext(file.filename)[1].lower()217    if ext not in ALLOWED_EXT:218        raise HTTPException(status_code=400, detail="Format audio non supporté")219 220    contents = await file.read()221 222    if not contents:223        raise HTTPException(status_code=400, detail="Fichier audio vide")224 225    tmp_path = None226 227    try:228        with tempfile.NamedTemporaryFile(delete=False, suffix=ext) as tmp:229            tmp.write(contents)230            tmp_path = tmp.name231 232        if os.path.getsize(tmp_path) < 1000:233            raise HTTPException(status_code=400, detail="Fichier audio trop petit")234 235        result = predict_emotion(tmp_path)236 237        return {238            "filename": file.filename,239            "status": "success",240            **result241        }242 243    except ValueError as e:244        raise HTTPException(status_code=400, detail=str(e))245 246    except Exception as e:247        raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}")248 249    finally:250        if tmp_path and os.path.exists(tmp_path):251            os.unlink(tmp_path)