|
|
@@ -3,10 +3,12 @@ from collections.abc import Generator
|
|
|
from uuid import UUID
|
|
|
|
|
|
import jwt
|
|
|
-from fastapi import Depends, HTTPException, Path, status
|
|
|
+from fastapi import Depends, HTTPException, Path, Query, 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
|
|
|
@@ -59,6 +61,70 @@ 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")
|