Files
stock-scanner/backend/app/api/auth.py

117 lines
4.2 KiB
Python

from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from sqlalchemy.orm import Session
from pydantic import BaseModel
from jose import JWTError, jwt
from app.db.database import get_db
from app.db.models import User
from app.core.security import (
public_key, decrypt_password, verify_password,
get_password_hash, create_access_token
)
from app.core.config import SECRET_KEY, ALGORITHM
router = APIRouter()
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="api/auth/login")
class AuthRequest(BaseModel):
username: str
password: str # Base64 encoded RSA-encrypted password
class ProfileUpdateRequest(BaseModel):
display_name: str = None
email_id: str = None
mobile_no: str = None
gender: str = None
password: str = None # Base64 encoded RSA-encrypted new password
class SectorMapUpdateRequest(BaseModel):
sector_map: dict
@router.get("/public-key")
def get_public_key():
return {"public_key": public_key.decode("utf-8")}
@router.post("/signup")
def signup(req: AuthRequest, db: Session = Depends(get_db)):
db_user = db.query(User).filter(User.username == req.username).first()
if db_user:
raise HTTPException(status_code=400, detail="Username already registered")
# Decrypt password
decrypted_password = decrypt_password(req.password)
hashed_password = get_password_hash(decrypted_password)
new_user = User(username=req.username, hashed_password=hashed_password)
db.add(new_user)
db.commit()
return {"message": "User created successfully"}
@router.post("/login")
def login(req: AuthRequest, db: Session = Depends(get_db)):
db_user = db.query(User).filter(User.username == req.username).first()
if not db_user:
raise HTTPException(status_code=400, detail="Invalid credentials")
decrypted_password = decrypt_password(req.password)
if not verify_password(decrypted_password, db_user.hashed_password):
raise HTTPException(status_code=400, detail="Invalid credentials")
access_token = create_access_token(data={"sub": db_user.username})
return {"access_token": access_token, "token_type": "bearer"}
def get_current_user(token: str = Depends(oauth2_scheme), db: Session = Depends(get_db)):
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
except JWTError:
raise credentials_exception
user = db.query(User).filter(User.username == username).first()
if user is None:
raise credentials_exception
return user
@router.get("/me")
def get_me(current_user: User = Depends(get_current_user)):
return {
"username": current_user.username,
"display_name": current_user.display_name,
"email_id": current_user.email_id,
"mobile_no": current_user.mobile_no,
"gender": current_user.gender,
"sector_map": current_user.sector_map
}
@router.put("/update-sectors")
def update_sectors(req: SectorMapUpdateRequest, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
current_user.sector_map = req.sector_map
db.commit()
return {"message": "Sectors updated successfully"}
@router.put("/update-profile")
def update_profile(req: ProfileUpdateRequest, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
if req.display_name is not None:
current_user.display_name = req.display_name
if req.email_id is not None:
current_user.email_id = req.email_id
if req.mobile_no is not None:
current_user.mobile_no = req.mobile_no
if req.gender is not None:
current_user.gender = req.gender
if req.password:
decrypted_password = decrypt_password(req.password)
hashed_password = get_password_hash(decrypted_password)
current_user.hashed_password = hashed_password
db.commit()
return {"message": "Profile updated successfully"}