diverse stuff
This commit is contained in:
parent
93de81cf86
commit
4b07853762
6 changed files with 61 additions and 20 deletions
Binary file not shown.
|
|
@ -1,3 +1,10 @@
|
||||||
|
from fastapi.security import HTTPBasicCredentials
|
||||||
|
from fastapi import Security
|
||||||
|
from fastapi.security import HTTPBearer
|
||||||
|
from fastapi import Header
|
||||||
|
from jwt.exceptions import PyJWTError
|
||||||
|
from datetime import timezone
|
||||||
|
import jwt
|
||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from backend.models.user import User
|
from backend.models.user import User
|
||||||
from backend.internal.database import get_session
|
from backend.internal.database import get_session
|
||||||
|
|
@ -13,23 +20,29 @@ SECRET_KEY = "secret-key"
|
||||||
ALGORITHM = "HS256"
|
ALGORITHM = "HS256"
|
||||||
ACCESS_TOKEN_EXPIRE_DAYS = 360
|
ACCESS_TOKEN_EXPIRE_DAYS = 360
|
||||||
|
|
||||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
|
|
||||||
|
|
||||||
def create_access_token(data: dict, expires_delta: timedelta | None = None):
|
def create_access_token(data: dict, expires_delta: timedelta | None = None):
|
||||||
to_encode = data.copy()
|
to_encode = data.copy()
|
||||||
expire = datetime.utcnow() + (expires_delta or timedelta(days=ACCESS_TOKEN_EXPIRE_DAYS))
|
expire = datetime.now(timezone.utc) + (expires_delta or timedelta(days=ACCESS_TOKEN_EXPIRE_DAYS))
|
||||||
to_encode.update({"exp": expire})
|
to_encode.update({"exp": expire})
|
||||||
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
|
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
|
||||||
|
|
||||||
async def get_current_user(token: str = Depends(oauth2_scheme), session: Session = Depends(get_session)):
|
security = HTTPBearer()
|
||||||
|
|
||||||
|
def get_token(credentials = Security(security)) -> str:
|
||||||
|
parts = credentials.credentials.split()
|
||||||
|
if len(parts) == 2 and parts[0].lower() == "bearer":
|
||||||
|
return parts[1]
|
||||||
|
raise HTTPException(status_code=403, detail="Invalid authorization format")
|
||||||
|
|
||||||
|
async def get_current_user(token: str = Depends(get_token), session: Session = Depends(get_session)):
|
||||||
try:
|
try:
|
||||||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||||||
user_code: str = payload.get("sub")
|
user_code: str | None = payload.get("sub")
|
||||||
if user_code is None:
|
if user_code is None:
|
||||||
raise HTTPException(status_code=401, detail="Invalid token")
|
raise HTTPException(status_code=401, detail="Invalid token")
|
||||||
except JWTError:
|
except PyJWTError:
|
||||||
raise HTTPException(status_code=401, detail="Invalid token")
|
raise HTTPException(status_code=401, detail="Invalid token")
|
||||||
user = session.exec(select(User).where(User.user_code == user_code))
|
user = session.exec(select(User).where(User.user_code == user_code)).first()
|
||||||
if user is None:
|
if user is None:
|
||||||
raise HTTPException(status_code=401, detail="User not found")
|
raise HTTPException(status_code=401, detail="User not found")
|
||||||
return user
|
return user
|
||||||
Binary file not shown.
|
|
@ -56,3 +56,8 @@ class NewUser(UserBase):
|
||||||
name = name,
|
name = name,
|
||||||
subscription = subscription
|
subscription = subscription
|
||||||
)
|
)
|
||||||
|
|
||||||
|
class UserInfo(UserBase):
|
||||||
|
name: str
|
||||||
|
email: str
|
||||||
|
subscription: bool
|
||||||
Binary file not shown.
|
|
@ -1,3 +1,6 @@
|
||||||
|
from backend.models.user import UserInfo
|
||||||
|
from backend.internal.authentication import get_current_user
|
||||||
|
from backend.internal.authentication import create_access_token
|
||||||
from starlette.status import HTTP_400_BAD_REQUEST,HTTP_401_UNAUTHORIZED
|
from starlette.status import HTTP_400_BAD_REQUEST,HTTP_401_UNAUTHORIZED
|
||||||
from fastapi import Form
|
from fastapi import Form
|
||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
|
|
@ -11,38 +14,58 @@ from sqlmodel import Session
|
||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
|
||||||
router = APIRouter(
|
router = APIRouter(
|
||||||
prefix="/users",
|
prefix="/",
|
||||||
tags=["users"],
|
tags=[""],
|
||||||
dependencies=[],
|
dependencies=[],
|
||||||
responses={404: {"description": "Not found"}},
|
responses={404: {"description": "Not found"}},
|
||||||
)
|
)
|
||||||
|
|
||||||
@router.post("/", status_code=201)
|
@router.post("/users", status_code=201)
|
||||||
def register_user(new_user: NewUser, session: Session = Depends(get_session)):
|
def register_user(new_user: NewUser, session: Session = Depends(get_session)) -> str:
|
||||||
user = User.model_validate(new_user)
|
user = User.model_validate(new_user)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
session.add(user)
|
session.add(user)
|
||||||
session.commit()
|
session.commit()
|
||||||
|
session.refresh(user)
|
||||||
# TODO: send mail
|
# TODO: send mail
|
||||||
except IntegrityError as e:
|
except IntegrityError as e:
|
||||||
session.rollback()
|
session.rollback()
|
||||||
raise HTTPException(HTTP_400_BAD_REQUEST)
|
raise HTTPException(HTTP_400_BAD_REQUEST)
|
||||||
|
return user.user_code
|
||||||
|
|
||||||
@router.post("/login") # TODO: rate limit
|
@router.post("/login", status_code=201) # TODO: rate limit
|
||||||
def login(user_code: str = Form(), session: Session = Depends(get_session)) -> str:
|
def login(user_code: str = Form(), session: Session = Depends(get_session)) -> str:
|
||||||
user = session.exec(select(User).where(User.user_code == user_code)).one
|
user = session.exec(select(User).where(User.user_code == user_code)).first()
|
||||||
|
|
||||||
if not user:
|
if not user:
|
||||||
raise HTTPException(HTTP_401_UNAUTHORIZED)
|
raise HTTPException(HTTP_401_UNAUTHORIZED)
|
||||||
return "token"
|
return create_access_token({"sub": user.user_code})
|
||||||
|
|
||||||
|
@router.post("users/forgot") # TODO: rate limit
|
||||||
|
def forgot_user_code(email: str = Form(), session: Session = Depends(get_session)):
|
||||||
|
user = session.exec(select(User).where(User.email == email)).first()
|
||||||
|
|
||||||
|
if user is not None:
|
||||||
|
pass # TODO: send Mail
|
||||||
|
|
||||||
|
@router.get("/me", status_code=201)
|
||||||
|
def get_current_user_info(current_user: User = Depends(get_current_user)) -> UserInfo:
|
||||||
|
return UserInfo.model_validate(current_user)
|
||||||
|
|
||||||
|
|
||||||
# Login Rate Limit! -> token for everything else
|
@router.delete("/me", status_code=201)
|
||||||
# Forgot Code <- email (rate limit)
|
def delete(current_user: User = Depends(get_current_user), session: Session = Depends(get_session)) :
|
||||||
# Delete User <- token
|
session.delete(current_user)
|
||||||
|
session.commit()
|
||||||
|
|
||||||
|
@router.put("me/subscribe", status_code=201)
|
||||||
|
def change_subscription(subscription: bool = Form(), current_user: User = Depends(get_current_user), session: Session = Depends(get_session)):
|
||||||
|
current_user.subscription = subscription
|
||||||
|
session.add(current_user)
|
||||||
|
session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
# TODO:
|
||||||
# change mail <- token, new mail
|
# change mail <- token, new mail
|
||||||
# subscribe
|
|
||||||
# unsubscribe
|
|
||||||
|
|
||||||
# delete all deactivated users after 48h
|
# delete all deactivated users after 48h
|
||||||
Loading…
Add table
Reference in a new issue