CoolFace
Apppublic

Hammad712/Auth

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
auth.py224 linesDownload Raw Back to root
1# auth.py2import os3import uuid4import logging5from datetime import datetime, timedelta6from urllib.parse import quote_plus7from typing import Optional8 9from dotenv import load_dotenv10from fastapi import APIRouter, HTTPException, Depends, Request, UploadFile, File, Form11from fastapi.responses import StreamingResponse12from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm13from jose import JWTError, jwt14from passlib.context import CryptContext15from pymongo import MongoClient16import gridfs17from bson import ObjectId  # Ensure ObjectId is imported18 19from models import User, UserUpdate, Token, LoginResponse20from config import CONNECTION_STRING, SECRET_KEY, ACCESS_TOKEN_EXPIRE_MINUTES, REFRESH_TOKEN_EXPIRE_DAYS21 22load_dotenv()23 24logger = logging.getLogger("uvicorn")25logger.setLevel(logging.INFO)26 27# Updated MongoDB initialization: now using CONNECTION_STRING from config.py28client = MongoClient(CONNECTION_STRING)29db = client.users_database30users_collection = db.users31# GridFS instance for storing avatars32fs = gridfs.GridFS(db, collection="avatars")33 34# OAuth2 setup35oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")36router = APIRouter(prefix="/auth", tags=["auth"])37 38# Password hashing39pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")40 41def verify_password(plain_password: str, hashed_password: str) -> bool:42    return pwd_context.verify(plain_password, hashed_password)43 44def get_password_hash(password: str) -> str:45    return pwd_context.hash(password)46 47def get_user(email: str) -> Optional[dict]:48    return users_collection.find_one({"email": email})49 50def authenticate_user(email: str, password: str) -> Optional[dict]:51    user = get_user(email)52    if not user or not verify_password(password, user["hashed_password"]):53        return None54    return user55 56def create_token(data: dict, expires_delta: timedelta = None) -> str:57    to_encode = data.copy()58    expire = datetime.utcnow() + (expires_delta or timedelta(minutes=15))59    to_encode.update({"exp": expire})60    algorithm = "HS256"61    return jwt.encode(to_encode, SECRET_KEY, algorithm=algorithm)62 63def create_access_token(email: str) -> str:64    return create_token({"sub": email}, timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES))65 66def create_refresh_token(email: str) -> str:67    return create_token({"sub": email}, timedelta(days=REFRESH_TOKEN_EXPIRE_DAYS))68 69def get_current_user(token: str = Depends(oauth2_scheme)) -> dict:70    try:71        payload = jwt.decode(token, SECRET_KEY, algorithms=["HS256"])72        email: str = payload.get("sub")73        if not email:74            raise HTTPException(status_code=401, detail="Invalid credentials")75        user = get_user(email)76        if not user:77            raise HTTPException(status_code=401, detail="User not found")78        return user79    except JWTError:80        raise HTTPException(status_code=401, detail="Invalid token")81 82async def save_avatar_file_to_gridfs(file: UploadFile) -> str:83    allowed_types = ["image/jpeg", "image/png", "image/gif"]84    if file.content_type not in allowed_types:85        logger.error(f"Unsupported file type: {file.content_type}")86        raise HTTPException(87            status_code=400,88            detail="Invalid image format. Only JPEG, PNG, and GIF are accepted."89        )90    try:91        contents = await file.read()92        file_id = fs.put(contents, filename=file.filename, contentType=file.content_type)93        logger.info(f"Avatar stored in GridFS with file_id: {file_id}")94        return str(file_id)95    except Exception as e:96        logger.exception("Failed to store avatar in GridFS")97        raise HTTPException(status_code=500, detail="Could not store avatar file in MongoDB.")98 99@router.post("/signup", response_model=Token)100async def signup(101    request: Request,102    name: str = Form(...),103    email: str = Form(...),104    password: str = Form(...),105    role: str = Form(...),  # <-- MODIFICATION: Added role106    avatar: Optional[UploadFile] = File(None)107):108    try:109        _ = User(name=name, email=email, password=password)110    except Exception as e:111        logger.error(f"Validation error during signup: {e}")112        raise HTTPException(status_code=400, detail=str(e))113    if get_user(email):114        logger.warning(f"Attempt to register already existing email: {email}")115        raise HTTPException(status_code=400, detail="Email already registered")116    117    hashed_password = get_password_hash(password)118    119    user_data = {120        "name": name,121        "email": email,122        "hashed_password": hashed_password,123        "role": role,  # <-- MODIFICATION: Added role to user data124        "chat_histories": []125    }126    127    if avatar:128        file_id = await save_avatar_file_to_gridfs(avatar)129        user_data["avatar"] = file_id130        131    users_collection.insert_one(user_data)132    logger.info(f"New user registered: {email} with role: {role}")133    134    return {135        "access_token": create_access_token(email),136        "refresh_token": create_refresh_token(email),137        "token_type": "bearer"138    }139 140@router.post("/login", response_model=LoginResponse)141async def login(request: Request, form_data: OAuth2PasswordRequestForm = Depends()):142    user = authenticate_user(form_data.username, form_data.password)143    if not user:144        logger.warning(f"Failed login attempt for: {form_data.username}")145        raise HTTPException(status_code=401, detail="Incorrect username or password")146    147    logger.info(f"User logged in: {user['email']}")148    149    avatar_url = None150    if "avatar" in user and user["avatar"]:151        avatar_url = f"/auth/avatar/{user['avatar']}"152        153    return {154        "access_token": create_access_token(user["email"]),155        "refresh_token": create_refresh_token(user["email"]),156        "token_type": "bearer",157        "name": user["name"],158        "avatar": avatar_url,159        "role": user.get("role", "user")  # <-- MODIFICATION: Return role (default to "user")160    }161 162@router.get("/user/data")163async def get_user_data(request: Request, current_user: dict = Depends(get_current_user)):164    avatar_url = None165    if "avatar" in current_user and current_user["avatar"]:166        avatar_url = f"/auth/avatar/{current_user['avatar']}"167        168    return {169        "name": current_user["name"],170        "email": current_user["email"],171        "avatar": avatar_url,172        "role": current_user.get("role", "user"),  # <-- MODIFICATION: Return role (default to "user")173        "chat_histories": current_user.get("chat_histories", [])174    }175 176@router.put("/user/update")177async def update_user(178    request: Request,179    name: Optional[str] = Form(None),180    email: Optional[str] = Form(None),181    password: Optional[str] = Form(None),182    avatar: Optional[UploadFile] = File(None),183    current_user: dict = Depends(get_current_user)184):185    update_data = {}186    if name is not None:187        update_data["name"] = name188    if email is not None:189        update_data["email"] = email190    if password is not None:191        try:192            _ = User(name=current_user["name"], email=current_user["email"], password=password)193        except Exception as e:194            logger.error(f"Password validation error during update: {e}")195            raise HTTPException(status_code=400, detail=str(e))196        update_data["hashed_password"] = get_password_hash(password)197        198    if avatar:199        file_id = await save_avatar_file_to_gridfs(avatar)200        update_data["avatar"] = file_id201        202    if not update_data:203        logger.info("No update parameters provided")204        raise HTTPException(status_code=400, detail="No update parameters provided")205        206    users_collection.update_one({"email": current_user["email"]}, {"$set": update_data})207    logger.info(f"User updated: {current_user['email']}")208    209    return {"message": "User updated successfully"}210 211@router.post("/logout")212async def logout(request: Request, current_user: dict = Depends(get_current_user)):213    logger.info(f"User logged out: {current_user['email']}")214    return {"message": "User logged out successfully"}215 216@router.get("/avatar/{file_id}")217async def get_avatar(file_id: str):218    try:219        # Convert the file_id string to an ObjectId before fetching220        file = fs.get(ObjectId(file_id))221        return StreamingResponse(file, media_type=file.content_type)222    except Exception as e:223        logger.error(f"Avatar not found for file_id {file_id}: {e}")224        raise HTTPException(status_code=404, detail="Avatar not found")