Spaces:
Runtime error
Runtime error
| from datetime import datetime, timedelta | |
| from typing import Optional | |
| from fastapi import APIRouter, Depends, HTTPException, status | |
| from fastapi.security import OAuth2PasswordBearer, OAuth2PasswordRequestForm | |
| from jose import JWTError, jwt | |
| import bcrypt | |
| from sqlalchemy.orm import Session | |
| from database import get_db | |
| import models.db_models as db_models | |
| import models.schemas as schemas | |
| import os | |
| from google.oauth2 import id_token | |
| from google.auth.transport import requests | |
| # Configuration | |
| SECRET_KEY = os.getenv("JWT_SECRET_KEY", "09d25e094faa6ca2556c818166b7a9563b93f7099f6f0f4caa6cf63b88e8d3e7") | |
| GOOGLE_CLIENT_ID = os.getenv("GOOGLE_CLIENT_ID", "YOUR_GOOGLE_CLIENT_ID") | |
| ALGORITHM = "HS256" | |
| ACCESS_TOKEN_EXPIRE_MINUTES = 60 * 24 * 7 # 1 week | |
| oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False) | |
| router = APIRouter(prefix="/auth", tags=["Authentication"]) | |
| def verify_password(plain_password: str, hashed_password: str) -> bool: | |
| # Handle the case where passlib originally generated the hash | |
| try: | |
| return bcrypt.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8')) | |
| except ValueError: | |
| return False | |
| def get_password_hash(password: str) -> str: | |
| return bcrypt.hashpw(password.encode('utf-8'), bcrypt.gensalt()).decode('utf-8') | |
| def create_access_token(data: dict, expires_delta: Optional[timedelta] = None): | |
| to_encode = data.copy() | |
| if expires_delta: | |
| expire = datetime.utcnow() + expires_delta | |
| else: | |
| expire = datetime.utcnow() + timedelta(minutes=15) | |
| to_encode.update({"exp": expire}) | |
| encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM) | |
| return encoded_jwt | |
| def get_current_user(token: str = Depends(oauth2_scheme), db: Session = Depends(get_db)): | |
| if not token or token == "null" or token == "undefined": | |
| return None | |
| credentials_exception = HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Could not validate credentials", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| try: | |
| payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM]) | |
| username: str = payload.get("sub") | |
| if username is None: | |
| raise credentials_exception | |
| token_data = schemas.TokenData(username=username) | |
| except JWTError: | |
| raise credentials_exception | |
| user = db.query(db_models.User).filter(db_models.User.username == token_data.username).first() | |
| if user is None: | |
| raise credentials_exception | |
| return user | |
| def register_user(user: schemas.UserCreate, db: Session = Depends(get_db)): | |
| db_user = db.query(db_models.User).filter( | |
| (db_models.User.username == user.username) | (db_models.User.email == user.email) | |
| ).first() | |
| if db_user: | |
| raise HTTPException(status_code=400, detail="Username or email already registered") | |
| hashed_password = get_password_hash(user.password) | |
| db_user = db_models.User( | |
| username=user.username, | |
| email=user.email, | |
| hashed_password=hashed_password | |
| ) | |
| db.add(db_user) | |
| db.commit() | |
| db.refresh(db_user) | |
| return db_user | |
| def login_for_access_token(form_data: OAuth2PasswordRequestForm = Depends(), db: Session = Depends(get_db)): | |
| user = db.query(db_models.User).filter(db_models.User.username == form_data.username).first() | |
| if not user or not verify_password(form_data.password, user.hashed_password): | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Incorrect username or password", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) | |
| access_token = create_access_token( | |
| data={"sub": user.username}, expires_delta=access_token_expires | |
| ) | |
| return {"access_token": access_token, "token_type": "bearer"} | |
| def read_users_me(current_user: db_models.User = Depends(get_current_user)): | |
| if not current_user: | |
| raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated") | |
| return current_user | |
| def google_auth(google_token: schemas.GoogleToken, db: Session = Depends(get_db)): | |
| try: | |
| # Verify the Google token | |
| idinfo = id_token.verify_oauth2_token( | |
| google_token.token, | |
| requests.Request(), | |
| GOOGLE_CLIENT_ID, | |
| clock_skew_in_seconds=10 | |
| ) | |
| email = idinfo['email'] | |
| name = idinfo.get('name', email.split('@')[0]) | |
| # Check if user exists | |
| user = db.query(db_models.User).filter(db_models.User.email == email).first() | |
| if not user: | |
| # Create a new user if they don't exist | |
| # Generate a random password since they use Google to login | |
| import secrets | |
| random_password = secrets.token_urlsafe(32) | |
| hashed_password = get_password_hash(random_password) | |
| # Ensure username is unique | |
| base_username = name.lower().replace(" ", "") | |
| username = base_username | |
| counter = 1 | |
| while db.query(db_models.User).filter(db_models.User.username == username).first(): | |
| username = f"{base_username}{counter}" | |
| counter += 1 | |
| user = db_models.User( | |
| username=username, | |
| email=email, | |
| hashed_password=hashed_password | |
| ) | |
| db.add(user) | |
| db.commit() | |
| db.refresh(user) | |
| access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) | |
| access_token = create_access_token( | |
| data={"sub": user.username}, expires_delta=access_token_expires | |
| ) | |
| return {"access_token": access_token, "token_type": "bearer"} | |
| except ValueError as e: | |
| # Invalid token | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail=f"Invalid Google token: {str(e)}", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |