|
|
@@ -3,12 +3,10 @@ from collections.abc import Generator
|
|
|
from uuid import UUID
|
|
|
|
|
|
import jwt
|
|
|
-from fastapi import Depends, HTTPException, Path, Query, status
|
|
|
+from fastapi import Depends, HTTPException, Path, status
|
|
|
from fastapi.security import OAuth2PasswordBearer
|
|
|
-from fastapi.security.utils import get_authorization_scheme_param
|
|
|
from sqlalchemy import exists, select
|
|
|
from sqlalchemy.orm import Session
|
|
|
-from starlette.requests import Request
|
|
|
|
|
|
from app.api.utils import get_project_organization_id
|
|
|
from app.core import config, security
|
|
|
@@ -61,70 +59,6 @@ async def get_current_user(
|
|
|
return user
|
|
|
|
|
|
|
|
|
-async def get_token_flexible(
|
|
|
- request: Request,
|
|
|
- token: str | None = Query(default=None),
|
|
|
-) -> str:
|
|
|
- """Same as reusable_oauth2, but also accepts ?token=... in the query
|
|
|
- string — needed because native EventSource cannot set custom headers."""
|
|
|
- auth_header = request.headers.get("Authorization")
|
|
|
- if auth_header:
|
|
|
- scheme, param = get_authorization_scheme_param(auth_header)
|
|
|
- if scheme.lower() == "bearer" and param:
|
|
|
- return param
|
|
|
- if token:
|
|
|
- return token
|
|
|
- raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Not authenticated")
|
|
|
-
|
|
|
-
|
|
|
-async def get_current_user_flexible(
|
|
|
- session: Session = Depends(get_session),
|
|
|
- token: str = Depends(get_token_flexible),
|
|
|
-) -> User:
|
|
|
- # identical body to get_current_user, just sourcing `token` differently
|
|
|
- try:
|
|
|
- payload = jwt.decode(token, config.settings.SECRET_KEY, algorithms=[security.JWT_ALGORITHM])
|
|
|
- except jwt.DecodeError:
|
|
|
- raise HTTPException(status.HTTP_403_FORBIDDEN, "Could not validate credentials.")
|
|
|
-
|
|
|
- token_data = security.JWTTokenPayload(**payload)
|
|
|
- if token_data.refresh:
|
|
|
- raise HTTPException(
|
|
|
- status.HTTP_403_FORBIDDEN, "Could not validate credentials, cannot use refresh token"
|
|
|
- )
|
|
|
-
|
|
|
- now = int(time.time())
|
|
|
- if now < token_data.issued_at or now > token_data.expires_at:
|
|
|
- raise HTTPException(
|
|
|
- status.HTTP_403_FORBIDDEN,
|
|
|
- "Could not validate credentials, token expired or not yet valid",
|
|
|
- )
|
|
|
-
|
|
|
- result = session.execute(select(User).where(User.id == token_data.sub))
|
|
|
- user = result.scalars().first()
|
|
|
- if not user:
|
|
|
- raise HTTPException(status_code=404, detail="User not found.")
|
|
|
- return user
|
|
|
-
|
|
|
-
|
|
|
-def require_org_role_sse(*allowed_roles: OrgRole):
|
|
|
- def dependency(
|
|
|
- project_id: UUID = Path(...),
|
|
|
- session: Session = Depends(get_session),
|
|
|
- current_user: User = Depends(get_current_user_flexible),
|
|
|
- ) -> User:
|
|
|
- if current_user.global_role == GlobalRole.SUPER_ADMIN:
|
|
|
- return current_user
|
|
|
- if not _has_org_role(session, current_user.id, project_id, *allowed_roles):
|
|
|
- get_project_organization_id(session, project_id)
|
|
|
- raise HTTPException(
|
|
|
- status.HTTP_403_FORBIDDEN, "Insufficient permissions for this organization"
|
|
|
- )
|
|
|
- return current_user
|
|
|
-
|
|
|
- return dependency
|
|
|
-
|
|
|
-
|
|
|
def require_super_admin(current_user: User = Depends(get_current_user)) -> User:
|
|
|
if current_user.global_role != GlobalRole.SUPER_ADMIN:
|
|
|
raise HTTPException(status.HTTP_403_FORBIDDEN, "Requires super_admin")
|