CoolFace
Apppublic

Xasik/Form-Doctor-AI

sourceHugging Facemitupdated 9mo agoView on Hugging Face
0likes
backend_api.py382 linesDownload Raw Back to root
1# backend_api.py - ИСПРАВЛЕННАЯ ПОЛНАЯ ВЕРСИЯ (БЕЗ ДУБЛЕЙ И КОНФЛИКТОВ)2 3from fastapi import FastAPI, UploadFile, File, HTTPException, Depends, status4from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm5from sqlalchemy.orm import Session6from passlib.context import CryptContext7from pydantic import BaseModel, field_validator8from typing import Optional, List9import shutil10import os11import sys12import tempfile13import cv214import mediapipe as mp15import numpy as np16import base6417import google.generativeai as genai # ИСПОЛЬЗУЕМ СТАБИЛЬНУЮ БИБЛИОТЕКУ18import ffmpeg 19import traceback 20import hashlib21from sqlalchemy import create_engine, Column, Integer, String, DateTime, ForeignKey, Float22from sqlalchemy.orm import sessionmaker, declarative_base, relationship23from datetime import datetime, timedelta24from jose import JWTError, jwt25from PIL import Image # НУЖНО ДЛЯ РАБОТЫ С КАРТИНКАМИ26 27# --- 1. КОНФИГУРАЦИЯ ---28SECRET_KEY = "YOUR_SUPER_SECRET_KEY_CHANGE_THIS" 29ALGORITHM = "HS256"30ACCESS_TOKEN_EXPIRE_MINUTES = 30031GEMINI_API_KEY = os.environ.get("GEMINI_API_KEY", "AIzaSyCi8ygU26S8xSeHENw2mzQjCZk0_lCDQgw")32 33# Настройка AI (Стабильная версия)34try:35    genai.configure(api_key=GEMINI_API_KEY)36except:37    pass38 39# Настройка пути к БД40if sys.platform.startswith('win'):41    DATABASE_URL = "sqlite:///./sql_app.db"42else:43    DATABASE_URL = "sqlite:////tmp/sql_app.db"44 45engine = create_engine(DATABASE_URL, connect_args={"check_same_thread": False})46Base = declarative_base()47SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)48 49# --- 2. МОДЕЛИ ДАННЫХ ---50class User(Base):51    __tablename__ = "users"52    id = Column(Integer, primary_key=True, index=True)53    email = Column(String, unique=True, index=True)54    hashed_password = Column(String)55    language_code = Column(String, default="ru")56    history_items = relationship("AnalysisHistory", back_populates="owner")57 58class AnalysisHistory(Base):59    __tablename__ = "history"60    id = Column(Integer, primary_key=True, index=True)61    user_id = Column(Integer, ForeignKey("users.id"))62    exercise = Column(String)63    analysis_date = Column(DateTime, default=datetime.utcnow)64    min_angle = Column(Float)65    report_text = Column(String)66    owner = relationship("User", back_populates="history_items")67    68Base.metadata.create_all(bind=engine)69 70# --- 3. СХЕМЫ PYDANTIC ---71class UserCreate(BaseModel):72    email: str73    password: str 74    language_code: Optional[str] = "ru"75    76    @field_validator('password', mode='before')77    @classmethod78    def check_password_length(cls, value):79        if len(value.encode('utf-8')) > 72:80            return value[:72]81        return value82 83class UserInDB(BaseModel):84    id: int85    email: str86    language_code: str87 88class Token(BaseModel):89    access_token: str90    token_type: str91 92class HistoryItem(BaseModel):93    exercise: str94    min_angle: float95    report_text: str96    analysis_date: datetime97    class Config: from_attributes = True98 99# --- 4. ИНИЦИАЛИЗАЦИЯ ---100mp_drawing = mp.solutions.drawing_utils101mp_pose = mp.solutions.pose102app = FastAPI(title="Form Doctor AI Backend")103pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")104oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")105 106# --- 5. ФУНКЦИИ БЕЗОПАСНОСТИ ---107def get_db():108    db = SessionLocal()109    try:110        yield db111    finally:112        db.close()113 114def get_user_by_email(db: Session, email: str):115    return db.query(User).filter(User.email == email).first()116 117def get_password_hash(password: str):118    return hashlib.sha256(password.encode()).hexdigest()119 120def verify_password(plain_password, hashed_password):121    return hashlib.sha256(plain_password.encode()).hexdigest() == hashed_password122 123def create_access_token(data: dict):124    to_encode = data.copy()125    expire = datetime.utcnow() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)126    to_encode.update({"exp": expire})127    return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)128 129def get_current_user(token: str = Depends(oauth2_scheme), db: Session = Depends(get_db)):130    credentials_exception = HTTPException(131        status_code=status.HTTP_401_UNAUTHORIZED,132        detail="Could not validate credentials",133        headers={"WWW-Authenticate": "Bearer"},134    )135    try:136        payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])137        email: str = payload.get("sub")138        if email is None: raise credentials_exception139    except JWTError: raise credentials_exception140    user = get_user_by_email(db, email=email)141    if user is None: raise credentials_exception142    return user143 144# --- 6. ФУНКЦИИ CV ---145def convert_to_mp4(input_path, output_path):146    try:147        (148            ffmpeg149            .input(input_path)150            .output(output_path, vcodec='libx264', acodec='aac', preset='medium', pix_fmt='yuv420p', vf='scale=1280:-2')151            .overwrite_output()152            .run(capture_stdout=True, capture_stderr=True)153        )154        return True155    except:156        return False157 158def calculate_angle(a, b, c):159    a = np.array(a)160    b = np.array(b)161    c = np.array(c)162    radians = np.arctan2(c[1]-b[1], c[0]-b[0]) - np.arctan2(a[1]-b[1], a[0]-b[0])163    angle = np.abs(radians*180.0/np.pi)164    if angle > 180.0:165        angle = 360 - angle166    return angle167 168def analyze_frame_results(image):169    with mp_pose.Pose(min_detection_confidence=0.5, min_tracking_confidence=0.5) as pose:170        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)171        image.flags.writeable = False172        results = pose.process(image)173        return results, image174 175def draw_visual_report(image, landmarks, knee, back, elbow):176    mp_drawing.draw_landmarks(177        image, landmarks, mp_pose.POSE_CONNECTIONS,178        mp_drawing.DrawingSpec(color=(245,117,66), thickness=2, circle_radius=2),179        mp_drawing.DrawingSpec(color=(245,66,230), thickness=2, circle_radius=2)180    )181    font = cv2.FONT_HERSHEY_SIMPLEX182    cv2.rectangle(image, (10, 10), (250, 110), (0,0,0), -1)183    cv2.putText(image, f"Knee: {int(knee)}", (20, 40), font, 0.7, (255, 255, 255), 2)184    cv2.putText(image, f"Back: {int(back)}", (20, 70), font, 0.7, (255, 255, 255), 2)185    cv2.putText(image, f"Elbow: {int(elbow)}", (20, 100), font, 0.7, (255, 255, 255), 2)186    return image187 188# Геттеры углов (безопасные)189def get_knee_angle(lm):190    try: return calculate_angle([lm[24].x, lm[24].y], [lm[26].x, lm[26].y], [lm[28].x, lm[28].y])191    except: return 180192def get_torso_back_angle(lm):193    try: return calculate_angle([lm[12].x, lm[12].y], [lm[24].x, lm[24].y], [lm[26].x, lm[26].y])194    except: return 180195def get_elbow_angle(lm):196    try: return calculate_angle([lm[12].x, lm[12].y], [lm[14].x, lm[14].y], [lm[16].x, lm[16].y])197    except: return 180198 199def recognize_exercise(angles_history, torso_back_angles_history):200    if not angles_history or not torso_back_angles_history: return "Неизвестное упражнение"201    min_knee = min(angles_history)202    min_torso = min(torso_back_angles_history)203    204    if min_knee < 95: return "Приседания"205    elif min_knee > 130 and min_torso < 140: return "Становая тяга"206    elif min_torso > 165: return "Армейский жим"207    return "Неизвестное упражнение"208 209# --- 7. ДВИЖОК АНАЛИЗА ---210def analyze_video(video_path, api_key): 211    cap = cv2.VideoCapture(video_path)212    213    angles_history = []214    torso_back_angles = []215    216    # Сбор кадров для ИИ217    key_frames_pil = [] 218    219    min_knee = 180; min_back = 180; min_elbow = 180220    visual_proof_b64 = None221    222    frame_count = 0223    total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))224    mid_frame_idx = int(total_frames / 2)225    226    while cap.isOpened():227        ret, frame = cap.read()228        if not ret: break229        230        # Оптимизация: каждый 3-й кадр231        if frame_count % 3 != 0:232            if frame_count != mid_frame_idx:233                frame_count += 1234                continue235 236        results, image = analyze_frame_results(frame)237        238        # Сбор 3-х кадров для Gemini (Формат Pillow Image для старой библиотеки)239        if frame_count in [0, mid_frame_idx, total_frames-5] or len(key_frames_pil) < 3:240             img_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)241             try:242                 pil_im = Image.fromarray(img_rgb)243                 if len(key_frames_pil) < 3:244                    key_frames_pil.append(pil_im)245             except: pass246 247        try:248            lm = results.pose_landmarks.landmark249            k = get_knee_angle(lm); b = get_torso_back_angle(lm); e = get_elbow_angle(lm)250            251            if k: angles_history.append(k); min_knee = min(min_knee, k)252            if b: torso_back_angles.append(b); min_back = min(min_back, b)253            if e: min_elbow = min(min_elbow, e)254            255            # Визуализация256            if frame_count == mid_frame_idx or (abs(frame_count - mid_frame_idx) < 5 and not visual_proof_b64):257                annotated = draw_visual_report(image.copy(), results.pose_landmarks, k, b, e)258                _, buf_ann = cv2.imencode('.jpg', annotated)259                visual_proof_b64 = base64.b64encode(buf_ann).decode('utf-8')260 261        except: pass262        frame_count += 1263    cap.release()264 265    if not visual_proof_b64 and key_frames_pil:266        # Если не смогли создать пруф, пробуем хоть что-то267        pass 268 269    if not angles_history: return 0, None, "Не удалось обнаружить человека.", "Error"270 271    exercise = recognize_exercise(angles_history, torso_back_angles)272    273    try:274        # --- ИИ ЧАСТЬ (СТАБИЛЬНАЯ БИБЛИОТЕКА) ---275        model = genai.GenerativeModel('gemini-1.5-flash')276        277        telemetry = f"Min Knee: {min_knee:.0f}, Min Back: {min_back:.0f}, Min Elbow: {min_elbow:.0f}"278        279        prompt = f"""280        Твоя роль: Эксперт по биомеханике (Powerlifting/Fitness).281        282        ДАННЫЕ:283        1. Телеметрия (углы): {telemetry}284        2. Видеоряд (см. кадры).285        3. Алгоритм предполагает: {exercise}.286 287        ЗАДАЧА:288        1. Идентифицируй упражнение.289        2. Оцени технику (0-100).290        3. Дай профессиональные рекомендации.291        292        ОТВЕТ (Markdown):293        ## 🏋️ [Название]294        **Оценка: [X]/100**295        ### 🔬 Анализ:296        ...297        ### ⚠️ Ошибки:298        ...299        ### 💡 Совет PRO:300        ...301        """302        303        # Старая библиотека принимает список [текст, картинка1, картинка2...]304        contents = [prompt] + key_frames_pil305        306        response = model.generate_content(contents)307        308        return min_knee, visual_proof_b64, response.text, "Auto"309 310    except Exception as e:311        return 0, None, f"Ошибка ИИ: {e}", "Error"312 313# --- 8. ЭНДПОИНТЫ (API) ---314 315@app.post("/register", response_model=UserInDB)316def register_new_user(user: UserCreate, db: Session = Depends(get_db)):317    clean_email = user.email.lower()318    if get_user_by_email(db, email=clean_email): 319        raise HTTPException(status_code=400, detail="Email занят.")320    hashed_password = get_password_hash(user.password)321    db_user = User(email=clean_email, hashed_password=hashed_password, language_code=user.language_code)322    db.add(db_user); db.commit(); db.refresh(db_user)323    return db_user324 325# ЗДЕСЬ БЫЛ ОШИБОЧНЫЙ ДУБЛИКАТ - ОН УДАЛЕН326 327@app.post("/token", response_model=Token)328def login(form_data: OAuth2PasswordRequestForm = Depends(), db: Session = Depends(get_db)):329    clean_email = form_data.username.lower()330    user = get_user_by_email(db, email=clean_email)331    if not user: raise HTTPException(status_code=400, detail="Неверный email")332    if get_password_hash(form_data.password) != user.hashed_password:333        raise HTTPException(status_code=400, detail="Неверный пароль")334    access_token = create_access_token(data={"sub": user.email})335    return {"access_token": access_token, "token_type": "bearer"}336 337@app.post("/analyze_form")338async def analyze_form_endpoint(339    video_file: UploadFile = File(...), 340    current_user: User = Depends(get_current_user), 341    db: Session = Depends(get_db)342):343    temp_input = ""; temp_output = ""344    try:345        with tempfile.NamedTemporaryFile(delete=False, suffix=".tmp") as tmp:346            shutil.copyfileobj(video_file.file, tmp)347            temp_input = tmp.name348        temp_output = temp_input + ".mp4"349 350        if not convert_to_mp4(temp_input, temp_output): temp_output = temp_input 351        352        deepest_angle, img_b64, report, name = analyze_video(temp_output, GEMINI_API_KEY)353        354        if report and "Ошибка" in report: raise HTTPException(status_code=500, detail=report)355 356        history_item = AnalysisHistory(357            user_id=current_user.id,358            exercise=name, 359            min_angle=deepest_angle,360            report_text=report[:3000] 361        )362        db.add(history_item); db.commit()363 364        return {365            "status": "success", 366            "min_angle": f"{deepest_angle:.0f}", 367            "error_image_base64": img_b64, 368            "report": report369        }370 371    except Exception as e:372        import traceback373        if os.path.exists(temp_input): os.unlink(temp_input)374        if os.path.exists(temp_output): os.unlink(temp_output)375        raise HTTPException(status_code=500, detail=f"Error: {traceback.format_exc()}")376    finally:377        if os.path.exists(temp_input): os.unlink(temp_input)378        if os.path.exists(temp_output) and temp_output != temp_input: os.unlink(temp_output)379 380@app.get("/history", response_model=List[HistoryItem])381def read_history(current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):382    return db.query(AnalysisHistory).filter(AnalysisHistory.user_id == current_user.id).all()