Added ref login
This commit is contained in:
+19
-15
@@ -6,7 +6,7 @@ from fastapi.security import OAuth2PasswordBearer
|
||||
import jwt
|
||||
from jwt.exceptions import PyJWTError
|
||||
|
||||
from .config import SECRET_KEY, ALGORITHM, ADMIN_USER
|
||||
from .config import SECRET_KEY, ALGORITHM
|
||||
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="auth/token")
|
||||
oauth2_scheme_optional = OAuth2PasswordBearer(tokenUrl="auth/token", auto_error=False)
|
||||
@@ -20,43 +20,47 @@ def create_access_token(data: dict, expires_delta: Optional[timedelta] = None):
|
||||
expire = datetime.now(timezone.utc) + timedelta(minutes=15)
|
||||
|
||||
to_encode.update({"exp": expire})
|
||||
|
||||
encoded_jwt = jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
|
||||
return encoded_jwt
|
||||
|
||||
|
||||
async def get_current_user(token: str = Depends(oauth2_scheme)):
|
||||
async def get_authenticated_user(token: str = Depends(oauth2_scheme)):
|
||||
"""Allows both Admins and Refs"""
|
||||
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 = payload.get("sub")
|
||||
|
||||
if username is None or username != ADMIN_USER:
|
||||
role = payload.get("role")
|
||||
if role not in ["admin", "ref"]:
|
||||
raise credentials_exception
|
||||
|
||||
return payload
|
||||
except PyJWTError:
|
||||
raise credentials_exception
|
||||
|
||||
return username
|
||||
|
||||
async def get_admin_user(token: str = Depends(oauth2_scheme)):
|
||||
"""Strictly allows ONLY Admins"""
|
||||
user = await get_authenticated_user(token)
|
||||
if user.get("role") != "admin":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail="Admin privileges required"
|
||||
)
|
||||
return user
|
||||
|
||||
|
||||
async def get_optional_user(
|
||||
token: Optional[str] = Depends(oauth2_scheme_optional),
|
||||
) -> Optional[str]:
|
||||
) -> Optional[dict]:
|
||||
"""Returns user payload if valid token exists, else None"""
|
||||
if not token:
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
|
||||
username = payload.get("sub")
|
||||
if username == ADMIN_USER:
|
||||
return username
|
||||
if payload.get("role") in ["admin", "ref"]:
|
||||
return payload
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
@@ -13,8 +13,13 @@ ACCESS_TOKEN_EXPIRE_MINUTES = 60 * 24 # 24 hours
|
||||
ADMIN_USER = os.getenv("ADMIN_USER", "admin")
|
||||
ADMIN_PASSWORD = os.getenv("ADMIN_PASSWORD", "admin")
|
||||
|
||||
# Ref Credentials
|
||||
REF_USER = os.getenv("REF_USER", "ref")
|
||||
REF_PASSWORD = os.getenv("REF_PASSWORD", "ref")
|
||||
|
||||
password_hash = PasswordHash.recommended()
|
||||
ADMIN_HASH = password_hash.hash(ADMIN_PASSWORD)
|
||||
REF_HASH = password_hash.hash(REF_PASSWORD)
|
||||
|
||||
|
||||
def verify_password(plain_password, hashed_password):
|
||||
|
||||
@@ -1,63 +0,0 @@
|
||||
# backend/app/core/utils.py
|
||||
from .brackets import Match
|
||||
|
||||
|
||||
class MermaidLive:
|
||||
|
||||
@staticmethod
|
||||
def _get_visual_target(match: Match, is_win=True) -> Match | None:
|
||||
current = match.next_win if is_win else match.next_loss
|
||||
while current and current.is_bye:
|
||||
current = current.next_win
|
||||
return current
|
||||
|
||||
@classmethod
|
||||
def export(cls, matches: list[Match]) -> None:
|
||||
print("\n--- COPY TO MERMAID.LIVE ---")
|
||||
print("graph LR")
|
||||
print(" classDef wb stroke:#01579b,stroke-width:2px;")
|
||||
print(" classDef lb stroke:#b71c1c,stroke-width:2px,stroke-dasharray: 5 5;")
|
||||
print(" classDef final stroke:#e65100,stroke-width:4px;")
|
||||
|
||||
wb_nodes = [
|
||||
m
|
||||
for m in matches
|
||||
if "WB" in m.name or "Semifinal" in m.name or "Winners Final" in m.name
|
||||
]
|
||||
|
||||
lb_nodes = [m for m in matches if "LB" in m.name or "Losers Final" in m.name]
|
||||
final_nodes = [
|
||||
m for m in matches if "Grand Final" in m.name or "3rd Place" in m.name
|
||||
]
|
||||
|
||||
print(" subgraph Winners Bracket")
|
||||
for m in wb_nodes:
|
||||
cls._print_node(m, "wb")
|
||||
print(" end")
|
||||
|
||||
if lb_nodes:
|
||||
print(" subgraph Losers Bracket")
|
||||
for m in lb_nodes:
|
||||
cls._print_node(m, "lb")
|
||||
print(" end")
|
||||
|
||||
print(" subgraph Championship / 3rd Place")
|
||||
for m in final_nodes:
|
||||
cls._print_node(m, "final")
|
||||
print(" end")
|
||||
|
||||
for m in matches:
|
||||
target_win = cls._get_visual_target(m, is_win=True)
|
||||
if target_win:
|
||||
print(f" M{m.id} --> M{target_win.id}")
|
||||
|
||||
target_loss = cls._get_visual_target(m, is_win=False)
|
||||
if target_loss:
|
||||
print(f" M{m.id} -.-> M{target_loss.id}")
|
||||
|
||||
@staticmethod
|
||||
def _print_node(m: Match, style: str) -> None:
|
||||
label = m.name
|
||||
if m.teams[0] and m.teams[1]:
|
||||
label += f" ({m.teams[0]} vs {m.teams[1]})"
|
||||
print(f' M{m.id}["{label}"]:::{style}')
|
||||
+19
-12
@@ -6,11 +6,13 @@ from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi.security import OAuth2PasswordRequestForm
|
||||
|
||||
from ..schemas import Token
|
||||
from ..core.auth import create_access_token, get_current_user
|
||||
from ..core.auth import create_access_token, get_authenticated_user
|
||||
from ..core.config import (
|
||||
ACCESS_TOKEN_EXPIRE_MINUTES,
|
||||
ADMIN_HASH,
|
||||
ADMIN_USER,
|
||||
REF_HASH,
|
||||
REF_USER,
|
||||
verify_password,
|
||||
)
|
||||
|
||||
@@ -21,14 +23,18 @@ router = APIRouter(prefix="/auth", tags=["Auth"])
|
||||
async def login_for_access_token(
|
||||
form_data: Annotated[OAuth2PasswordRequestForm, Depends()],
|
||||
):
|
||||
if form_data.username != ADMIN_USER:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Incorrect username or password",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
role = None
|
||||
|
||||
if not verify_password(form_data.password, ADMIN_HASH):
|
||||
if form_data.username == ADMIN_USER and verify_password(
|
||||
form_data.password, ADMIN_HASH
|
||||
):
|
||||
role = "admin"
|
||||
elif form_data.username == REF_USER and verify_password(
|
||||
form_data.password, REF_HASH
|
||||
):
|
||||
role = "ref"
|
||||
|
||||
if not role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Incorrect username or password",
|
||||
@@ -37,12 +43,13 @@ async def login_for_access_token(
|
||||
|
||||
access_token_expires = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
access_token = create_access_token(
|
||||
data={"sub": form_data.username}, expires_delta=access_token_expires
|
||||
data={"sub": form_data.username, "role": role},
|
||||
expires_delta=access_token_expires,
|
||||
)
|
||||
|
||||
return {"access_token": access_token, "token_type": "bearer"}
|
||||
return {"access_token": access_token, "token_type": "bearer", "role": role}
|
||||
|
||||
|
||||
@router.get("/check")
|
||||
async def check_auth(user: str = Depends(get_current_user)):
|
||||
return {"is_admin": True, "user": user}
|
||||
async def check_auth(user: dict = Depends(get_authenticated_user)):
|
||||
return {"role": user.get("role"), "user": user.get("sub")}
|
||||
|
||||
@@ -4,7 +4,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from ... import crud, schemas
|
||||
from ...constants import SUCCESS
|
||||
from ...core.auth import get_current_user
|
||||
from ...core.auth import get_admin_user
|
||||
from ...core.websocket_manager import send_ws_update
|
||||
from ...database import get_db
|
||||
from . import router
|
||||
@@ -20,7 +20,7 @@ async def create_court(
|
||||
id: str,
|
||||
court: schemas.CourtCreate,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
new_court = crud.create_court(db, id, court)
|
||||
if not new_court:
|
||||
@@ -35,7 +35,7 @@ async def update_courts(
|
||||
id: str,
|
||||
courts: list[str],
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
t = crud.update_tournament_courts(db, id, courts)
|
||||
if not t:
|
||||
@@ -50,7 +50,7 @@ async def delete_court(
|
||||
id: str,
|
||||
court_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
success = crud.delete_court(db, id, court_id)
|
||||
if not success:
|
||||
|
||||
@@ -6,13 +6,13 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from ... import crud, schemas
|
||||
from ...database import get_db
|
||||
from ...core.auth import get_current_user
|
||||
from ...core.auth import get_admin_user
|
||||
from . import router
|
||||
|
||||
|
||||
@router.get("/{id}/settings", response_model=schemas.TournamentSettingsResponse)
|
||||
def get_tournament(
|
||||
id: str, db: Session = Depends(get_db), user: str = Depends(get_current_user)
|
||||
id: str, db: Session = Depends(get_db), user: dict = Depends(get_admin_user)
|
||||
):
|
||||
t = crud.get_tournament(db, id)
|
||||
if not t:
|
||||
|
||||
@@ -4,7 +4,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from ... import crud, schemas
|
||||
from ...constants import SUCCESS
|
||||
from ...core.auth import get_current_user
|
||||
from ...core.auth import get_admin_user
|
||||
from ...core.websocket_manager import send_ws_update
|
||||
from ...database import get_db
|
||||
from . import router
|
||||
@@ -20,7 +20,7 @@ async def create_team(
|
||||
id: str,
|
||||
team: schemas.TeamCreate,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
new_team = crud.create_team(db, id, team)
|
||||
if not new_team:
|
||||
@@ -35,7 +35,7 @@ async def update_teams(
|
||||
id: str,
|
||||
teams: list[str],
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
t = crud.update_tournament_teams(db, id, teams)
|
||||
if not t:
|
||||
@@ -50,7 +50,7 @@ async def delete_team(
|
||||
id: str,
|
||||
team_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
success = crud.delete_team(db, id, team_id)
|
||||
if not success:
|
||||
|
||||
@@ -6,7 +6,7 @@ from ... import crud, schemas
|
||||
from ...constants import SUCCESS
|
||||
from ...core.websocket_manager import send_ws_update
|
||||
from ...database import get_db
|
||||
from ...core.auth import get_current_user
|
||||
from ...core.auth import get_admin_user
|
||||
from . import router
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from . import router
|
||||
async def create_tournament(
|
||||
data: schemas.TournamentCreate,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
new_t = crud.create_tournament(db, data)
|
||||
await send_ws_update(new_t.id)
|
||||
@@ -39,7 +39,7 @@ async def update_settings(
|
||||
id: str,
|
||||
data: schemas.TournamentUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
user: str = Depends(get_current_user),
|
||||
user: dict = Depends(get_admin_user),
|
||||
):
|
||||
t = crud.update_tournament_details(db, id, data)
|
||||
if not t:
|
||||
@@ -51,7 +51,7 @@ async def update_settings(
|
||||
|
||||
@router.delete("/{id}")
|
||||
async def delete_tournament(
|
||||
id: str, db: Session = Depends(get_db), user: str = Depends(get_current_user)
|
||||
id: str, db: Session = Depends(get_db), user: dict = Depends(get_admin_user)
|
||||
):
|
||||
success = crud.delete_tournament(db, id)
|
||||
if not success:
|
||||
|
||||
@@ -110,3 +110,4 @@ class TournamentSettingsResponse(TournamentOut):
|
||||
class Token(BaseModel):
|
||||
access_token: str
|
||||
token_type: str
|
||||
role: str | None
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
from typing import AsyncGenerator, Generator
|
||||
|
||||
import pytest
|
||||
from app.core.auth import get_current_user, get_optional_user
|
||||
from app.core.auth import get_admin_user, get_optional_user
|
||||
from app.database import Base, get_db
|
||||
|
||||
# Import your app and models
|
||||
@@ -52,7 +52,7 @@ async def client(db: Session) -> AsyncGenerator[AsyncClient, None]:
|
||||
pass
|
||||
|
||||
# Strict Auth: Always requires a token (simulated by header presence)
|
||||
def override_get_current_user(request: Request):
|
||||
def override_get_admin_user(request: Request):
|
||||
if "Authorization" not in request.headers:
|
||||
# Let FastAPI raise the 401 naturally if header is missing
|
||||
raise pytest.skip("Auth header missing in strict auth test")
|
||||
@@ -65,7 +65,7 @@ async def client(db: Session) -> AsyncGenerator[AsyncClient, None]:
|
||||
return None
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
app.dependency_overrides[get_current_user] = override_get_current_user
|
||||
app.dependency_overrides[get_admin_user] = override_get_admin_user
|
||||
app.dependency_overrides[get_optional_user] = override_get_optional_user
|
||||
|
||||
async with AsyncClient(
|
||||
|
||||
Reference in New Issue
Block a user