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 @router.post("/register", response_model=schemas.UserResponse) 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 @router.post("/login", response_model=schemas.Token) 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"} @router.get("/me", response_model=schemas.UserResponse) 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 @router.post("/google", response_model=schemas.Token) 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"}, )