From a5aff8f99d17b212cbe5d9c0d6371180903aa54f Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 07:55:12 +0000 Subject: [PATCH 01/32] Initial plan From 5ade87437a85e1e25ce2cfac32f69a5ff1e7c17a Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 08:16:42 +0000 Subject: [PATCH 02/32] Strengthen TRE API authentication: new auth layer, remove AccessService abstraction --- CHANGELOG.md | 1 + api_app/api/routes/resource_helpers.py | 11 +- api_app/api/routes/workspace_users.py | 12 +- api_app/api/routes/workspaces.py | 6 +- api_app/auth/__init__.py | 0 api_app/auth/dependencies.py | 82 ++++++ api_app/auth/exceptions.py | 22 ++ api_app/auth/models.py | 43 +++ api_app/auth/rbac.py | 89 ++++++ api_app/auth/registry.py | 49 ++++ api_app/auth/token_validator.py | 71 +++++ api_app/db/repositories/airlock_requests.py | 4 +- api_app/services/aad_authentication.py | 276 +++++++++--------- api_app/services/access_service.py | 37 --- api_app/services/airlock.py | 6 +- api_app/services/authentication.py | 20 +- api_app/tests_ma/auth/__init__.py | 0 api_app/tests_ma/auth/test_rbac.py | 151 ++++++++++ api_app/tests_ma/auth/test_token_validator.py | 146 +++++++++ api_app/tests_ma/test_api/conftest.py | 14 +- .../test_airlock_request_repository.py | 12 +- .../test_services/test_aad_access_service.py | 3 +- 22 files changed, 832 insertions(+), 223 deletions(-) create mode 100644 api_app/auth/__init__.py create mode 100644 api_app/auth/dependencies.py create mode 100644 api_app/auth/exceptions.py create mode 100644 api_app/auth/models.py create mode 100644 api_app/auth/rbac.py create mode 100644 api_app/auth/registry.py create mode 100644 api_app/auth/token_validator.py delete mode 100644 api_app/services/access_service.py create mode 100644 api_app/tests_ma/auth/__init__.py create mode 100644 api_app/tests_ma/auth/test_rbac.py create mode 100644 api_app/tests_ma/auth/test_token_validator.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ef9339b2b8..8ffc8818fd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ ENHANCEMENTS: * Add Windows Server 2025 image support to Guacamole. ([#4890](https://github.com/microsoft/AzureTRE/issues/4890)) * Add support for setting resource processor VMSS SKU via environment variables ([#4936](https://github.com/microsoft/AzureTRE/issues/4936)) * Exclude recovery service vaults from e2e tests ([#4920](https://github.com/microsoft/AzureTRE/issues/4920)) +* Strengthen TRE API authentication: introduce layered `auth/` package with typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, and composable RBAC factories; remove the `AccessService` abstraction that is no longer needed now that Entra ID is the only auth provider. ## (0.28.0) (March 2, 2026) **BREAKING CHANGES** diff --git a/api_app/api/routes/resource_helpers.py b/api_app/api/routes/resource_helpers.py index 28e1ca000c..a32e84bc3f 100644 --- a/api_app/api/routes/resource_helpers.py +++ b/api_app/api/routes/resource_helpers.py @@ -24,7 +24,7 @@ send_resource_request_message, RequestAction, ) -from services.authentication import get_access_service +from services.authentication import get_aad_service from services.logging import logger @@ -157,13 +157,8 @@ def construct_location_header(operation: Operation) -> str: def get_identity_role_assignments(user): - access_service = get_access_service() - return access_service.get_identity_role_assignments(user.id) - - -def get_app_user_roles_assignments_emails(app_obj_id): - access_service = get_access_service() - return access_service.get_app_user_role_assignments_emails(app_obj_id) + aad_service = get_aad_service() + return aad_service.get_identity_role_assignments(user.id) async def send_uninstall_message( diff --git a/api_app/api/routes/workspace_users.py b/api_app/api/routes/workspace_users.py index b90a6b07ea..90a92fa62a 100644 --- a/api_app/api/routes/workspace_users.py +++ b/api_app/api/routes/workspace_users.py @@ -2,7 +2,7 @@ from api.dependencies.workspaces import get_workspace_by_id_from_path from models.schemas.workspace_users import UserRoleAssignmentRequest from resources import strings -from services.authentication import get_access_service +from services.authentication import get_aad_service from models.schemas.users import UsersInResponse, AssignableUsersInResponse, WorkspaceUserOperationResponse from models.schemas.roles import RolesInResponse from services.authentication import get_current_admin_user, get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin @@ -12,25 +12,25 @@ @workspaces_users_shared_router.get("/workspaces/{workspace_id}/users", response_model=UsersInResponse, name=strings.API_GET_WORKSPACE_USERS) -async def get_workspace_users(workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_access_service)) -> UsersInResponse: +async def get_workspace_users(workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_aad_service)) -> UsersInResponse: users = access_service.get_workspace_users(workspace) return UsersInResponse(users=users) @workspaces_users_admin_router.get("/workspaces/{workspace_id}/assignable-users", response_model=AssignableUsersInResponse, name=strings.API_GET_ASSIGNABLE_USERS) -async def get_assignable_users(filter: str = "", maxResultCount: int = 5, access_service=Depends(get_access_service)) -> AssignableUsersInResponse: +async def get_assignable_users(filter: str = "", maxResultCount: int = 5, access_service=Depends(get_aad_service)) -> AssignableUsersInResponse: assignable_users = access_service.get_assignable_users(filter, maxResultCount) return AssignableUsersInResponse(assignable_users=assignable_users) @workspaces_users_admin_router.get("/workspaces/{workspace_id}/roles", response_model=RolesInResponse, name=strings.API_GET_WORKSPACE_ROLES) -async def get_workspace_roles(workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_access_service)) -> RolesInResponse: +async def get_workspace_roles(workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_aad_service)) -> RolesInResponse: roles = access_service.get_workspace_roles(workspace) return RolesInResponse(roles=roles) @workspaces_users_admin_router.post("/workspaces/{workspace_id}/users/assign", status_code=status.HTTP_202_ACCEPTED, name=strings.API_ASSIGN_WORKSPACE_USER) -async def assign_workspace_user(response: Response, userRoleAssignmentRequest: UserRoleAssignmentRequest, workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_access_service)) -> WorkspaceUserOperationResponse: +async def assign_workspace_user(response: Response, userRoleAssignmentRequest: UserRoleAssignmentRequest, workspace=Depends(get_workspace_by_id_from_path), access_service=Depends(get_aad_service)) -> WorkspaceUserOperationResponse: for user_id in userRoleAssignmentRequest.user_ids: access_service.assign_workspace_user( @@ -46,7 +46,7 @@ async def assign_workspace_user(response: Response, userRoleAssignmentRequest: U async def remove_workspace_user_assignment(user_id: str, role_id: str, workspace=Depends(get_workspace_by_id_from_path), - access_service=Depends(get_access_service)) -> WorkspaceUserOperationResponse: + access_service=Depends(get_aad_service)) -> WorkspaceUserOperationResponse: access_service.remove_workspace_role_user_assignment( user_id, diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 4f8ee42bb2..78f3965697 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -23,9 +23,9 @@ from models.schemas.resource import ResourceHistoryInList, ResourcePatch from models.schemas.resource_template import ResourceTemplateInformationInList from resources import strings -from services.access_service import AuthConfigValidationError +from services.aad_authentication import AuthConfigValidationError from services.authentication import get_current_admin_user, \ - get_access_service, get_current_workspace_owner_user, get_current_workspace_owner_or_researcher_user, get_current_tre_user_or_tre_admin, \ + get_aad_service, get_current_workspace_owner_user, get_current_workspace_owner_or_researcher_user, get_current_tre_user_or_tre_admin, \ get_current_workspace_owner_or_tre_admin, \ get_current_workspace_owner_or_researcher_user_or_airlock_manager, \ get_current_workspace_owner_or_airlock_manager, \ @@ -67,7 +67,7 @@ async def retrieve_users_active_workspaces(request: Request, user=Depends(get_cu except Exception: workspaces = await workspace_repo.get_active_workspaces() - access_service = get_access_service() + access_service = get_aad_service() user_role_assignments = get_identity_role_assignments(user) def _safe_get_workspace_role(user, workspace, user_role_assignments): diff --git a/api_app/auth/__init__.py b/api_app/auth/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/api_app/auth/dependencies.py b/api_app/auth/dependencies.py new file mode 100644 index 0000000000..c94c51def2 --- /dev/null +++ b/api_app/auth/dependencies.py @@ -0,0 +1,82 @@ +from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer + +from auth.exceptions import AuthError, TokenExpired, TokenInvalid, TokenSignatureInvalid +from auth.models import AuthenticatedUser +from auth.registry import get_core_validator, get_workspace_validator +from models.domain.workspace import Workspace +from resources import strings +from services.logging import logger + +_bearer = HTTPBearer(auto_error=True) + + +def _to_http_exception(exc: AuthError) -> HTTPException: + if isinstance(exc, TokenExpired): + return HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.EXPIRED_SIGNATURE, + ) + if isinstance(exc, TokenSignatureInvalid): + return HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_SIGNATURE, + ) + return HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_TOKEN, + ) + + +async def get_authenticated_user( + credentials: HTTPAuthorizationCredentials = Depends(_bearer), +) -> AuthenticatedUser: + """Validate the bearer token against the core TRE app registration. + + Returns an immutable :class:`AuthenticatedUser` or raises HTTP 401. + """ + try: + return get_core_validator().validate(credentials.credentials) + except AuthError as exc: + logger.debug("Core token validation failed: %s", exc) + raise _to_http_exception(exc) + + +async def get_workspace_authenticated_user( + credentials: HTTPAuthorizationCredentials = Depends(_bearer), + workspace: Workspace = Depends(lambda: None), +) -> AuthenticatedUser: + """Validate the bearer token for a workspace-scoped request. + + Tries the workspace app registration first. If that audience validation + fails (but not on expiry or signature errors), falls back to the core app + registration so that TREAdmin users can always reach workspace endpoints. + + The *workspace* parameter should be provided via + ``Depends(get_workspace_by_id_from_path)`` when wiring routes. + """ + token = credentials.credentials + + if workspace is not None: + client_id = workspace.properties.get("client_id", "") + if client_id: + try: + return get_workspace_validator(client_id).validate(token) + except (TokenExpired, TokenSignatureInvalid) as exc: + # Hard failures — the token IS for this workspace but is bad. + logger.debug("Workspace token hard failure: %s", exc) + raise _to_http_exception(exc) + except TokenInvalid as exc: + # Wrong audience — fall through to core validator. + logger.debug( + "Workspace token invalid (likely wrong audience), " + "trying core validator: %s", + exc, + ) + + # Fall back to core app registration (allows TREAdmin access). + try: + return get_core_validator().validate(token) + except AuthError as exc: + logger.debug("Core token validation failed: %s", exc) + raise _to_http_exception(exc) diff --git a/api_app/auth/exceptions.py b/api_app/auth/exceptions.py new file mode 100644 index 0000000000..6c03d988cf --- /dev/null +++ b/api_app/auth/exceptions.py @@ -0,0 +1,22 @@ +class AuthError(Exception): + """Base class for all authentication and authorisation errors.""" + + +class TokenExpired(AuthError): + """The JWT has passed its expiry time.""" + + +class TokenSignatureInvalid(AuthError): + """The JWT signature does not match the signing key.""" + + +class TokenInvalid(AuthError): + """The JWT is structurally invalid or fails claims validation.""" + + +class InsufficientPermissions(AuthError): + """The authenticated user does not hold the required role.""" + + +class WorkspaceNotFound(AuthError): + """The requested workspace could not be found.""" diff --git a/api_app/auth/models.py b/api_app/auth/models.py new file mode 100644 index 0000000000..f9a7dc9b3a --- /dev/null +++ b/api_app/auth/models.py @@ -0,0 +1,43 @@ +from enum import StrEnum +from typing import List, Optional, Union + +from pydantic import BaseModel + + +class TRERole(StrEnum): + Admin = "TREAdmin" + User = "TREUser" + AirlockAutomation = "TREAirlockAutomation" + + +class WorkspaceAccessRole(StrEnum): + Owner = "WorkspaceOwner" + Researcher = "WorkspaceResearcher" + AirlockManager = "AirlockManager" + + +class AuthenticatedUser(BaseModel): + """Immutable, validated user derived from a JWT. + + Fields map directly to standard JWT claims; ``id`` holds the ``oid`` + claim (the stable object identifier in Entra ID). The model is frozen so + roles cannot be escalated after creation. + """ + + id: str + name: str + email: Optional[str] = None + roles: List[str] = [] + audience: str = "" + is_workspace_token: bool = False + + class Config: + frozen = True + + def has_any_role(self, *roles: Union[TRERole, WorkspaceAccessRole]) -> bool: + """Return *True* if the user holds at least one of *roles*.""" + role_values = {r.value for r in roles} + return bool(role_values & set(self.roles)) + + def is_tre_admin(self) -> bool: + return TRERole.Admin in self.roles diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py new file mode 100644 index 0000000000..773beb9e00 --- /dev/null +++ b/api_app/auth/rbac.py @@ -0,0 +1,89 @@ +from typing import Callable, Union + +from fastapi import Depends, HTTPException, status + +from auth.dependencies import get_authenticated_user, get_workspace_authenticated_user +from auth.models import AuthenticatedUser, TRERole, WorkspaceAccessRole +from resources import strings + + +def require_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Callable: + """Factory that returns a FastAPI dependency enforcing at least one of *roles*. + + The dependency validates the bearer token against the core app registration + and raises HTTP 403 if the user does not hold at least one required role. + """ + role_values = frozenset(r.value for r in roles) + role_names = [r.value for r in roles] + + async def _check( + user: AuthenticatedUser = Depends(get_authenticated_user), + ) -> AuthenticatedUser: + if not (set(user.roles) & role_values): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {role_names}", + headers={"WWW-Authenticate": "Bearer"}, + ) + return user + + return _check + + +def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Callable: + """Factory that returns a dependency enforcing workspace-scoped *roles*. + + TREAdmin users are always allowed regardless of workspace role. For other + users the token is validated against the workspace app registration and the + role list is checked. + + Workspace context is resolved by pairing with + ``Depends(get_workspace_by_id_from_path)`` on the route. + """ + role_values = frozenset(r.value for r in roles) + # TREAdmin can access any workspace endpoint + allowed_values = role_values | {TRERole.Admin.value} + role_names = [r.value for r in roles] + + async def _check( + user: AuthenticatedUser = Depends(get_workspace_authenticated_user), + ) -> AuthenticatedUser: + if not (set(user.roles) & allowed_values): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {role_names}", + headers={"WWW-Authenticate": "Bearer"}, + ) + return user + + return _check + + +# --------------------------------------------------------------------------- +# Pre-built role checks — replace the module-level AzureADAuthorization +# singletons that previously lived in services/authentication.py. +# --------------------------------------------------------------------------- + +require_tre_user = require_roles(TRERole.User) +require_tre_admin = require_roles(TRERole.Admin) +require_tre_user_or_admin = require_roles(TRERole.User, TRERole.Admin) +require_workspace_owner = require_workspace_roles(WorkspaceAccessRole.Owner) +require_workspace_researcher = require_workspace_roles(WorkspaceAccessRole.Researcher) +require_airlock_manager = require_workspace_roles(WorkspaceAccessRole.AirlockManager) +require_workspace_owner_or_researcher = require_workspace_roles( + WorkspaceAccessRole.Owner, WorkspaceAccessRole.Researcher +) +require_workspace_owner_or_airlock_manager = require_workspace_roles( + WorkspaceAccessRole.Owner, WorkspaceAccessRole.AirlockManager +) +require_workspace_owner_or_researcher_or_airlock_manager = require_workspace_roles( + WorkspaceAccessRole.Owner, + WorkspaceAccessRole.Researcher, + WorkspaceAccessRole.AirlockManager, +) +require_workspace_owner_or_researcher_or_airlock_manager_or_admin = require_workspace_roles( + WorkspaceAccessRole.Owner, + WorkspaceAccessRole.Researcher, + WorkspaceAccessRole.AirlockManager, +) +require_workspace_owner_or_admin = require_workspace_roles(WorkspaceAccessRole.Owner) diff --git a/api_app/auth/registry.py b/api_app/auth/registry.py new file mode 100644 index 0000000000..5764e37871 --- /dev/null +++ b/api_app/auth/registry.py @@ -0,0 +1,49 @@ +from functools import lru_cache + +from auth.token_validator import TokenValidator, TokenValidatorConfig +from core import config + + +def _jwks_uri() -> str: + # Direct JWKS endpoint — PyJWKClient fetches this and parses it as a key set. + return ( + f"{config.AAD_AUTHORITY_URL.rstrip('/')}" + f"/{config.AAD_TENANT_ID}/discovery/v2.0/keys" + ) + + +def _issuer() -> str: + return ( + f"{config.AAD_AUTHORITY_URL.rstrip('/')}" + f"/{config.AAD_TENANT_ID}/v2.0" + ) + + +@lru_cache(maxsize=1) +def get_core_validator() -> TokenValidator: + """Singleton :class:`TokenValidator` for the core TRE app registration.""" + return TokenValidator( + TokenValidatorConfig( + jwks_uri=_jwks_uri(), + audience=config.API_AUDIENCE, + issuer=_issuer(), + is_workspace_token=False, + ) + ) + + +@lru_cache(maxsize=256) +def get_workspace_validator(client_id: str) -> TokenValidator: + """Per-workspace :class:`TokenValidator`, cached by *client_id*. + + All workspace validators share the same JWKS URI so a single HTTP fetch + serves all audiences; only the audience validation differs. + """ + return TokenValidator( + TokenValidatorConfig( + jwks_uri=_jwks_uri(), + audience=client_id, + issuer=_issuer(), + is_workspace_token=True, + ) + ) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py new file mode 100644 index 0000000000..05b66ed2a1 --- /dev/null +++ b/api_app/auth/token_validator.py @@ -0,0 +1,71 @@ +from dataclasses import dataclass + +import jwt +from jwt import PyJWKClient + +from auth.exceptions import TokenExpired, TokenInvalid, TokenSignatureInvalid +from auth.models import AuthenticatedUser + + +@dataclass(frozen=True) +class TokenValidatorConfig: + jwks_uri: str + audience: str + issuer: str + is_workspace_token: bool = False + + +class TokenValidator: + """Single responsibility: validate a JWT and return a typed user. + + Uses :class:`jwt.PyJWKClient` which handles JWKS caching and key rotation + automatically — keys removed from the JWKS endpoint are evicted from the + cache, preventing unbounded growth. + """ + + def __init__(self, config: TokenValidatorConfig) -> None: + self._config = config + self._jwks_client = PyJWKClient(config.jwks_uri, cache_keys=True, lifespan=300) + + def validate(self, token: str) -> AuthenticatedUser: + """Validate *token* and return an :class:`AuthenticatedUser`. + + No silent exceptions — every failure mode raises a typed error. + + Raises: + TokenExpired: token has passed its expiry time. + TokenSignatureInvalid: signature cannot be verified. + TokenInvalid: any other validation failure. + """ + try: + signing_key = self._jwks_client.get_signing_key_from_jwt(token) + except Exception as exc: + raise TokenInvalid("Cannot obtain signing key") from exc + + try: + claims = jwt.decode( + token, + signing_key.key, + algorithms=["RS256"], + audience=self._config.audience, + issuer=self._config.issuer, + options={ + "verify_signature": True, + "verify_exp": True, + "verify_aud": True, + "verify_iss": True, + }, + ) + except jwt.ExpiredSignatureError as exc: + raise TokenExpired("Token expired") from exc + except jwt.InvalidTokenError as exc: + raise TokenInvalid(f"Token invalid: {exc}") from exc + + return AuthenticatedUser( + id=claims["oid"], + name=claims.get("name", ""), + email=claims.get("email") or claims.get("preferred_username"), + roles=claims.get("roles", []), + audience=self._config.audience, + is_workspace_token=self._config.is_workspace_token, + ) diff --git a/api_app/db/repositories/airlock_requests.py b/api_app/db/repositories/airlock_requests.py index 0990b90ef4..2cc37144bb 100644 --- a/api_app/db/repositories/airlock_requests.py +++ b/api_app/db/repositories/airlock_requests.py @@ -8,7 +8,7 @@ from fastapi import HTTPException, status from pydantic import parse_obj_as from db.repositories.workspaces import WorkspaceRepository -from services.authentication import get_access_service +from services.authentication import get_aad_service from models.domain.authentication import User from db.errors import EntityDoesNotExist from models.domain.airlock_request import AirlockFile, AirlockRequest, AirlockRequestStatus, \ @@ -162,7 +162,7 @@ async def get_airlock_request_by_id(self, airlock_request_id: UUID4) -> AirlockR async def get_airlock_requests_for_airlock_manager(self, user_id: str, type: Optional[AirlockRequestType] = None, status: Optional[AirlockRequestStatus] = None, order_by: Optional[str] = None, order_ascending=True) -> List[AirlockRequest]: workspace_repo = await WorkspaceRepository.create() - access_service = get_access_service() + access_service = get_aad_service() workspaces = await workspace_repo.get_active_workspaces() user_role_assignments = access_service.get_identity_role_assignments(user_id) diff --git a/api_app/services/aad_authentication.py b/api_app/services/aad_authentication.py index 4363ea3009..e6d147a571 100644 --- a/api_app/services/aad_authentication.py +++ b/api_app/services/aad_authentication.py @@ -1,203 +1,197 @@ -import base64 from collections import defaultdict from enum import Enum from typing import List, Optional -import jwt -import requests -from fastapi import Request, HTTPException, status +import requests +from fastapi import HTTPException, Request, status +from fastapi.security import OAuth2AuthorizationCodeBearer from msal import ConfidentialClientApplication +from semantic_version import Version -from services.access_service import AccessService, AuthConfigValidationError, UserRoleAssignmentError +from auth.exceptions import AuthError, TokenExpired, TokenInvalid, TokenSignatureInvalid +from auth.registry import get_core_validator, get_workspace_validator from core import config from db.errors import EntityDoesNotExist from models.domain.authentication import User, RoleAssignment -from models.domain.workspace_users import AssignedUser, AssignmentType, AssignableUser, Role from models.domain.workspace import Workspace, WorkspaceRole +from models.domain.workspace_users import AssignableUser, AssignedUser, AssignmentType, Role from resources import strings from db.repositories.workspaces import WorkspaceRepository from services.logging import logger -from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.hazmat.backends import default_backend -from cryptography.hazmat.primitives import serialization -from semantic_version import Version - MICROSOFT_GRAPH_URL = config.MICROSOFT_GRAPH_URL.strip("/") GRAPH_REQUEST_TIMEOUT = 10 USER_MANAGEMENT_MINIMUM_BASE_TEMPLATE_VERSION = "2.1.0" +class AuthConfigValidationError(Exception): + """Raised when the input auth information is invalid.""" + + +class UserRoleAssignmentError(Exception): + """Raised when a user role assignment fails.""" + + +def _authenticated_user_to_user(validated) -> User: + """Convert an :class:`~auth.models.AuthenticatedUser` to the legacy :class:`~models.domain.authentication.User`.""" + return User( + id=validated.id, + name=validated.name, + email=validated.email or "", + roles=list(validated.roles), + ) + + class PrincipalType(Enum): User = "User" Group = "Group" ServicePrincipal = "ServicePrincipal" -class AzureADAuthorization(AccessService): - _jwt_keys: dict = {} +class AzureADAuthorization(OAuth2AuthorizationCodeBearer): + """FastAPI security dependency that validates Entra ID JWTs. + + Uses :mod:`auth.token_validator` (backed by :class:`jwt.PyJWKClient`) for + JWT validation so key management is handled automatically. + """ - require_one_of_roles = None - aad_instance = config.AAD_AUTHORITY_URL + require_one_of_roles: Optional[list] = None + aad_instance: str = config.AAD_AUTHORITY_URL TRE_CORE_ROLES = ['TREAdmin', 'TREUser', 'TREAirlockAutomation'] - WORKSPACE_ROLES_DICT = {'WorkspaceOwner': 'app_role_id_workspace_owner', 'WorkspaceResearcher': 'app_role_id_workspace_researcher', 'AirlockManager': 'app_role_id_workspace_airlock_manager'} + WORKSPACE_ROLES_DICT = { + 'WorkspaceOwner': 'app_role_id_workspace_owner', + 'WorkspaceResearcher': 'app_role_id_workspace_researcher', + 'AirlockManager': 'app_role_id_workspace_airlock_manager', + } def __init__(self, auto_error: bool = True, require_one_of_roles: Optional[list] = None): - super(AzureADAuthorization, self).__init__( + super().__init__( authorizationUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/authorize", tokenUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/token", refreshUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/token", scheme_name="oauth2", - auto_error=auto_error + auto_error=auto_error, ) self.require_one_of_roles = require_one_of_roles async def __call__(self, request: Request) -> User: + token: str = await super().__call__(request) - token: str = await super(AzureADAuthorization, self).__call__(request) + decoded_user = None - decoded_token = None - - # Try workspace app registration if appropriate - if 'workspace_id' in request.path_params and any(role in self.require_one_of_roles for role in self.WORKSPACE_ROLES_DICT.keys()): - # as we have a workspace_id not given, try decoding token - logger.debug("Workspace ID was provided. Getting Workspace API app registration") + # Try workspace app registration first when a workspace_id is present + # and the route requires workspace-scoped roles. + if 'workspace_id' in request.path_params and any( + role in self.require_one_of_roles for role in self.WORKSPACE_ROLES_DICT + ): + logger.debug("Workspace ID present — attempting workspace app registration") try: - # get the app reg id - which might be blank if the workspace hasn't fully created yet. - # if it's blank, don't use workspace auth, use core auth - and a TRE Admin can still get it app_reg_id = await self._fetch_ws_app_reg_id_from_ws_id(request) - if app_reg_id != "": - decoded_token = self._decode_token(token, app_reg_id) - except HTTPException as h: - raise h - except Exception as e: - logger.debug(e) - logger.debug("Failed to decode using workspace_id, trying with TRE API app registration") - pass - - # Try TRE API app registration if appropriate - if decoded_token is None and any(role in self.require_one_of_roles for role in self.TRE_CORE_ROLES): + if app_reg_id: + try: + validated = get_workspace_validator(app_reg_id).validate(token) + decoded_user = self._get_user_from_token(validated) + except (TokenExpired, TokenSignatureInvalid) as exc: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.EXPIRED_SIGNATURE + if isinstance(exc, TokenExpired) + else strings.INVALID_SIGNATURE, + ) + except TokenInvalid: + logger.debug( + "Workspace token invalid, will try core app registration" + ) + except HTTPException: + raise + except Exception as exc: + logger.debug("Failed to resolve workspace app registration: %s", exc) + + # Try core app registration for TRE core roles. + if decoded_user is None and any( + role in self.require_one_of_roles for role in self.TRE_CORE_ROLES + ): try: - decoded_token = self._decode_token(token, config.API_AUDIENCE) - except jwt.exceptions.InvalidSignatureError: - logger.debug("Failed to decode using TRE API app registration (Invalid Signatrue)") - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.INVALID_SIGNATURE) - except jwt.exceptions.ExpiredSignatureError: - logger.debug("Failed to decode using TRE API app registration (Expired Signature)") - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.EXPIRED_SIGNATURE) - except jwt.exceptions.InvalidTokenError: - # any other token validation exception, we want to catch all of these... - logger.debug("Failed to decode using TRE API app registration (Invalid token)") - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.INVALID_TOKEN) - except Exception as e: - # Unexpected token decoding/validation exception. making sure we are not crashing (with 500) - logger.debug(e) - pass - - # Failed to decode token using either app registration - if decoded_token is None: - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.AUTH_UNABLE_TO_VALIDATE_TOKEN) - - try: - user = self._get_user_from_token(decoded_token) - except Exception as e: - logger.debug(e) - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.ACCESS_UNABLE_TO_GET_ROLE_ASSIGNMENTS_FOR_USER, headers={"WWW-Authenticate": "Bearer"}) + validated = get_core_validator().validate(token) + decoded_user = self._get_user_from_token(validated) + except TokenExpired: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.EXPIRED_SIGNATURE, + ) + except TokenSignatureInvalid: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_SIGNATURE, + ) + except TokenInvalid: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_TOKEN, + ) + except AuthError as exc: + logger.debug("Core token validation failed: %s", exc) + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.AUTH_UNABLE_TO_VALIDATE_TOKEN, + ) + + if decoded_user is None: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.AUTH_UNABLE_TO_VALIDATE_TOKEN, + ) - try: - if not any(role in self.require_one_of_roles for role in user.roles): - raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=f'{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {self.require_one_of_roles}', headers={"WWW-Authenticate": "Bearer"}) - except Exception as e: - logger.debug(e) - raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=f'{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {self.require_one_of_roles}', headers={"WWW-Authenticate": "Bearer"}) + if not any(role in self.require_one_of_roles for role in decoded_user.roles): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {self.require_one_of_roles}", + headers={"WWW-Authenticate": "Bearer"}, + ) - return user + return decoded_user @staticmethod async def _fetch_ws_app_reg_id_from_ws_id(request: Request) -> str: - workspace_id = None - if "workspace_id" not in request.path_params: - logger.error("Neither a workspace ID nor a default app registration id were provided") - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS) + workspace_id = request.path_params.get('workspace_id') + if not workspace_id: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, + ) try: - workspace_id = request.path_params['workspace_id'] ws_repo = await WorkspaceRepository.create() workspace = await ws_repo.get_workspace_by_id(workspace_id) - - ws_app_reg_id = "" - if "client_id" in workspace.properties: - ws_app_reg_id = workspace.properties['client_id'] - - return ws_app_reg_id + return workspace.properties.get('client_id', '') except EntityDoesNotExist: - logger.exception(strings.WORKSPACE_DOES_NOT_EXIST) - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=strings.WORKSPACE_DOES_NOT_EXIST) - except Exception: - logger.exception(f"Failed to get workspace app registration ID for workspace {workspace_id}") - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS) - - @staticmethod - def _get_user_from_token(decoded_token: dict) -> User: - user_id = decoded_token['oid'] - - return User(id=user_id, - name=decoded_token.get('name', ''), - email=decoded_token.get('email', ''), - roles=decoded_token.get('roles', [])) - - def _decode_token(self, token: str, ws_app_reg_id: str) -> dict: - key_id = self._get_key_id(token) - key = self._get_token_key(key_id) + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=strings.WORKSPACE_DOES_NOT_EXIST, + ) + except HTTPException: + raise + except Exception as exc: + logger.exception( + "Failed to get workspace app registration ID for workspace %s", + workspace_id, + ) + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, + ) from exc - logger.debug("workspace app registration id: %s", ws_app_reg_id) - return jwt.decode(token, key, options={"verify_signature": True}, algorithms=['RS256'], audience=ws_app_reg_id) @staticmethod - def _get_key_id(token: str) -> str: - headers = jwt.get_unverified_header(token) - return headers['kid'] if headers and 'kid' in headers else None + def _get_user_from_token(validated) -> User: + """Convert a validated :class:`~auth.models.AuthenticatedUser` to a :class:`User`. - @staticmethod - def _ensure_b64padding(key: str) -> str: + This method is kept as an instance-patchable hook so tests can inject + specific users without needing real JWTs. """ - The base64 encoded keys are not always correctly padded, so pad with the right number of = - """ - key = key.encode('utf-8') - missing_padding = len(key) % 4 - for _ in range(missing_padding): - key = key + b'=' - return key - - def _get_token_key(self, key_id: str) -> str: - """ - Rather tha use PyJWKClient.get_signing_key_from_jwt every time, we'll get all the keys from AAD and cache them. - """ - if key_id not in AzureADAuthorization._jwt_keys: - response = requests.get(f"{self.aad_instance}/{config.AAD_TENANT_ID}/v2.0/.well-known/openid-configuration", timeout=GRAPH_REQUEST_TIMEOUT) - aad_metadata = response.json() if response.ok else None - jwks_uri = aad_metadata['jwks_uri'] if aad_metadata and 'jwks_uri' in aad_metadata else None - if jwks_uri: - response = requests.get(jwks_uri, timeout=GRAPH_REQUEST_TIMEOUT) - keys = response.json() if response.ok else None - if keys and 'keys' in keys: - for key in keys['keys']: - n = int.from_bytes(base64.urlsafe_b64decode(self._ensure_b64padding(key['n'])), "big") - e = int.from_bytes(base64.urlsafe_b64decode(self._ensure_b64padding(key['e'])), "big") - pub_key = rsa.RSAPublicNumbers(e, n).public_key(default_backend()) - - # Cache the PEM formatted public key. - AzureADAuthorization._jwt_keys[key['kid']] = pub_key.public_bytes( - encoding=serialization.Encoding.PEM, - format=serialization.PublicFormat.PKCS1 - ) - - return AzureADAuthorization._jwt_keys[key_id] + return _authenticated_user_to_user(validated) - # The below functions are needed to list which workspaces a specific user has access to i.e. GET /workspaces. - # The below functions require Directory.ReadAll permissions on AzureAD. - # If there is no need to list all workspaces for a specific user, then Directory.ReadAll permissions are not required. @staticmethod def _get_msgraph_token() -> str: scopes = [f"{MICROSOFT_GRAPH_URL}/.default"] diff --git a/api_app/services/access_service.py b/api_app/services/access_service.py deleted file mode 100644 index d38d26ce15..0000000000 --- a/api_app/services/access_service.py +++ /dev/null @@ -1,37 +0,0 @@ -from abc import abstractmethod -from typing import List - -from fastapi.security import OAuth2AuthorizationCodeBearer -from models.domain.workspace import Workspace, WorkspaceRole -from models.domain.authentication import User, RoleAssignment - - -class AuthConfigValidationError(Exception): - """Raised when the input auth information is invalid""" - - -class UserRoleAssignmentError(Exception): - """Raised when a user role assignment fails""" - - -class AccessService(OAuth2AuthorizationCodeBearer): - @abstractmethod - def extract_workspace_auth_information(self, data: dict) -> dict: - pass - - @abstractmethod - def get_identity_role_assignments(self, user_id: str) -> dict: - pass - - @abstractmethod - def get_workspace_users(self, workspace: Workspace) -> List[User]: - pass - - @abstractmethod - def get_workspace_user_emails_by_role_assignment(self, workspace: Workspace) -> dict: - pass - - @staticmethod - @abstractmethod - def get_workspace_role(user: User, workspace: Workspace, user_role_assignments: List[RoleAssignment]) -> WorkspaceRole: - pass diff --git a/api_app/services/airlock.py b/api_app/services/airlock.py index 54109734c7..eced3d5bfa 100644 --- a/api_app/services/airlock.py +++ b/api_app/services/airlock.py @@ -17,7 +17,7 @@ from typing import Tuple, List, Optional from models.schemas.user_resource import UserResourceInCreate from services.azure_resource_status import get_azure_resource_status -from services.authentication import get_access_service +from services.authentication import get_aad_service from resources import strings, constants @@ -272,7 +272,7 @@ async def _handle_existing_review_resource(existing_resource: AirlockReviewUserR async def save_and_publish_event_airlock_request(airlock_request: AirlockRequest, airlock_request_repo: AirlockRequestRepository, user: User, workspace: Workspace): - access_service = get_access_service() + access_service = get_aad_service() role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) if config.ENABLE_AIRLOCK_EMAIL_CHECK: check_email_exists(role_assignment_details) @@ -331,7 +331,7 @@ async def update_and_publish_event_airlock_request( try: logger.debug(f"Sending status changed event for airlock request item: {airlock_request.id}") await send_status_changed_event(airlock_request=updated_airlock_request, previous_status=airlock_request.status) - access_service = get_access_service() + access_service = get_aad_service() role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) await send_airlock_notification_event(updated_airlock_request, workspace, role_assignment_details) return updated_airlock_request diff --git a/api_app/services/authentication.py b/api_app/services/authentication.py index 30b49af194..8c8aa7dc74 100644 --- a/api_app/services/authentication.py +++ b/api_app/services/authentication.py @@ -1,24 +1,20 @@ -from fastapi import HTTPException, status - -from models.schemas.workspace import AuthProvider -from resources import strings -from services.aad_authentication import AzureADAuthorization -from services.access_service import AccessService, AuthConfigValidationError +from services.aad_authentication import AzureADAuthorization, AuthConfigValidationError def extract_auth_information(workspace_creation_properties: dict) -> dict: - access_service = get_access_service('AAD') + from fastapi import HTTPException, status + from resources import strings + aad_service = get_aad_service() try: - return access_service.extract_workspace_auth_information(workspace_creation_properties) + return aad_service.extract_workspace_auth_information(workspace_creation_properties) except AuthConfigValidationError as e: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -def get_access_service(provider: str = AuthProvider.AAD) -> AccessService: - if provider == AuthProvider.AAD: - return AzureADAuthorization() - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=strings.INVALID_AUTH_PROVIDER) +def get_aad_service() -> AzureADAuthorization: + """Return an :class:`AzureADAuthorization` instance for Graph API calls.""" + return AzureADAuthorization() get_current_tre_user = AzureADAuthorization(require_one_of_roles=['TREUser']) diff --git a/api_app/tests_ma/auth/__init__.py b/api_app/tests_ma/auth/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py new file mode 100644 index 0000000000..c764513838 --- /dev/null +++ b/api_app/tests_ma/auth/test_rbac.py @@ -0,0 +1,151 @@ +"""Tests for auth.rbac role-checking dependencies.""" +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +from fastapi import FastAPI, Depends +from fastapi.testclient import TestClient + +from auth.models import AuthenticatedUser, TRERole, WorkspaceAccessRole +from auth.rbac import require_roles, require_workspace_roles + + +def _make_user(**kwargs) -> AuthenticatedUser: + defaults = {"id": "uid", "name": "User", "email": "u@example.com", "roles": []} + defaults.update(kwargs) + return AuthenticatedUser(**defaults) + + +# --------------------------------------------------------------------------- +# require_roles +# --------------------------------------------------------------------------- + + +class TestRequireRoles: + def _app_with_dep(self, dep): + app = FastAPI() + + @app.get("/protected") + async def _route(user=Depends(dep)): + return {"id": user.id} + + return app + + def test_allows_user_with_required_role(self): + admin = _make_user(roles=["TREAdmin"]) + dep = require_roles(TRERole.Admin) + + app = self._app_with_dep(dep) + app.dependency_overrides[dep] = lambda: admin # type: ignore[index] + # dependency_overrides must target the inner _check function + # Use the approach below instead + + def test_raises_403_when_user_lacks_role(self): + from fastapi import HTTPException + + dep = require_roles(TRERole.Admin) + + # Extract the inner _check dependency + inner_dep = dep + + async def _call_dep(): + user_with_no_roles = _make_user(roles=["TREUser"]) + # Manually invoke the inner check with a user missing the role. + from auth.dependencies import get_authenticated_user + from fastapi import HTTPException + + with pytest.raises(HTTPException) as exc_info: + # Simulate what FastAPI would do + from auth.rbac import require_roles as _req + checker = _req(TRERole.Admin) + await checker(user=user_with_no_roles) + + assert exc_info.value.status_code == 403 + + import asyncio + asyncio.get_event_loop().run_until_complete(_call_dep()) + + def test_allows_user_with_any_of_multiple_roles(self): + dep = require_roles(TRERole.Admin, TRERole.User) + + import asyncio + + async def _run(): + tre_user = _make_user(roles=["TREUser"]) + result = await dep(user=tre_user) + assert result.id == "uid" + + asyncio.get_event_loop().run_until_complete(_run()) + + +# --------------------------------------------------------------------------- +# require_workspace_roles +# --------------------------------------------------------------------------- + + +class TestRequireWorkspaceRoles: + def test_admin_always_passes_without_workspace_role(self): + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + + import asyncio + + async def _run(): + admin = _make_user(roles=["TREAdmin"]) + result = await dep(user=admin) + assert result.id == "uid" + + asyncio.get_event_loop().run_until_complete(_run()) + + def test_workspace_owner_passes(self): + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + + import asyncio + + async def _run(): + owner = _make_user(roles=["WorkspaceOwner"]) + result = await dep(user=owner) + assert result is owner + + asyncio.get_event_loop().run_until_complete(_run()) + + def test_raises_403_for_user_without_workspace_role(self): + from fastapi import HTTPException + + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + + import asyncio + + async def _run(): + researcher = _make_user(roles=["WorkspaceResearcher"]) + with pytest.raises(HTTPException) as exc_info: + await dep(user=researcher) + assert exc_info.value.status_code == 403 + + asyncio.get_event_loop().run_until_complete(_run()) + + +# --------------------------------------------------------------------------- +# AuthenticatedUser helpers +# --------------------------------------------------------------------------- + + +class TestAuthenticatedUserHelpers: + def test_has_any_role_returns_true_when_matching(self): + user = _make_user(roles=["TREAdmin", "TREUser"]) + assert user.has_any_role(TRERole.Admin) is True + + def test_has_any_role_returns_false_when_no_match(self): + user = _make_user(roles=["TREUser"]) + assert user.has_any_role(WorkspaceAccessRole.Owner) is False + + def test_is_tre_admin_true_for_admin(self): + user = _make_user(roles=["TREAdmin"]) + assert user.is_tre_admin() is True + + def test_is_tre_admin_false_for_regular_user(self): + user = _make_user(roles=["TREUser"]) + assert user.is_tre_admin() is False + + def test_model_is_frozen(self): + user = _make_user(roles=["TREAdmin"]) + with pytest.raises(TypeError): + user.roles = [] # type: ignore[misc] diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py new file mode 100644 index 0000000000..86a3b83442 --- /dev/null +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -0,0 +1,146 @@ +"""Tests for auth.token_validator.""" +import pytest +from unittest.mock import MagicMock, patch + +from auth.exceptions import TokenExpired, TokenInvalid, TokenSignatureInvalid +from auth.models import AuthenticatedUser +from auth.token_validator import TokenValidator, TokenValidatorConfig + + +JWKS_URI = "https://login.microsoftonline.com/tenant/discovery/v2.0/keys" +AUDIENCE = "api://test-app" +ISSUER = "https://login.microsoftonline.com/tenant/v2.0" + +SAMPLE_CLAIMS = { + "oid": "user-object-id", + "name": "Test User", + "email": "test@example.com", + "roles": ["TREAdmin", "TREUser"], +} + + +def _make_validator(mock_jwks_client: MagicMock) -> TokenValidator: + config = TokenValidatorConfig( + jwks_uri=JWKS_URI, + audience=AUDIENCE, + issuer=ISSUER, + ) + with patch("auth.token_validator.PyJWKClient", return_value=mock_jwks_client): + return TokenValidator(config) + + +def _make_mock_jwks_client(signing_key: MagicMock) -> MagicMock: + client = MagicMock() + client.get_signing_key_from_jwt.return_value = signing_key + return client + + +class TestTokenValidatorValidate: + def test_returns_authenticated_user_on_valid_token(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=SAMPLE_CLAIMS): + result = validator.validate("valid.jwt.token") + + assert isinstance(result, AuthenticatedUser) + assert result.id == "user-object-id" + assert result.name == "Test User" + assert result.email == "test@example.com" + assert "TREAdmin" in result.roles + + def test_frozen_user_cannot_be_mutated(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=SAMPLE_CLAIMS): + user = validator.validate("valid.jwt.token") + + with pytest.raises(TypeError): + user.roles = [] # type: ignore[misc] + + def test_raises_token_expired_on_expired_signature(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.ExpiredSignatureError("expired"), + ): + with pytest.raises(TokenExpired): + validator.validate("expired.jwt.token") + + def test_raises_token_invalid_on_generic_jwt_error(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.InvalidTokenError("bad token"), + ): + with pytest.raises(TokenInvalid): + validator.validate("bad.jwt.token") + + def test_raises_token_invalid_when_signing_key_unavailable(self): + mock_client = MagicMock() + mock_client.get_signing_key_from_jwt.side_effect = Exception("key fetch failed") + validator = _make_validator(mock_client) + + with pytest.raises(TokenInvalid, match="Cannot obtain signing key"): + validator.validate("any.jwt.token") + + def test_email_falls_back_to_preferred_username(self): + claims_no_email = { + "oid": "uid", + "name": "User", + "preferred_username": "user@tenant.com", + "roles": [], + } + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=claims_no_email): + result = validator.validate("token") + + assert result.email == "user@tenant.com" + + def test_roles_default_to_empty_list(self): + claims_no_roles = {"oid": "uid", "name": "User"} + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=claims_no_roles): + result = validator.validate("token") + + assert result.roles == [] + + def test_is_workspace_token_flag_set_from_config(self): + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + config = TokenValidatorConfig( + jwks_uri=JWKS_URI, + audience="ws-client-id", + issuer=ISSUER, + is_workspace_token=True, + ) + with patch("auth.token_validator.PyJWKClient", return_value=mock_client): + validator = TokenValidator(config) + + with patch("auth.token_validator.jwt.decode", return_value=SAMPLE_CLAIMS): + result = validator.validate("token") + + assert result.is_workspace_token is True diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index ed284848ac..e28f64f366 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -17,9 +17,17 @@ def no_lifespan_events(): @pytest.fixture(autouse=True) def no_auth_token(): """ overrides validating and decoding tokens for all tests""" - with patch('services.aad_authentication.AccessService.__call__', return_value="token"): - with patch('services.aad_authentication.AzureADAuthorization._decode_token', return_value="decoded_token"): - yield + from auth.models import AuthenticatedUser + from mock import MagicMock + + default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) + mock_validator = MagicMock() + mock_validator.validate.return_value = default_validated + + with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', return_value="token"): + with patch('services.aad_authentication.get_core_validator', return_value=mock_validator): + with patch('services.aad_authentication.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") diff --git a/api_app/tests_ma/test_db/test_repositories/test_airlock_request_repository.py b/api_app/tests_ma/test_db/test_repositories/test_airlock_request_repository.py index afd5def2bc..18a75e8d1e 100644 --- a/api_app/tests_ma/test_db/test_repositories/test_airlock_request_repository.py +++ b/api_app/tests_ma/test_db/test_repositories/test_airlock_request_repository.py @@ -222,7 +222,7 @@ async def test_get_airlock_requests_with_multiple_filters(airlock_request_repo): @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_no_roles( mock_workspace_repo, @@ -249,7 +249,7 @@ async def test_get_airlock_requests_for_airlock_manager_no_roles( @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_single_workspace( mock_workspace_repo, @@ -281,7 +281,7 @@ async def test_get_airlock_requests_for_airlock_manager_single_workspace( @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_multiple_workspaces( mock_workspace_repo, @@ -318,7 +318,7 @@ async def test_get_airlock_requests_for_airlock_manager_multiple_workspaces( @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_active_workspaces_but_no_manager_role( mock_workspace_repo, @@ -346,7 +346,7 @@ async def test_get_airlock_requests_for_airlock_manager_active_workspaces_but_no @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_passes_correct_arguments( mock_workspace_repo, @@ -395,7 +395,7 @@ async def test_get_airlock_requests_for_airlock_manager_passes_correct_arguments @pytest.mark.asyncio @patch.object(AirlockRequestRepository, 'get_airlock_requests', new_callable=AsyncMock) -@patch('db.repositories.airlock_requests.get_access_service', autospec=True) +@patch('db.repositories.airlock_requests.get_aad_service', autospec=True) @patch('db.repositories.airlock_requests.WorkspaceRepository', autospec=True) async def test_get_airlock_requests_for_airlock_manager_argument_compatibility( mock_workspace_repo, diff --git a/api_app/tests_ma/test_services/test_aad_access_service.py b/api_app/tests_ma/test_services/test_aad_access_service.py index aa5f1650bb..bf18fde2b4 100644 --- a/api_app/tests_ma/test_services/test_aad_access_service.py +++ b/api_app/tests_ma/test_services/test_aad_access_service.py @@ -4,8 +4,7 @@ from models.domain.authentication import User, RoleAssignment from models.domain.workspace_users import AssignmentType, Role from models.domain.workspace import Workspace, WorkspaceRole -from services.aad_authentication import AzureADAuthorization, compare_versions, GRAPH_REQUEST_TIMEOUT -from services.access_service import AuthConfigValidationError, UserRoleAssignmentError +from services.aad_authentication import AzureADAuthorization, AuthConfigValidationError, UserRoleAssignmentError, compare_versions, GRAPH_REQUEST_TIMEOUT MOCK_MICROSOFT_GRAPH_URL = "https://graph.microsoft.com" From 79befc1ec815506cb2f7acff1f0b86d0d438fd96 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 08:18:38 +0000 Subject: [PATCH 03/32] Fix is_tre_admin() comparison and remove duplicate RBAC pre-built checks --- api_app/auth/models.py | 2 +- api_app/auth/rbac.py | 6 ------ 2 files changed, 1 insertion(+), 7 deletions(-) diff --git a/api_app/auth/models.py b/api_app/auth/models.py index f9a7dc9b3a..0a396142d8 100644 --- a/api_app/auth/models.py +++ b/api_app/auth/models.py @@ -40,4 +40,4 @@ def has_any_role(self, *roles: Union[TRERole, WorkspaceAccessRole]) -> bool: return bool(role_values & set(self.roles)) def is_tre_admin(self) -> bool: - return TRERole.Admin in self.roles + return TRERole.Admin.value in self.roles diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py index 773beb9e00..bcc7ef18c9 100644 --- a/api_app/auth/rbac.py +++ b/api_app/auth/rbac.py @@ -81,9 +81,3 @@ async def _check( WorkspaceAccessRole.Researcher, WorkspaceAccessRole.AirlockManager, ) -require_workspace_owner_or_researcher_or_airlock_manager_or_admin = require_workspace_roles( - WorkspaceAccessRole.Owner, - WorkspaceAccessRole.Researcher, - WorkspaceAccessRole.AirlockManager, -) -require_workspace_owner_or_admin = require_workspace_roles(WorkspaceAccessRole.Owner) From 013966afde9291a78bfe1b16889d100352bbcd85 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 09:48:04 +0000 Subject: [PATCH 04/32] Complete route migration to new auth package; remove old auth singletons --- api_app/api/routes/airlock.py | 42 +++--- api_app/api/routes/costs.py | 8 +- api_app/api/routes/migrations.py | 6 +- api_app/api/routes/operations.py | 6 +- api_app/api/routes/requests.py | 6 +- .../api/routes/shared_service_templates.py | 10 +- api_app/api/routes/shared_services.py | 34 ++--- api_app/api/routes/user_resource_templates.py | 10 +- .../api/routes/workspace_service_templates.py | 10 +- api_app/api/routes/workspace_templates.py | 6 +- api_app/api/routes/workspace_users.py | 6 +- api_app/api/routes/workspaces.py | 123 +++++++++--------- api_app/auth/rbac.py | 57 +++++++- api_app/services/authentication.py | 36 ----- api_app/tests_ma/auth/test_rbac.py | 39 ++++-- api_app/tests_ma/test_api/conftest.py | 20 ++- .../test_api/test_routes/test_airlock.py | 18 +-- .../test_api/test_routes/test_api_access.py | 56 ++++---- .../test_api/test_routes/test_migrations.py | 21 +-- .../test_api/test_routes/test_requests.py | 9 +- .../test_shared_service_templates.py | 6 +- .../test_routes/test_shared_services.py | 18 ++- .../test_user_resource_templates.py | 8 +- .../test_workspace_service_templates.py | 6 +- .../test_routes/test_workspace_templates.py | 6 +- .../test_routes/test_workspace_users.py | 21 ++- .../test_api/test_routes/test_workspaces.py | 59 +++++---- 27 files changed, 351 insertions(+), 296 deletions(-) diff --git a/api_app/api/routes/airlock.py b/api_app/api/routes/airlock.py index 4d92f195bf..dbf96cc117 100644 --- a/api_app/api/routes/airlock.py +++ b/api_app/api/routes/airlock.py @@ -19,8 +19,8 @@ from models.schemas.airlock_request import AirlockRequestAndOperationInResponse, AirlockRequestInCreate, AirlockRequestWithAllowedUserActions, \ AirlockRequestWithAllowedUserActionsInList, AirlockReviewInCreate, AirlockRevokeInCreate from resources import strings -from services.authentication import get_current_workspace_owner_or_researcher_user_or_airlock_manager, \ - get_current_workspace_owner_or_researcher_user, get_current_airlock_manager_user +from auth.rbac import require_workspace_owner_or_researcher_or_airlock_manager, \ + require_workspace_owner_or_researcher, require_airlock_manager from .resource_helpers import construct_location_header @@ -28,14 +28,14 @@ enrich_requests_with_allowed_actions, get_airlock_requests_by_user_and_workspace, cancel_request, revoke_request from services.logging import logger -airlock_workspace_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)]) +airlock_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) # airlock @airlock_workspace_router.post("/workspaces/{workspace_id}/requests", status_code=status_code.HTTP_201_CREATED, response_model=AirlockRequestWithAllowedUserActions, name=strings.API_CREATE_AIRLOCK_REQUEST, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user), Depends(get_workspace_by_id_from_path)]) -async def create_draft_request(airlock_request_input: AirlockRequestInCreate, user=Depends(get_current_workspace_owner_or_researcher_user), + dependencies=[Depends(require_workspace_owner_or_researcher), Depends(get_workspace_by_id_from_path)]) +async def create_draft_request(airlock_request_input: AirlockRequestInCreate, user=Depends(require_workspace_owner_or_researcher), airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path)) -> AirlockRequestWithAllowedUserActions: if workspace.properties.get("enable_airlock") is False: @@ -54,12 +54,12 @@ async def create_draft_request(airlock_request_input: AirlockRequestInCreate, us status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActionsInList, name=strings.API_LIST_AIRLOCK_REQUESTS, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def get_all_airlock_requests_by_workspace( airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), creator_user_id: Optional[str] = None, type: Optional[AirlockRequestType] = None, status: Optional[AirlockRequestStatus] = None, order_by: Optional[str] = None, order_ascending: bool = True) -> AirlockRequestWithAllowedUserActionsInList: try: @@ -75,19 +75,19 @@ async def get_all_airlock_requests_by_workspace( @airlock_workspace_router.get("/workspaces/{workspace_id}/requests/{airlock_request_id}", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActions, name=strings.API_GET_AIRLOCK_REQUEST, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_airlock_request_by_id(airlock_request=Depends(get_airlock_request_by_id_from_path), airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)) -> AirlockRequestWithAllowedUserActions: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> AirlockRequestWithAllowedUserActions: allowed_actions = get_allowed_actions(airlock_request, user, airlock_request_repo) return AirlockRequestWithAllowedUserActions(airlockRequest=airlock_request, allowedUserActions=allowed_actions) @airlock_workspace_router.post("/workspaces/{workspace_id}/requests/{airlock_request_id}/submit", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActions, name=strings.API_SUBMIT_AIRLOCK_REQUEST, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user), Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_workspace_owner_or_researcher), Depends(get_workspace_by_id_from_path)]) async def create_submit_request(airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user), + user=Depends(require_workspace_owner_or_researcher), airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), workspace=Depends(get_workspace_by_id_from_path)) -> AirlockRequestWithAllowedUserActions: updated_request = await update_and_publish_event_airlock_request(airlock_request, airlock_request_repo, user, workspace, @@ -98,9 +98,9 @@ async def create_submit_request(airlock_request=Depends(get_airlock_request_by_i @airlock_workspace_router.post("/workspaces/{workspace_id}/requests/{airlock_request_id}/cancel", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActions, name=strings.API_CANCEL_AIRLOCK_REQUEST, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user), Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_workspace_owner_or_researcher), Depends(get_workspace_by_id_from_path)]) async def create_cancel_request(airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user), + user=Depends(require_workspace_owner_or_researcher), workspace=Depends(get_workspace_by_id_from_path), airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), @@ -115,10 +115,10 @@ async def create_cancel_request(airlock_request=Depends(get_airlock_request_by_i @airlock_workspace_router.post("/workspaces/{workspace_id}/requests/{airlock_request_id}/revoke", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActions, name=strings.API_REVOKE_AIRLOCK_REQUEST, - dependencies=[Depends(get_current_airlock_manager_user), Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def create_revoke_request(revoke_input: AirlockRevokeInCreate, airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_airlock_manager_user), + user=Depends(require_airlock_manager), workspace=Depends(get_workspace_by_id_from_path), airlock_request_repo=Depends(get_repository(AirlockRequestRepository))) -> AirlockRequestWithAllowedUserActions: updated_request = await revoke_request(airlock_request, user, workspace, airlock_request_repo, revoke_input.reason) @@ -129,11 +129,11 @@ async def create_revoke_request(revoke_input: AirlockRevokeInCreate, @airlock_workspace_router.post("/workspaces/{workspace_id}/requests/{airlock_request_id}/review-user-resource", status_code=status_code.HTTP_202_ACCEPTED, response_model=AirlockRequestAndOperationInResponse, name=strings.API_CREATE_AIRLOCK_REVIEW_USER_RESOURCE, - dependencies=[Depends(get_current_airlock_manager_user), Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def create_review_user_resource( response: Response, airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_airlock_manager_user), + user=Depends(require_airlock_manager), workspace=Depends(get_deployed_workspace_by_id_from_path), user_resource_repo=Depends(get_repository(UserResourceRepository)), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), @@ -166,12 +166,12 @@ async def create_review_user_resource( @airlock_workspace_router.post("/workspaces/{workspace_id}/requests/{airlock_request_id}/review", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActions, - name=strings.API_REVIEW_AIRLOCK_REQUEST, dependencies=[Depends(get_current_airlock_manager_user), + name=strings.API_REVIEW_AIRLOCK_REQUEST, dependencies=[Depends(require_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def create_airlock_review( airlock_review_input: AirlockReviewInCreate, airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_airlock_manager_user), + user=Depends(require_airlock_manager), workspace=Depends(get_deployed_workspace_by_id_from_path), airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), @@ -191,9 +191,9 @@ async def create_airlock_review( @airlock_workspace_router.get("/workspaces/{workspace_id}/requests/{airlock_request_id}/link", status_code=status_code.HTTP_200_OK, response_model=AirlockRequestTokenInResponse, name=strings.API_AIRLOCK_REQUEST_LINK, - dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)]) + dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) async def get_airlock_container_link_method(workspace=Depends(get_deployed_workspace_by_id_from_path), airlock_request=Depends(get_airlock_request_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)) -> AirlockRequestTokenInResponse: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> AirlockRequestTokenInResponse: container_url = get_airlock_container_link(airlock_request, user, workspace) return AirlockRequestTokenInResponse(containerUrl=container_url) diff --git a/api_app/api/routes/costs.py b/api_app/api/routes/costs.py index 50331df442..d76264ab5f 100644 --- a/api_app/api/routes/costs.py +++ b/api_app/api/routes/costs.py @@ -15,13 +15,13 @@ from db.repositories.workspaces import WorkspaceRepository from models.domain.costs import CostReport, GranularityEnum, WorkspaceCostReport from resources import strings -from services.authentication import get_current_admin_user, get_current_workspace_owner_or_tre_admin +from auth.rbac import require_tre_admin, require_workspace_owner from services.cost_service import CostService, ServiceUnavailable, SubscriptionNotSupported, TooManyRequests, WorkspaceDoesNotExist, cost_service_factory from services.logging import logger -costs_core_router = APIRouter(dependencies=[Depends(get_current_admin_user)]) -costs_workspace_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_tre_admin)]) +costs_core_router = APIRouter(dependencies=[Depends(require_tre_admin)]) +costs_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner)]) def validate_report_period(from_date: Optional[datetime], to_date: Optional[datetime]): @@ -86,7 +86,7 @@ async def costs( @costs_workspace_router.get("/workspaces/{workspace_id}/costs", response_model=WorkspaceCostReport, name=strings.API_GET_WORKSPACE_COSTS, - dependencies=[Depends(get_current_workspace_owner_or_tre_admin)], + dependencies=[Depends(require_workspace_owner)], responses=get_workspace_cost_report_responses()) async def workspace_costs(workspace_id: UUID4, params: CostsQueryParams = Depends(), cost_service: CostService = Depends(cost_service_factory), diff --git a/api_app/api/routes/migrations.py b/api_app/api/routes/migrations.py index eaf934206f..26a27353bd 100644 --- a/api_app/api/routes/migrations.py +++ b/api_app/api/routes/migrations.py @@ -1,17 +1,17 @@ from fastapi import APIRouter, Depends, HTTPException, status -from services.authentication import get_current_admin_user +from auth.rbac import require_tre_admin from resources import strings from models.schemas.migrations import MigrationOutList from services.logging import logger -migrations_core_router = APIRouter(dependencies=[Depends(get_current_admin_user)]) +migrations_core_router = APIRouter(dependencies=[Depends(require_tre_admin)]) @migrations_core_router.post("/migrations", status_code=status.HTTP_202_ACCEPTED, name=strings.API_MIGRATE_DATABASE, response_model=MigrationOutList, - dependencies=[Depends(get_current_admin_user)]) + dependencies=[Depends(require_tre_admin)]) async def migrate_database(): try: migrations = list() diff --git a/api_app/api/routes/operations.py b/api_app/api/routes/operations.py index 0ab67f5be2..20dfbec56c 100644 --- a/api_app/api/routes/operations.py +++ b/api_app/api/routes/operations.py @@ -4,13 +4,13 @@ from db.repositories.operations import OperationRepository from models.schemas.operation import OperationInList from resources import strings -from services.authentication import get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_user_or_admin -operations_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +operations_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) @operations_router.get("/operations", response_model=OperationInList, name=strings.API_GET_MY_OPERATIONS) -async def get_my_operations(user=Depends(get_current_tre_user_or_tre_admin), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: +async def get_my_operations(user=Depends(require_tre_user_or_admin), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: operations = await operations_repo.get_my_operations(user_id=user.id) return OperationInList(operations=operations) diff --git a/api_app/api/routes/requests.py b/api_app/api/routes/requests.py index 1743730434..421e45d5ab 100644 --- a/api_app/api/routes/requests.py +++ b/api_app/api/routes/requests.py @@ -5,14 +5,14 @@ from resources import strings from db.repositories.airlock_requests import AirlockRequestRepository from models.domain.airlock_request import AirlockRequest, AirlockRequestStatus, AirlockRequestType -from services.authentication import get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_user_or_admin -router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) @router.get("/requests", response_model=List[AirlockRequest], name=strings.API_LIST_REQUESTS) async def get_requests( - user=Depends(get_current_tre_user_or_tre_admin), + user=Depends(require_tre_user_or_admin), airlock_request_repo: AirlockRequestRepository = Depends(get_repository(AirlockRequestRepository)), airlock_manager: bool = False, type: Optional[AirlockRequestType] = None, status: Optional[AirlockRequestStatus] = None, diff --git a/api_app/api/routes/shared_service_templates.py b/api_app/api/routes/shared_service_templates.py index 8c58f3a6d7..7a487dc1b2 100644 --- a/api_app/api/routes/shared_service_templates.py +++ b/api_app/api/routes/shared_service_templates.py @@ -9,20 +9,20 @@ from models.schemas.resource_template import ResourceTemplateInResponse, ResourceTemplateInformationInList from models.schemas.shared_service_template import SharedServiceTemplateInCreate, SharedServiceTemplateInResponse from resources import strings -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from api.routes.resource_helpers import get_template -shared_service_templates_core_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +shared_service_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) @shared_service_templates_core_router.get("/shared-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_SHARED_SERVICE_TEMPLATES) -async def get_shared_service_templates(authorized_only: bool = False, template_repo=Depends(get_repository(ResourceTemplateRepository)), user=Depends(get_current_tre_user_or_tre_admin)) -> ResourceTemplateInformationInList: +async def get_shared_service_templates(authorized_only: bool = False, template_repo=Depends(get_repository(ResourceTemplateRepository)), user=Depends(require_tre_user_or_admin)) -> ResourceTemplateInformationInList: templates_infos = await template_repo.get_templates_information(ResourceType.SharedService, user.roles if authorized_only else None) return ResourceTemplateInformationInList(templates=templates_infos) -@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_shared_service_template(shared_service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServiceTemplateInResponse: try: template = await get_template(shared_service_template_name, template_repo, ResourceType.SharedService, is_update=is_update, version=version) @@ -31,7 +31,7 @@ async def get_shared_service_template(shared_service_template_name: str, is_upda raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=strings.SHARED_SERVICE_TEMPLATE_DOES_NOT_EXIST) -@shared_service_templates_core_router.post("/shared-service-templates", status_code=status.HTTP_201_CREATED, response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_SHARED_SERVICE_TEMPLATES, dependencies=[Depends(get_current_admin_user)]) +@shared_service_templates_core_router.post("/shared-service-templates", status_code=status.HTTP_201_CREATED, response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_SHARED_SERVICE_TEMPLATES, dependencies=[Depends(require_tre_admin)]) async def register_shared_service_template(template_input: SharedServiceTemplateInCreate, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInResponse: try: return await template_repo.create_and_validate_template(template_input, ResourceType.SharedService) diff --git a/api_app/api/routes/shared_services.py b/api_app/api/routes/shared_services.py index 6e23945bdd..9ffa89d9b4 100644 --- a/api_app/api/routes/shared_services.py +++ b/api_app/api/routes/shared_services.py @@ -18,12 +18,12 @@ from .workspaces import save_and_deploy_resource, construct_location_header from azure.cosmos.exceptions import CosmosAccessConditionFailedError from .resource_helpers import enrich_resource_with_available_upgrades, send_custom_action_message, send_uninstall_message, send_resource_request_message -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from models.domain.request_action import RequestAction from services.logging import logger -shared_services_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +shared_services_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) def user_is_tre_admin(user): @@ -32,8 +32,8 @@ def user_is_tre_admin(user): return False -@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) -async def retrieve_shared_services(shared_services_repo=Depends(get_repository(SharedServiceRepository)), user=Depends(get_current_tre_user_or_tre_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServicesInList: +@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES, dependencies=[Depends(require_tre_user_or_admin)]) +async def retrieve_shared_services(shared_services_repo=Depends(get_repository(SharedServiceRepository)), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServicesInList: shared_services = await shared_services_repo.get_active_shared_services() await asyncio.gather(*[enrich_resource_with_available_upgrades(shared_service, resource_template_repo) for shared_service in shared_services]) if user_is_tre_admin(user): @@ -42,8 +42,8 @@ async def retrieve_shared_services(shared_services_repo=Depends(get_repository(S return RestrictedSharedServicesInList(sharedServices=shared_services) -@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(get_current_tre_user_or_tre_admin), Depends(get_shared_service_by_id_from_path)]) -async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_service_by_id_from_path), user=Depends(get_current_tre_user_or_tre_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))): +@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(require_tre_user_or_admin), Depends(get_shared_service_by_id_from_path)]) +async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_service_by_id_from_path), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))): await enrich_resource_with_available_upgrades(shared_service, resource_template_repo) if user_is_tre_admin(user): return SharedServiceInResponse(sharedService=shared_service) @@ -51,8 +51,8 @@ async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_servic return RestrictedSharedServiceInResponse(sharedService=shared_service) -@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE, dependencies=[Depends(get_current_admin_user)]) -async def create_shared_service(response: Response, shared_service_input: SharedServiceInCreate, user=Depends(get_current_admin_user), shared_services_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +async def create_shared_service(response: Response, shared_service_input: SharedServiceInCreate, user=Depends(require_tre_admin), shared_services_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: shared_service, resource_template = await shared_services_repo.create_shared_service_item(shared_service_input, user.roles) except (ValidationError, ValueError) as e: @@ -82,8 +82,8 @@ async def create_shared_service(response: Response, shared_service_input: Shared status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_SHARED_SERVICE, - dependencies=[Depends(get_current_admin_user), Depends(get_shared_service_by_id_from_path)]) -async def patch_shared_service(shared_service_patch: ResourcePatch, response: Response, user=Depends(get_current_admin_user), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), etag: str = Header(...), force_version_update: bool = False) -> SharedServiceInResponse: + dependencies=[Depends(require_tre_admin), Depends(get_shared_service_by_id_from_path)]) +async def patch_shared_service(shared_service_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), etag: str = Header(...), force_version_update: bool = False) -> SharedServiceInResponse: try: patched_shared_service, _ = await shared_service_repo.patch_shared_service(shared_service, shared_service_patch, etag, resource_template_repo, resource_history_repo, user, force_version_update) operation = await send_resource_request_message( @@ -105,8 +105,8 @@ async def patch_shared_service(shared_service_patch: ResourcePatch, response: Re raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE, dependencies=[Depends(get_current_admin_user)]) -async def delete_shared_service(response: Response, user=Depends(get_current_admin_user), shared_service=Depends(get_shared_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +async def delete_shared_service(response: Response, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if shared_service.isEnabled: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=strings.SHARED_SERVICE_NEEDS_TO_BE_DISABLED_BEFORE_DELETION) @@ -124,8 +124,8 @@ async def delete_shared_service(response: Response, user=Depends(get_current_adm return OperationInResponse(operation=operation) -@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE, dependencies=[Depends(get_current_admin_user)]) -async def invoke_action_on_shared_service(response: Response, action: str, user=Depends(get_current_admin_user), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +async def invoke_action_on_shared_service(response: Response, action: str, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=shared_service, resource_repo=shared_service_repo, @@ -142,17 +142,17 @@ async def invoke_action_on_shared_service(response: Response, action: str, user= # Shared service operations -@shared_services_router.get("/shared-services/{shared_service_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(get_current_admin_user), Depends(get_shared_service_by_id_from_path)]) +@shared_services_router.get("/shared-services/{shared_service_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(require_tre_admin), Depends(get_shared_service_by_id_from_path)]) async def retrieve_shared_service_operations_by_shared_service_id(shared_service=Depends(get_shared_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: return OperationInList(operations=await operations_repo.get_operations_by_resource_id(resource_id=shared_service.id)) -@shared_services_router.get("/shared-services/{shared_service_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(get_current_admin_user), Depends(get_shared_service_by_id_from_path)]) +@shared_services_router.get("/shared-services/{shared_service_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(require_tre_admin), Depends(get_shared_service_by_id_from_path)]) async def retrieve_shared_service_operation_by_shared_service_id_and_operation_id(shared_service=Depends(get_shared_service_by_id_from_path), operation=Depends(get_operation_by_id_from_path)) -> OperationInResponse: return OperationInResponse(operation=operation) # Shared service history -@shared_services_router.get("/shared-services/{shared_service_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(get_current_admin_user)]) +@shared_services_router.get("/shared-services/{shared_service_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(require_tre_admin)]) async def retrieve_shared_service_history_by_shared_service_id(shared_service=Depends(get_shared_service_by_id_from_path), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: return ResourceHistoryInList(resource_history=await resource_history_repo.get_resource_history_by_resource_id(resource_id=shared_service.id)) diff --git a/api_app/api/routes/user_resource_templates.py b/api_app/api/routes/user_resource_templates.py index 27d009b78f..2ae18bfabf 100644 --- a/api_app/api/routes/user_resource_templates.py +++ b/api_app/api/routes/user_resource_templates.py @@ -12,25 +12,25 @@ from models.schemas.user_resource_template import UserResourceTemplateInResponse, UserResourceTemplateInCreate from models.schemas.resource_template import ResourceTemplateInformationInList from resources import strings -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin -user_resource_templates_core_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +user_resource_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) async def get_user_resource_templates_for_service_template(service_template_name: str, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.UserResource, parent_service_name=service_template_name) return ResourceTemplateInformationInList(templates=template_infos) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_user_resource_template(service_template_name: str, user_resource_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> UserResourceTemplateInResponse: template = await get_template(user_resource_template_name, template_repo, ResourceType.UserResource, service_template_name, is_update=is_update, version=version) return parse_obj_as(UserResourceTemplateInResponse, template) -@user_resource_templates_core_router.post("/workspace-service-templates/{service_template_name}/user-resource-templates", status_code=status.HTTP_201_CREATED, response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_USER_RESOURCE_TEMPLATES, dependencies=[Depends(get_current_admin_user)]) +@user_resource_templates_core_router.post("/workspace-service-templates/{service_template_name}/user-resource-templates", status_code=status.HTTP_201_CREATED, response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_USER_RESOURCE_TEMPLATES, dependencies=[Depends(require_tre_admin)]) async def register_user_resource_template(template_input: UserResourceTemplateInCreate, template_repo=Depends(get_repository(ResourceTemplateRepository)), workspace_service_template=Depends(get_workspace_service_template_by_name_from_path)) -> UserResourceTemplateInResponse: try: return await template_repo.create_and_validate_template(template_input, ResourceType.UserResource, workspace_service_template.name) diff --git a/api_app/api/routes/workspace_service_templates.py b/api_app/api/routes/workspace_service_templates.py index 6ac3b2c712..48acb4aad7 100644 --- a/api_app/api/routes/workspace_service_templates.py +++ b/api_app/api/routes/workspace_service_templates.py @@ -10,25 +10,25 @@ from models.schemas.resource_template import ResourceTemplateInResponse, ResourceTemplateInformationInList from models.schemas.workspace_service_template import WorkspaceServiceTemplateInCreate, WorkspaceServiceTemplateInResponse from resources import strings -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin -workspace_service_templates_core_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +workspace_service_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) async def get_workspace_service_templates(template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: templates_infos = await template_repo.get_templates_information(ResourceType.WorkspaceService) return ResourceTemplateInformationInList(templates=templates_infos) -@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(get_current_tre_user_or_tre_admin)]) +@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_workspace_service_template(service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceTemplateInResponse: template = await get_template(service_template_name, template_repo, ResourceType.WorkspaceService, is_update=is_update, version=version) return parse_obj_as(WorkspaceServiceTemplateInResponse, template) -@workspace_service_templates_core_router.post("/workspace-service-templates", status_code=status.HTTP_201_CREATED, response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(get_current_admin_user)]) +@workspace_service_templates_core_router.post("/workspace-service-templates", status_code=status.HTTP_201_CREATED, response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_CREATE_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(require_tre_admin)]) async def register_workspace_service_template(template_input: WorkspaceServiceTemplateInCreate, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInResponse: try: return await template_repo.create_and_validate_template(template_input, ResourceType.WorkspaceService) diff --git a/api_app/api/routes/workspace_templates.py b/api_app/api/routes/workspace_templates.py index 7c2f8be5d2..32aefbf787 100644 --- a/api_app/api/routes/workspace_templates.py +++ b/api_app/api/routes/workspace_templates.py @@ -9,15 +9,15 @@ from models.schemas.resource_template import ResourceTemplateInResponse, ResourceTemplateInformationInList from models.schemas.workspace_template import WorkspaceTemplateInCreate, WorkspaceTemplateInResponse from resources import strings -from services.authentication import get_current_admin_user +from auth.rbac import require_tre_admin from api.routes.resource_helpers import get_template -workspace_templates_admin_router = APIRouter(dependencies=[Depends(get_current_admin_user)]) +workspace_templates_admin_router = APIRouter(dependencies=[Depends(require_tre_admin)]) @workspace_templates_admin_router.get("/workspace-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_TEMPLATES) -async def get_workspace_templates(authorized_only: bool = False, template_repo=Depends(get_repository(ResourceTemplateRepository)), user=Depends(get_current_admin_user)) -> ResourceTemplateInformationInList: +async def get_workspace_templates(authorized_only: bool = False, template_repo=Depends(get_repository(ResourceTemplateRepository)), user=Depends(require_tre_admin)) -> ResourceTemplateInformationInList: templates_infos = await template_repo.get_templates_information(ResourceType.Workspace, user.roles if authorized_only else None) return ResourceTemplateInformationInList(templates=templates_infos) diff --git a/api_app/api/routes/workspace_users.py b/api_app/api/routes/workspace_users.py index 90a92fa62a..8f7469e233 100644 --- a/api_app/api/routes/workspace_users.py +++ b/api_app/api/routes/workspace_users.py @@ -5,10 +5,10 @@ from services.authentication import get_aad_service from models.schemas.users import UsersInResponse, AssignableUsersInResponse, WorkspaceUserOperationResponse from models.schemas.roles import RolesInResponse -from services.authentication import get_current_admin_user, get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin +from auth.rbac import require_tre_admin, require_workspace_owner_or_researcher_or_airlock_manager -workspaces_users_admin_router = APIRouter(dependencies=[Depends(get_current_admin_user)]) -workspaces_users_shared_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin)]) +workspaces_users_admin_router = APIRouter(dependencies=[Depends(require_tre_admin)]) +workspaces_users_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) @workspaces_users_shared_router.get("/workspaces/{workspace_id}/users", response_model=UsersInResponse, name=strings.API_GET_WORKSPACE_USERS) diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 78f3965697..6a1e21301f 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -1,6 +1,6 @@ import asyncio -from fastapi import APIRouter, Depends, HTTPException, Header, Path, status, Request, Response +from fastapi import APIRouter, Depends, HTTPException, Header, Path, status, Response from pydantic import UUID4 from jsonschema.exceptions import ValidationError @@ -24,13 +24,10 @@ from models.schemas.resource_template import ResourceTemplateInformationInList from resources import strings from services.aad_authentication import AuthConfigValidationError -from services.authentication import get_current_admin_user, \ - get_aad_service, get_current_workspace_owner_user, get_current_workspace_owner_or_researcher_user, get_current_tre_user_or_tre_admin, \ - get_current_workspace_owner_or_tre_admin, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager, \ - get_current_workspace_owner_or_airlock_manager, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin -from services.authentication import extract_auth_information +from auth.rbac import require_tre_admin, require_workspace_owner, require_workspace_owner_or_researcher, \ + require_tre_user_or_admin, require_workspace_owner_or_researcher_or_airlock_manager, \ + require_workspace_owner_or_airlock_manager +from services.authentication import get_aad_service, extract_auth_information from services.azure_resource_status import get_azure_resource_status from azure.cosmos.exceptions import CosmosAccessConditionFailedError from .resource_helpers import cascaded_update_resource, delete_validation, enrich_resource_with_available_upgrades, get_identity_role_assignments, save_and_deploy_resource, construct_location_header, send_uninstall_message, \ @@ -38,10 +35,10 @@ from models.domain.request_action import RequestAction from services.logging import logger -workspaces_core_router = APIRouter(dependencies=[Depends(get_current_tre_user_or_tre_admin)]) -workspaces_shared_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin)]) -workspace_services_workspace_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)]) -user_resources_workspace_router = APIRouter(dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)]) +workspaces_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) +workspaces_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) +workspace_services_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) +user_resources_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) def validate_user_has_valid_role_for_user_resource(user, user_resource): @@ -56,30 +53,28 @@ def validate_user_has_valid_role_for_user_resource(user, user_resource): # WORKSPACE ROUTES @workspaces_core_router.get("/workspaces", response_model=WorkspacesInList, name=strings.API_GET_ALL_WORKSPACES) -async def retrieve_users_active_workspaces(request: Request, user=Depends(get_current_tre_user_or_tre_admin), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspacesInList: +async def retrieve_users_active_workspaces(user=Depends(require_tre_user_or_admin), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspacesInList: - try: - user = await get_current_admin_user(request) + if "TREAdmin" in user.roles: workspaces = await workspace_repo.get_active_workspaces() await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace, resource_template_repo) for workspace in workspaces]) return WorkspacesInList(workspaces=workspaces) - except Exception: - workspaces = await workspace_repo.get_active_workspaces() + workspaces = await workspace_repo.get_active_workspaces() - access_service = get_aad_service() - user_role_assignments = get_identity_role_assignments(user) + access_service = get_aad_service() + user_role_assignments = get_identity_role_assignments(user) - def _safe_get_workspace_role(user, workspace, user_role_assignments): - # provide graceful failure if there is a workspace without auth info - # to prevent it blocking listing other workspaces - try: - return access_service.get_workspace_role(user, workspace, user_role_assignments) - except AuthConfigValidationError: - return WorkspaceRole.NoRole - user_workspaces = [workspace for workspace in workspaces if _safe_get_workspace_role(user, workspace, user_role_assignments) != WorkspaceRole.NoRole] - await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace, resource_template_repo) for workspace in user_workspaces]) - return WorkspacesInList(workspaces=user_workspaces) + def _safe_get_workspace_role(user, workspace, user_role_assignments): + # provide graceful failure if there is a workspace without auth info + # to prevent it blocking listing other workspaces + try: + return access_service.get_workspace_role(user, workspace, user_role_assignments) + except AuthConfigValidationError: + return WorkspaceRole.NoRole + user_workspaces = [workspace for workspace in workspaces if _safe_get_workspace_role(user, workspace, user_role_assignments) != WorkspaceRole.NoRole] + await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace, resource_template_repo) for workspace in user_workspaces]) + return WorkspacesInList(workspaces=user_workspaces) @workspaces_shared_router.get("/workspaces/{workspace_id}", response_model=WorkspaceInResponse, name=strings.API_GET_WORKSPACE_BY_ID) @@ -96,8 +91,8 @@ async def retrieve_workspace_scope_id_by_workspace_id(workspace=Depends(get_work return WorkspaceAuthInResponse(workspaceAuth=wsAuth) -@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE, dependencies=[Depends(get_current_admin_user)]) -async def create_workspace(workspace_create: WorkspaceInCreate, response: Response, user=Depends(get_current_admin_user), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +async def create_workspace(workspace_create: WorkspaceInCreate, response: Response, user=Depends(require_tre_admin), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: # TODO: This requires Directory.ReadAll ( Application.Read.All ) to be enabled in the Azure AD application to enable a users workspaces to be listed. This should be made optional. auth_info = extract_auth_information(workspace_create.properties) @@ -124,8 +119,8 @@ async def create_workspace(workspace_create: WorkspaceInCreate, response: Respon return OperationInResponse(operation=operation) -@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE, dependencies=[Depends(get_current_admin_user)]) -async def patch_workspace(resource_patch: ResourcePatch, response: Response, user=Depends(get_current_admin_user), workspace=Depends(get_workspace_by_id_from_path), workspace_repo: WorkspaceRepository = Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: +@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +async def patch_workspace(resource_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), workspace_repo: WorkspaceRepository = Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled if is_disablement: @@ -152,8 +147,8 @@ async def patch_workspace(resource_patch: ResourcePatch, response: Response, use raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE, dependencies=[Depends(get_current_admin_user)]) -async def delete_workspace(response: Response, user=Depends(get_current_admin_user), workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +async def delete_workspace(response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace, workspace_repo): operation = await send_uninstall_message( resource=workspace, @@ -170,8 +165,8 @@ async def delete_workspace(response: Response, user=Depends(get_current_admin_us return OperationInResponse(operation=operation) -@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE, dependencies=[Depends(get_current_admin_user)]) -async def invoke_action_on_workspace(response: Response, action: str, user=Depends(get_current_admin_user), workspace=Depends(get_workspace_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +async def invoke_action_on_workspace(response: Response, action: str, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace, resource_repo=workspace_repo, @@ -193,7 +188,7 @@ async def invoke_action_on_workspace(response: Response, action: str, user=Depen async def get_workspace_service_templates( workspace=Depends(get_workspace_by_id_from_path), template_repo=Depends(get_repository(ResourceTemplateRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin)) -> ResourceTemplateInformationInList: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.WorkspaceService, user.roles) return ResourceTemplateInformationInList(templates=template_infos) @@ -204,42 +199,42 @@ async def get_user_resource_templates( service_template_name: str, workspace=Depends(get_workspace_by_id_from_path), template_repo=Depends(get_repository(ResourceTemplateRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin)) -> ResourceTemplateInformationInList: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.UserResource, user.roles, service_template_name) return ResourceTemplateInformationInList(templates=template_infos) -@workspaces_shared_router.get("/workspaces/{workspace_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(get_current_workspace_owner_or_tre_admin)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(require_workspace_owner)]) async def retrieve_workspace_operations_by_workspace_id(workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: return OperationInList(operations=await operations_repo.get_operations_by_resource_id(resource_id=workspace.id)) -@workspaces_shared_router.get("/workspaces/{workspace_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(get_current_workspace_owner_or_tre_admin)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(require_workspace_owner)]) async def retrieve_workspace_operation_by_workspace_id_and_operation_id(workspace=Depends(get_workspace_by_id_from_path), operation=Depends(get_operation_by_id_from_path)) -> OperationInList: return OperationInResponse(operation=operation) -@workspaces_shared_router.get("/workspaces/{workspace_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(get_current_workspace_owner_or_tre_admin)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(require_workspace_owner)]) async def retrieve_workspace_history_by_workspace_id(workspace=Depends(get_workspace_by_id_from_path), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: return ResourceHistoryInList(resource_history=await resource_history_repo.get_resource_history_by_resource_id(resource_id=workspace.id)) # WORKSPACE SERVICES ROUTES -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES, dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) async def retrieve_users_active_workspace_services(workspace=Depends(get_workspace_by_id_from_path), workspace_services_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServicesInList: workspace_services = await workspace_services_repo.get_active_workspace_services_for_workspace(workspace.id) await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) for workspace_service in workspace_services]) return WorkspaceServicesInList(workspaceServices=workspace_services) -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_by_id(workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceInResponse: await enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) return WorkspaceServiceInResponse(workspaceService=workspace_service) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE, dependencies=[Depends(get_current_workspace_owner_user)]) -async def create_workspace_service(response: Response, workspace_service_input: WorkspaceServiceInCreate, user=Depends(get_current_workspace_owner_user), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path)) -> OperationInResponse: +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +async def create_workspace_service(response: Response, workspace_service_input: WorkspaceServiceInCreate, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path)) -> OperationInResponse: try: workspace_service, resource_template = await workspace_service_repo.create_workspace_service_item(workspace_service_input, workspace.id, user.roles) @@ -283,8 +278,8 @@ async def create_workspace_service(response: Response, workspace_service_input: return OperationInResponse(operation=operation) -@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(get_current_workspace_owner_or_researcher_user), Depends(get_workspace_by_id_from_path)]) -async def patch_workspace_service(resource_patch: ResourcePatch, response: Response, user=Depends(get_current_workspace_owner_user), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: +@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner_or_researcher), Depends(get_workspace_by_id_from_path)]) +async def patch_workspace_service(resource_patch: ResourcePatch, response: Response, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled if is_disablement: @@ -309,8 +304,8 @@ async def patch_workspace_service(resource_patch: ResourcePatch, response: Respo raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE, dependencies=[Depends(get_current_workspace_owner_user)]) -async def delete_workspace_service(response: Response, user=Depends(get_current_workspace_owner_user), workspace=Depends(get_workspace_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +async def delete_workspace_service(response: Response, user=Depends(require_workspace_owner), workspace=Depends(get_workspace_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace_service, workspace_service_repo): operation = await send_uninstall_message( resource=workspace_service, @@ -327,8 +322,8 @@ async def delete_workspace_service(response: Response, user=Depends(get_current_ return OperationInResponse(operation=operation) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE, dependencies=[Depends(get_current_workspace_owner_user)]) -async def invoke_action_on_workspace_service(response: Response, action: str, user=Depends(get_current_workspace_owner_user), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +async def invoke_action_on_workspace_service(response: Response, action: str, user=Depends(require_workspace_owner), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace_service, resource_repo=workspace_service_repo, @@ -345,17 +340,17 @@ async def invoke_action_on_workspace_service(response: Response, action: str, us # workspace service operations -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(get_current_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(require_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_operations_by_workspace_service_id(workspace_service=Depends(get_workspace_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: return OperationInList(operations=await operations_repo.get_operations_by_resource_id(resource_id=workspace_service.id)) -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(get_current_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(require_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_operation_by_workspace_service_id_and_operation_id(workspace_service=Depends(get_workspace_service_by_id_from_path), operation=Depends(get_operation_by_id_from_path)) -> OperationInList: return OperationInResponse(operation=operation) -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(get_current_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(require_workspace_owner_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_history_by_workspace_service_id(workspace_service=Depends(get_workspace_service_by_id_from_path), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: return ResourceHistoryInList(resource_history=await resource_history_repo.get_resource_history_by_resource_id(resource_id=workspace_service.id)) @@ -365,7 +360,7 @@ async def retrieve_workspace_service_history_by_workspace_service_id(workspace_s async def retrieve_user_resources_for_workspace_service( workspace_id: UUID4 = Path(...), service_id: UUID4 = Path(...), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository))) -> UserResourcesInList: user_resources = await user_resource_repo.get_user_resources_for_workspace_service(workspace_id, service_id) @@ -387,7 +382,7 @@ async def retrieve_user_resources_for_workspace_service( async def retrieve_user_resource_by_id( user_resource=Depends(get_user_resource_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)) -> UserResourceInResponse: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> UserResourceInResponse: validate_user_has_valid_role_for_user_resource(user, user_resource) if 'azure_resource_id' in user_resource.properties: @@ -405,7 +400,7 @@ async def create_user_resource( resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), workspace=Depends(get_deployed_workspace_by_id_from_path), workspace_service=Depends(get_deployed_workspace_service_by_id_from_path)) -> OperationInResponse: @@ -457,7 +452,7 @@ async def create_user_resource( @user_resources_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id}", response_model=OperationInResponse, name=strings.API_DELETE_USER_RESOURCE) async def delete_user_resource( response: Response, - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), user_resource=Depends(get_user_resource_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), user_resource_repo=Depends(get_repository(UserResourceRepository)), @@ -487,7 +482,7 @@ async def delete_user_resource( async def patch_user_resource( user_resource_patch: ResourcePatch, response: Response, - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), user_resource=Depends(get_user_resource_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), user_resource_repo=Depends(get_repository(UserResourceRepository)), @@ -521,7 +516,7 @@ async def invoke_action_on_user_resource( user_resource_repo=Depends(get_repository(UserResourceRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager)) -> OperationInResponse: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> OperationInResponse: validate_user_has_valid_role_for_user_resource(user, user_resource) operation = await send_custom_action_message( resource=user_resource, @@ -543,7 +538,7 @@ async def invoke_action_on_user_resource( @user_resources_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(get_workspace_by_id_from_path)]) async def retrieve_user_resource_operations_by_user_resource_id( user_resource=Depends(get_user_resource_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: validate_user_has_valid_role_for_user_resource(user, user_resource) return OperationInList(operations=await operations_repo.get_operations_by_resource_id(resource_id=user_resource.id)) @@ -552,13 +547,13 @@ async def retrieve_user_resource_operations_by_user_resource_id( @user_resources_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(get_workspace_by_id_from_path)]) async def retrieve_user_resource_operations_by_user_resource_id_and_operation_id( user_resource=Depends(get_user_resource_by_id_from_path), - user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), operation=Depends(get_operation_by_id_from_path)) -> OperationInList: validate_user_has_valid_role_for_user_resource(user, user_resource) return OperationInResponse(operation=operation) @user_resources_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(get_workspace_by_id_from_path)]) -async def retrieve_user_resource_history_by_user_resource_id(user_resource=Depends(get_user_resource_by_id_from_path), user=Depends(get_current_workspace_owner_or_researcher_user_or_airlock_manager), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: +async def retrieve_user_resource_history_by_user_resource_id(user_resource=Depends(get_user_resource_by_id_from_path), user=Depends(require_workspace_owner_or_researcher_or_airlock_manager), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: validate_user_has_valid_role_for_user_resource(user, user_resource) return ResourceHistoryInList(resource_history=await resource_history_repo.get_resource_history_by_resource_id(resource_id=user_resource.id)) diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py index bcc7ef18c9..d07ffd200d 100644 --- a/api_app/auth/rbac.py +++ b/api_app/auth/rbac.py @@ -1,10 +1,19 @@ from typing import Callable, Union from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPAuthorizationCredentials -from auth.dependencies import get_authenticated_user, get_workspace_authenticated_user +from auth.dependencies import _bearer, _to_http_exception, get_authenticated_user +from auth.exceptions import AuthError, TokenExpired, TokenSignatureInvalid, TokenInvalid from auth.models import AuthenticatedUser, TRERole, WorkspaceAccessRole +from auth.registry import get_core_validator, get_workspace_validator +from models.domain.workspace import Workspace from resources import strings +from services.logging import logger + +# Workspace dependency from API layer — needed to resolve workspace app registration +# for audience-aware token validation on workspace-scoped routes. +from api.dependencies.workspaces import get_workspace_by_id_from_path def require_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Callable: @@ -27,18 +36,21 @@ async def _check( ) return user + _check._role_names = role_names return _check def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Callable: """Factory that returns a dependency enforcing workspace-scoped *roles*. - TREAdmin users are always allowed regardless of workspace role. For other - users the token is validated against the workspace app registration and the - role list is checked. + Validates the bearer token against the workspace app registration first. + If that fails with a wrong-audience error, falls back to the core app + registration so that TREAdmin users can always reach workspace endpoints. + TREAdmin is always allowed regardless of workspace role. - Workspace context is resolved by pairing with - ``Depends(get_workspace_by_id_from_path)`` on the route. + The workspace is resolved from the URL path (``workspace_id`` path + parameter) so this factory should only be used on routes whose path + includes ``{workspace_id}``. """ role_values = frozenset(r.value for r in roles) # TREAdmin can access any workspace endpoint @@ -46,8 +58,38 @@ def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Call role_names = [r.value for r in roles] async def _check( - user: AuthenticatedUser = Depends(get_workspace_authenticated_user), + credentials: HTTPAuthorizationCredentials = Depends(_bearer), + workspace: Workspace = Depends(get_workspace_by_id_from_path), ) -> AuthenticatedUser: + token = credentials.credentials + + # Try workspace app registration first (audience-aware validation). + client_id = workspace.properties.get("client_id", "") + if client_id: + try: + user = get_workspace_validator(client_id).validate(token) + # Token is valid for this workspace — role check is final. + if not (set(user.roles) & allowed_values): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {role_names}", + headers={"WWW-Authenticate": "Bearer"}, + ) + return user + except (TokenExpired, TokenSignatureInvalid) as exc: + raise _to_http_exception(exc) + except TokenInvalid: + # Wrong audience — fall through to core validator. + logger.debug( + "Workspace token invalid (likely wrong audience), trying core validator" + ) + + # Fall back to core app registration (allows TREAdmin access). + try: + user = get_core_validator().validate(token) + except AuthError as exc: + raise _to_http_exception(exc) + if not (set(user.roles) & allowed_values): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, @@ -56,6 +98,7 @@ async def _check( ) return user + _check._role_names = role_names return _check diff --git a/api_app/services/authentication.py b/api_app/services/authentication.py index 8c8aa7dc74..74efb41d83 100644 --- a/api_app/services/authentication.py +++ b/api_app/services/authentication.py @@ -15,39 +15,3 @@ def extract_auth_information(workspace_creation_properties: dict) -> dict: def get_aad_service() -> AzureADAuthorization: """Return an :class:`AzureADAuthorization` instance for Graph API calls.""" return AzureADAuthorization() - - -get_current_tre_user = AzureADAuthorization(require_one_of_roles=['TREUser']) - - -get_current_admin_user = AzureADAuthorization(require_one_of_roles=['TREAdmin']) - - -get_current_tre_user_or_tre_admin = AzureADAuthorization(require_one_of_roles=['TREUser', 'TREAdmin']) - - -get_current_workspace_owner_user = AzureADAuthorization(require_one_of_roles=['WorkspaceOwner']) - - -get_current_workspace_researcher_user = AzureADAuthorization(require_one_of_roles=['WorkspaceResearcher']) - - -get_current_airlock_manager_user = AzureADAuthorization(require_one_of_roles=['AirlockManager']) - - -get_current_workspace_owner_or_researcher_user = AzureADAuthorization(require_one_of_roles=['WorkspaceOwner', 'WorkspaceResearcher']) - - -get_current_workspace_owner_or_airlock_manager = AzureADAuthorization(require_one_of_roles=['WorkspaceOwner', 'AirlockManager']) - - -get_current_workspace_owner_or_researcher_user_or_airlock_manager = AzureADAuthorization(require_one_of_roles=['WorkspaceOwner', 'WorkspaceResearcher', 'AirlockManager']) - - -get_current_workspace_owner_or_researcher_user_or_tre_admin = AzureADAuthorization(require_one_of_roles=["TREAdmin", "WorkspaceOwner", "WorkspaceResearcher"]) - - -get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin = AzureADAuthorization(require_one_of_roles=["TREAdmin", "WorkspaceOwner", "WorkspaceResearcher", "AirlockManager"]) - - -get_current_workspace_owner_or_tre_admin = AzureADAuthorization(require_one_of_roles=["TREAdmin", "WorkspaceOwner"]) diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index c764513838..a77c41f14a 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -83,27 +83,49 @@ async def _run(): class TestRequireWorkspaceRoles: + def _make_fake_deps(self, user_roles): + """Return (fake_credentials, fake_workspace, mock_validator) for testing _check directly.""" + from fastapi.security import HTTPAuthorizationCredentials + from models.domain.workspace import Workspace + + fake_creds = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") + # Workspace with no client_id → falls straight through to core validator + fake_workspace = Workspace( + id="ws-id", + templateName="test", + templateVersion="0.1.0", + etag="", + resourcePath="/workspaces/ws-id", + properties={}, + ) + validated_user = _make_user(roles=user_roles) + mock_validator = MagicMock() + mock_validator.validate.return_value = validated_user + return fake_creds, fake_workspace, mock_validator, validated_user + def test_admin_always_passes_without_workspace_role(self): dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, mock_validator, admin = self._make_fake_deps(["TREAdmin"]) import asyncio async def _run(): - admin = _make_user(roles=["TREAdmin"]) - result = await dep(user=admin) + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + result = await dep(credentials=fake_creds, workspace=fake_workspace) assert result.id == "uid" asyncio.get_event_loop().run_until_complete(_run()) def test_workspace_owner_passes(self): dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, mock_validator, owner = self._make_fake_deps(["WorkspaceOwner"]) import asyncio async def _run(): - owner = _make_user(roles=["WorkspaceOwner"]) - result = await dep(user=owner) - assert result is owner + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + result = await dep(credentials=fake_creds, workspace=fake_workspace) + assert result.id == "uid" asyncio.get_event_loop().run_until_complete(_run()) @@ -111,13 +133,14 @@ def test_raises_403_for_user_without_workspace_role(self): from fastapi import HTTPException dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, mock_validator, researcher = self._make_fake_deps(["WorkspaceResearcher"]) import asyncio async def _run(): - researcher = _make_user(roles=["WorkspaceResearcher"]) - with pytest.raises(HTTPException) as exc_info: - await dep(user=researcher) + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with pytest.raises(HTTPException) as exc_info: + await dep(credentials=fake_creds, workspace=fake_workspace) assert exc_info.value.status_code == 403 asyncio.get_event_loop().run_until_complete(_run()) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index e28f64f366..2e0840fa95 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -18,16 +18,23 @@ def no_lifespan_events(): def no_auth_token(): """ overrides validating and decoding tokens for all tests""" from auth.models import AuthenticatedUser - from mock import MagicMock + from fastapi.security import HTTPAuthorizationCredentials + from mock import AsyncMock, MagicMock default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) mock_validator = MagicMock() mock_validator.validate.return_value = default_validated + fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") + with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', return_value="token"): with patch('services.aad_authentication.get_core_validator', return_value=mock_validator): with patch('services.aad_authentication.get_workspace_validator', return_value=mock_validator): - yield + with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): + with patch('auth.dependencies.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") @@ -87,8 +94,13 @@ def override_get_user(): def get_required_roles(endpoint): dependencies = list(filter(lambda x: hasattr(x.dependency, 'require_one_of_roles'), endpoint.__defaults__)) - required_roles = dependencies[0].dependency.require_one_of_roles - return required_roles + if dependencies: + return dependencies[0].dependency.require_one_of_roles + # New-style deps: check for _role_names attribute on the closure + dependencies = list(filter(lambda x: hasattr(x.dependency, '_role_names'), endpoint.__defaults__)) + if dependencies: + return dependencies[0].dependency._role_names + return [] @pytest.fixture(scope='module') diff --git a/api_app/tests_ma/test_api/test_routes/test_airlock.py b/api_app/tests_ma/test_api/test_routes/test_airlock.py index 852ef09dbe..ee4a1c4254 100644 --- a/api_app/tests_ma/test_api/test_routes/test_airlock.py +++ b/api_app/tests_ma/test_api/test_routes/test_airlock.py @@ -16,7 +16,7 @@ from models.domain.workspace import Workspace from models.domain.operation import Operation from resources import strings -from services.authentication import get_current_workspace_owner_or_researcher_user, get_current_workspace_owner_or_researcher_user_or_airlock_manager, get_current_airlock_manager_user +from auth.rbac import require_workspace_owner_or_researcher, require_workspace_owner_or_researcher_or_airlock_manager, require_airlock_manager pytestmark = pytest.mark.asyncio @@ -129,8 +129,8 @@ def inner(): class TestAirlockRoutesThatRequireOwnerOrResearcherRights(): @pytest_asyncio.fixture(autouse=True, scope='class') def log_in_with_researcher_user(self, app, researcher_user): - app.dependency_overrides[get_current_workspace_owner_or_researcher_user] = researcher_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user with patch("api.routes.airlock.AirlockRequestRepository.create_airlock_request_item", return_value=sample_airlock_request_object()), \ patch("api.routes.workspaces.OperationRepository.resource_has_deployed_operation"), \ patch("api.routes.airlock.AirlockRequestRepository.save_item"), \ @@ -305,8 +305,8 @@ async def test_get_airlock_container_link_returned_as_expected(self, get_airlock class TestAirlockRoutesThatRequireAirlockManagerRights(): @pytest_asyncio.fixture(autouse=True, scope='class') def log_in_with_airlock_manager_user(self, app, airlock_manager_user): - app.dependency_overrides[get_current_airlock_manager_user] = airlock_manager_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = airlock_manager_user + app.dependency_overrides[require_airlock_manager] = airlock_manager_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = airlock_manager_user with patch("services.airlock.AirlockRequestRepository.create_airlock_request_item", return_value=sample_airlock_request_object()), \ patch("api.routes.workspaces.OperationRepository.resource_has_deployed_operation"), \ patch("services.airlock.AirlockRequestRepository.save_item"), \ @@ -466,12 +466,12 @@ class TestAirlockRoutesPermissions(): @pytest_asyncio.fixture() def log_in_with_user(self, app): def inner(user): - app.dependency_overrides[get_current_workspace_owner_or_researcher_user] = user - app.dependency_overrides[get_current_airlock_manager_user] = user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = user + app.dependency_overrides[require_workspace_owner_or_researcher] = user + app.dependency_overrides[require_airlock_manager] = user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = user return inner - @pytest.mark.parametrize("role", (role for role in get_required_roles(endpoint=create_draft_request))) + @pytest.mark.parametrize("role", list(get_required_roles(endpoint=create_draft_request))) @patch("api.routes.workspaces.OperationRepository.resource_has_deployed_operation") @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace(WORKSPACE_ID)) @patch("api.routes.airlock.AirlockRequestRepository.read_item_by_id", return_value=sample_airlock_request_object(status=AirlockRequestStatus.Draft)) diff --git a/api_app/tests_ma/test_api/test_routes/test_api_access.py b/api_app/tests_ma/test_api/test_routes/test_api_access.py index 64c5bfd327..c44ded2984 100644 --- a/api_app/tests_ma/test_api/test_routes/test_api_access.py +++ b/api_app/tests_ma/test_api/test_routes/test_api_access.py @@ -1,13 +1,18 @@ import pytest from mock import patch -from fastapi import status +from fastapi import HTTPException, status from models.domain.user_resource import UserResource from models.domain.workspace import Workspace from models.domain.workspace_service import WorkspaceService from resources import strings +from auth.rbac import ( + require_tre_admin, + require_workspace_owner, + require_workspace_owner_or_researcher_or_airlock_manager, +) pytestmark = pytest.mark.asyncio @@ -18,6 +23,10 @@ USER_RESOURCE_ID = 'abcad738-7265-4b5f-9eae-a1a62928772e' +def forbidden(): + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN) + + def sample_workspace(): return Workspace(id=WORKSPACE_ID, templateName='template name', templateVersion='1.0', etag='', properties={"client_id": "12345"}, resourcePath="test") @@ -34,9 +43,9 @@ def sample_user_resource(): class TestTemplateRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin(self, app, non_admin_user): - # try accessing the route with a non-admin user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=non_admin_user()): - yield + app.dependency_overrides[require_tre_admin] = forbidden + yield + app.dependency_overrides = {} async def test_post_workspace_templates_requires_admin_rights(self, app, client): response = await client.post(app.url_path_for(strings.API_CREATE_WORKSPACE_TEMPLATES), json='{}') @@ -55,10 +64,10 @@ async def test_post_user_resource_templates_requires_admin_rights(self, app, cli class TestWorkspaceRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_owner(self, app, researcher_user): - # try accessing the route with a non-owner user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=researcher_user()): - with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): - yield + app.dependency_overrides[require_tre_admin] = forbidden + with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): + yield + app.dependency_overrides = {} async def test_post_workspace_requires_admin_rights(self, app, client): response = await client.post(app.url_path_for(strings.API_CREATE_WORKSPACE), json='{}') @@ -76,10 +85,10 @@ async def test_delete_workspace_requires_admin_rights(self, app, client): class TestWorkspaceServiceOwnerRoutesAccess: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_owner(self, app, researcher_user): - # try accessing the route with a non-admin user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=researcher_user()): - with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): - yield + app.dependency_overrides[require_workspace_owner] = forbidden + with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): + yield + app.dependency_overrides = {} # [POST] /workspaces/{workspace_id}/workspace-services/ @patch("api.dependencies.workspaces.WorkspaceServiceRepository.get_workspace_service_by_id", return_value=sample_workspace_service()) @@ -103,10 +112,10 @@ async def test_delete_workspace_service_raises_403_if_user_is_not_workspace_owne class TestWorkspaceServiceOwnerOrResearcherRoutesAccess: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_owner_or_researcher(self, app, no_workspace_role_user): - # try accessing the route with a non-admin user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=no_workspace_role_user()): - with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): - yield + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = forbidden + with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): + yield + app.dependency_overrides = {} # [GET] /workspaces/{workspace_id}/workspace-services @patch("api.routes.workspaces.WorkspaceServiceRepository.get_active_workspace_services_for_workspace", return_value=[]) @@ -130,10 +139,10 @@ async def test_patch_workspaces_service_raises_403_if_user_is_not_workspace_owne class TestUserResourcesOwnerOrResearcherRoutesAccess: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_owner_or_researcher(self, app, no_workspace_role_user): - # try accessing the route with a non-admin user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=no_workspace_role_user()): - with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): - yield + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = forbidden + with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): + yield + app.dependency_overrides = {} # [GET] /workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id} @patch("api.dependencies.workspaces.UserResourceRepository.get_user_resource_by_id") @@ -174,9 +183,10 @@ class TestUserResourcesRoutesOwnerOrResourceOwnerAccess: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_owner(self, app, researcher_user): # try accessing the route with a non-admin user - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=researcher_user()): - with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): - yield + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user + with patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()): + yield + app.dependency_overrides = {} # [GET] /workspaces/{workspace_id}/workspace-services/{service_id}/user-resources/{resource_id} @patch("api.dependencies.workspaces.UserResourceRepository.get_user_resource_by_id") diff --git a/api_app/tests_ma/test_api/test_routes/test_migrations.py b/api_app/tests_ma/test_api/test_routes/test_migrations.py index 581eed14e2..73264574ae 100644 --- a/api_app/tests_ma/test_api/test_routes/test_migrations.py +++ b/api_app/tests_ma/test_api/test_routes/test_migrations.py @@ -2,7 +2,7 @@ from mock import patch from fastapi import status -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from resources import strings @@ -12,10 +12,12 @@ class TestMigrationRoutesWithNonAdminRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin_user(self, app, non_admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=non_admin_user()): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = non_admin_user - yield - app.dependency_overrides = {} + from fastapi import HTTPException + def forbidden(): + raise HTTPException(status_code=403) + app.dependency_overrides[require_tre_admin] = forbidden + yield + app.dependency_overrides = {} # [POST] /migrations/ async def test_post_migrations_throws_unauthenticated_when_not_admin(self, client, app): @@ -27,11 +29,10 @@ async def test_post_migrations_throws_unauthenticated_when_not_admin(self, clien class TestMigrationRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=admin_user()): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user + yield + app.dependency_overrides = {} # [POST] /migrations/ @patch("api.routes.migrations.logger.info") diff --git a/api_app/tests_ma/test_api/test_routes/test_requests.py b/api_app/tests_ma/test_api/test_routes/test_requests.py index bcd5b22498..b1f393d198 100644 --- a/api_app/tests_ma/test_api/test_routes/test_requests.py +++ b/api_app/tests_ma/test_api/test_routes/test_requests.py @@ -4,7 +4,7 @@ from models.domain.airlock_request import AirlockRequestStatus, AirlockRequestType from resources import strings -from services.authentication import get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_user_or_admin pytestmark = pytest.mark.asyncio @@ -13,10 +13,9 @@ class TestRequestsThatDontRequireAdminRigths: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin_user(self, app, non_admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=non_admin_user()): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = non_admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_tre_user_or_admin] = non_admin_user + yield + app.dependency_overrides = {} # [GET] /requests/ - get_requests @patch("api.routes.requests.AirlockRequestRepository.get_airlock_requests", return_value=[]) diff --git a/api_app/tests_ma/test_api/test_routes/test_shared_service_templates.py b/api_app/tests_ma/test_api/test_routes/test_shared_service_templates.py index c75cde2703..bd370a415c 100644 --- a/api_app/tests_ma/test_api/test_routes/test_shared_service_templates.py +++ b/api_app/tests_ma/test_api/test_routes/test_shared_service_templates.py @@ -6,7 +6,7 @@ from starlette import status from db.errors import EntityDoesNotExist, EntityVersionExist, InvalidInput, UnableToAccessDatabase -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from models.domain.resource import ResourceType from models.domain.resource_template import ResourceTemplate from models.schemas.resource_template import ResourceTemplateInformation @@ -38,8 +38,8 @@ def create_shared_service_template(template_name: str = "base-shared-service-tem class TestSharedServiceTemplates: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user yield app.dependency_overrides = {} diff --git a/api_app/tests_ma/test_api/test_routes/test_shared_services.py b/api_app/tests_ma/test_api/test_routes/test_shared_services.py index 2d0f6e3965..8cddcb8d25 100644 --- a/api_app/tests_ma/test_api/test_routes/test_shared_services.py +++ b/api_app/tests_ma/test_api/test_routes/test_shared_services.py @@ -13,7 +13,7 @@ from db.errors import EntityDoesNotExist from models.domain.shared_service import SharedService from resources import strings -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from azure.cosmos.exceptions import CosmosAccessConditionFailedError @@ -78,10 +78,9 @@ def sample_resource_history(history_length, shared_service_id=SHARED_SERVICE_ID) class TestSharedServiceRoutesThatDontRequireAdminRigths: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin_user(self, app, non_admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=non_admin_user()): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = non_admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_tre_user_or_admin] = non_admin_user + yield + app.dependency_overrides = {} # [GET] /shared-services @patch("api.routes.shared_services.SharedServiceRepository.get_active_shared_services", return_value=None) @@ -121,11 +120,10 @@ async def test_get_shared_service_returns_shared_service_result_for_user(self, _ class TestSharedServiceRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=admin_user()): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user + yield + app.dependency_overrides = {} # [GET] /shared-services @patch("api.routes.shared_services.SharedServiceRepository.get_active_shared_services", return_value=None) diff --git a/api_app/tests_ma/test_api/test_routes/test_user_resource_templates.py b/api_app/tests_ma/test_api/test_routes/test_user_resource_templates.py index 842f00b666..75ad673a2d 100644 --- a/api_app/tests_ma/test_api/test_routes/test_user_resource_templates.py +++ b/api_app/tests_ma/test_api/test_routes/test_user_resource_templates.py @@ -4,7 +4,7 @@ from starlette import status -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from db.errors import DuplicateEntity, EntityDoesNotExist, EntityVersionExist, InvalidInput, UnableToAccessDatabase from models.domain.resource import ResourceType from models.domain.user_resource_template import UserResourceTemplate @@ -36,8 +36,8 @@ def create_user_resource_template(template_name: str = "vm-resource-template", p class TestUserResourceTemplatesRequiringAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user yield app.dependency_overrides = {} @@ -106,7 +106,7 @@ async def test_creating_a_user_resource_template_raises_http_422_if_step_ids_are class TestUserResourceTemplatesNotRequiringAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, researcher_user): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = researcher_user + app.dependency_overrides[require_tre_user_or_admin] = researcher_user yield app.dependency_overrides = {} diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_service_templates.py b/api_app/tests_ma/test_api/test_routes/test_workspace_service_templates.py index f60041f111..a9ba696ab7 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_service_templates.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_service_templates.py @@ -5,7 +5,7 @@ from pydantic import parse_obj_as from starlette import status -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from db.errors import EntityDoesNotExist, EntityVersionExist, InvalidInput, UnableToAccessDatabase from models.domain.resource import ResourceType from models.domain.resource_template import ResourceTemplate @@ -58,8 +58,8 @@ def create_user_resource_template(template_name: str = "vm-resource-template", p class TestWorkspaceServiceTemplatesRequiringAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user yield app.dependency_overrides = {} diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_templates.py b/api_app/tests_ma/test_api/test_routes/test_workspace_templates.py index 49178b7999..61a4da4ae3 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_templates.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_templates.py @@ -4,7 +4,7 @@ from pydantic import parse_obj_as from starlette import status -from services.authentication import get_current_admin_user, get_current_tre_user_or_tre_admin +from auth.rbac import require_tre_admin, require_tre_user_or_admin from models.domain.resource import ResourceType from resources import strings @@ -39,8 +39,8 @@ class TestWorkspaceTemplate: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user yield app.dependency_overrides = {} diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py index 64fb2a7d68..66df4d520e 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py @@ -6,10 +6,10 @@ from models.domain.workspace_users import AssignmentType, Role from tests_ma.test_api.test_routes.test_resource_helpers import FAKE_CREATE_TIMESTAMP from tests_ma.test_api.conftest import create_admin_user -from services.authentication import get_current_admin_user, \ - get_current_tre_user_or_tre_admin, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin +from auth.rbac import require_tre_admin, \ + require_tre_user_or_admin, \ + require_workspace_owner_or_researcher_or_airlock_manager, \ + require_workspace_owner_or_researcher_or_airlock_manager from models.domain.workspace import Workspace from resources import strings @@ -47,13 +47,12 @@ def sample_workspace(workspace_id=WORKSPACE_ID, auth_info: dict = {}) -> Workspa class TestWorkspaceUserRoutesWithTreAdmin: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=admin_user()): - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin] = admin_user - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user + yield + app.dependency_overrides = {} @pytest.mark.parametrize("auth_class", ["aad_authentication.AzureADAuthorization"]) @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", return_value=sample_workspace()) diff --git a/api_app/tests_ma/test_api/test_routes/test_workspaces.py b/api_app/tests_ma/test_api/test_routes/test_workspaces.py index f97312f977..c599d9ce01 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspaces.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspaces.py @@ -23,12 +23,12 @@ from models.domain.workspace_service import WorkspaceService from resources import strings from models.schemas.resource_template import ResourceTemplateInformation -from services.authentication import get_current_admin_user, \ - get_current_tre_user_or_tre_admin, get_current_workspace_owner_user, \ - get_current_workspace_owner_or_researcher_user, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager, \ - get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin, \ - get_current_workspace_owner_or_airlock_manager +from auth.rbac import require_tre_admin, \ + require_tre_user_or_admin, require_workspace_owner, \ + require_workspace_owner_or_researcher, \ + require_workspace_owner_or_researcher_or_airlock_manager, \ + require_workspace_owner_or_airlock_manager, \ + require_airlock_manager from azure.cosmos.exceptions import CosmosAccessConditionFailedError @@ -249,8 +249,9 @@ def disabled_user_resource(): class TestWorkspaceRoutesThatDontRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin_user(self, app, non_admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=non_admin_user()): - yield + app.dependency_overrides[require_tre_user_or_admin] = non_admin_user + yield + app.dependency_overrides = {} # [GET] /workspaces @patch("api.routes.workspaces.WorkspaceRepository.get_active_workspaces") @@ -289,12 +290,20 @@ async def test_get_workspaces_returns_correct_data_when_resources_exist(self, _, @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id") @patch("api.routes.workspaces.get_identity_role_assignments") async def test_get_workspace_by_id_get_as_tre_user_returns_403(self, access_service_mock, get_workspace_mock, app, client): + from fastapi import HTTPException auth_info_user_in_workspace_owner_role = {'sp_id': 'ab123', 'client_id': 'cl123', 'app_role_id_workspace_owner': 'ab124', 'app_role_id_workspace_researcher': 'ab125', 'app_role_id_workspace_airlock_manager': 'ab130'} get_workspace_mock.return_value = sample_workspace(auth_info=auth_info_user_in_workspace_owner_role) access_service_mock.return_value = [RoleAssignment('ab123', 'ab124')] - response = await client.get(app.url_path_for(strings.API_GET_WORKSPACE_BY_ID, workspace_id=WORKSPACE_ID)) - assert response.status_code == status.HTTP_403_FORBIDDEN + def forbidden(): + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN) + + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = forbidden + try: + response = await client.get(app.url_path_for(strings.API_GET_WORKSPACE_BY_ID, workspace_id=WORKSPACE_ID)) + assert response.status_code == status.HTTP_403_FORBIDDEN + finally: + app.dependency_overrides.pop(require_workspace_owner_or_researcher_or_airlock_manager, None) # [GET] /workspaces/{workspace_id} @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", side_effect=EntityDoesNotExist) @@ -341,13 +350,15 @@ async def test_get_workspaces_scope_id_returns_empty_if_no_scope_id(self, worksp class TestWorkspaceRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - with patch('services.aad_authentication.AzureADAuthorization._get_user_from_token', return_value=admin_user()): - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin] = admin_user - app.dependency_overrides[get_current_tre_user_or_tre_admin] = admin_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = admin_user - app.dependency_overrides[get_current_admin_user] = admin_user - yield - app.dependency_overrides = {} + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user + app.dependency_overrides[require_workspace_owner_or_researcher] = admin_user + app.dependency_overrides[require_workspace_owner_or_airlock_manager] = admin_user + app.dependency_overrides[require_workspace_owner] = admin_user + app.dependency_overrides[require_airlock_manager] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = admin_user + app.dependency_overrides[require_tre_admin] = admin_user + yield + app.dependency_overrides = {} # [GET] /workspaces @patch("api.routes.workspaces.WorkspaceRepository.get_active_workspaces") @@ -708,10 +719,10 @@ class TestWorkspaceServiceRoutesThatRequireOwnerRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_owner_user(self, app, owner_user): # The following ws services requires the WS app registration - app.dependency_overrides[get_current_workspace_owner_user] = owner_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = owner_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user] = owner_user - app.dependency_overrides[get_current_workspace_owner_or_airlock_manager] = owner_user + app.dependency_overrides[require_workspace_owner] = owner_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = owner_user + app.dependency_overrides[require_workspace_owner_or_researcher] = owner_user + app.dependency_overrides[require_workspace_owner_or_airlock_manager] = owner_user yield app.dependency_overrides = {} @@ -1356,9 +1367,9 @@ class TestWorkspaceServiceRoutesThatRequireOwnerOrResearcherRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_researcher_user(self, app, researcher_user): # The following ws services requires the WS app registration - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager] = researcher_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user] = researcher_user - app.dependency_overrides[get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user yield app.dependency_overrides = {} From 447f2c41c31a32ffa37d79e949d5f56c05ee582c Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 09:49:27 +0000 Subject: [PATCH 05/32] Remove duplicate dependency overrides in test fixtures --- api_app/tests_ma/test_api/test_routes/test_workspace_users.py | 1 - api_app/tests_ma/test_api/test_routes/test_workspaces.py | 1 - 2 files changed, 2 deletions(-) diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py index 66df4d520e..6236a5a38c 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py @@ -49,7 +49,6 @@ class TestWorkspaceUserRoutesWithTreAdmin: def _prepare(self, app, admin_user): app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user app.dependency_overrides[require_tre_user_or_admin] = admin_user - app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user app.dependency_overrides[require_tre_admin] = admin_user yield app.dependency_overrides = {} diff --git a/api_app/tests_ma/test_api/test_routes/test_workspaces.py b/api_app/tests_ma/test_api/test_routes/test_workspaces.py index c599d9ce01..8f87c46f97 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspaces.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspaces.py @@ -1369,7 +1369,6 @@ def log_in_with_researcher_user(self, app, researcher_user): # The following ws services requires the WS app registration app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user app.dependency_overrides[require_workspace_owner_or_researcher] = researcher_user - app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user yield app.dependency_overrides = {} From 52a2c76dc411fc483859af5815388c0ae82a1bb7 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 09:51:10 +0000 Subject: [PATCH 06/32] Fix duplicate import in test_workspace_users.py --- api_app/tests_ma/test_api/test_routes/test_workspace_users.py | 1 - 1 file changed, 1 deletion(-) diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py index 6236a5a38c..f6f1fe8ebd 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py @@ -8,7 +8,6 @@ from tests_ma.test_api.conftest import create_admin_user from auth.rbac import require_tre_admin, \ require_tre_user_or_admin, \ - require_workspace_owner_or_researcher_or_airlock_manager, \ require_workspace_owner_or_researcher_or_airlock_manager from models.domain.workspace import Workspace From 962dfecf062d684b47ffc1f30b6724fdec3b467b Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 12:45:42 +0000 Subject: [PATCH 07/32] Fix linting issues: unused import, extra blank line, unhandled signature error --- api_app/auth/token_validator.py | 2 ++ api_app/services/aad_authentication.py | 1 - api_app/services/authentication.py | 1 - 3 files changed, 2 insertions(+), 2 deletions(-) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index 05b66ed2a1..f7a139238c 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -58,6 +58,8 @@ def validate(self, token: str) -> AuthenticatedUser: ) except jwt.ExpiredSignatureError as exc: raise TokenExpired("Token expired") from exc + except jwt.InvalidSignatureError as exc: + raise TokenSignatureInvalid("Token signature invalid") from exc except jwt.InvalidTokenError as exc: raise TokenInvalid(f"Token invalid: {exc}") from exc diff --git a/api_app/services/aad_authentication.py b/api_app/services/aad_authentication.py index e6d147a571..c893afa516 100644 --- a/api_app/services/aad_authentication.py +++ b/api_app/services/aad_authentication.py @@ -182,7 +182,6 @@ async def _fetch_ws_app_reg_id_from_ws_id(request: Request) -> str: detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, ) from exc - @staticmethod def _get_user_from_token(validated) -> User: """Convert a validated :class:`~auth.models.AuthenticatedUser` to a :class:`User`. diff --git a/api_app/services/authentication.py b/api_app/services/authentication.py index 74efb41d83..9dd2378cbe 100644 --- a/api_app/services/authentication.py +++ b/api_app/services/authentication.py @@ -4,7 +4,6 @@ def extract_auth_information(workspace_creation_properties: dict) -> dict: from fastapi import HTTPException, status - from resources import strings aad_service = get_aad_service() try: return aad_service.extract_workspace_auth_information(workspace_creation_properties) From 72fa930e446cf61bde456afc0902f4cb64456657 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 15:08:39 +0000 Subject: [PATCH 08/32] Add security tests for auth layer; fix missing-oid KeyError in TokenValidator --- api_app/auth/token_validator.py | 19 +++-- api_app/tests_ma/auth/test_rbac.py | 66 ++++++++++++++- api_app/tests_ma/auth/test_token_validator.py | 81 +++++++++++++++++++ 3 files changed, 155 insertions(+), 11 deletions(-) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index f7a139238c..eedd92e3af 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -63,11 +63,14 @@ def validate(self, token: str) -> AuthenticatedUser: except jwt.InvalidTokenError as exc: raise TokenInvalid(f"Token invalid: {exc}") from exc - return AuthenticatedUser( - id=claims["oid"], - name=claims.get("name", ""), - email=claims.get("email") or claims.get("preferred_username"), - roles=claims.get("roles", []), - audience=self._config.audience, - is_workspace_token=self._config.is_workspace_token, - ) + try: + return AuthenticatedUser( + id=claims["oid"], + name=claims.get("name", ""), + email=claims.get("email") or claims.get("preferred_username"), + roles=claims.get("roles", []), + audience=self._config.audience, + is_workspace_token=self._config.is_workspace_token, + ) + except KeyError as exc: + raise TokenInvalid("Token is missing required claim: oid") from exc diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index a77c41f14a..20793508fd 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -145,10 +145,70 @@ async def _run(): asyncio.get_event_loop().run_until_complete(_run()) + def test_wrong_workspace_token_rejected_with_401(self): + """A token issued for workspace A must not grant access to workspace B. -# --------------------------------------------------------------------------- -# AuthenticatedUser helpers -# --------------------------------------------------------------------------- + When the workspace validator raises TokenInvalid (wrong audience) the + code falls back to the core validator. If the token is also invalid + for the core audience the result must be HTTP 401, not a silent pass. + """ + from fastapi import HTTPException + from auth.exceptions import TokenInvalid as _TokenInvalid + + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, _, _ = self._make_fake_deps(["WorkspaceOwner"]) + + # Workspace validator: wrong audience (workspace B rejects a workspace A token) + ws_validator_mock = MagicMock() + ws_validator_mock.validate.side_effect = _TokenInvalid("wrong audience") + + # Core validator: also rejects the token (it's not a core token) + core_validator_mock = MagicMock() + core_validator_mock.validate.side_effect = _TokenInvalid("not a core token") + + import asyncio + + async def _run(): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator_mock): + with patch('auth.rbac.get_core_validator', return_value=core_validator_mock): + with pytest.raises(HTTPException) as exc_info: + await dep(credentials=fake_creds, workspace=fake_workspace) + assert exc_info.value.status_code == 401 + + asyncio.get_event_loop().run_until_complete(_run()) + + def test_non_admin_workspace_user_cannot_elevate_via_core_fallback(self): + """A workspace user (non-admin) whose token is rejected by the workspace validator + must not be able to use the core-validator fallback to gain access. + + Even if the core validator 'accepts' the token, the user must hold at + least one of the required workspace roles (or TREAdmin) to proceed. + """ + from fastapi import HTTPException + from auth.exceptions import TokenInvalid as _TokenInvalid + + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, _, _ = self._make_fake_deps([]) + + # Workspace validator rejects the token (wrong audience) + ws_validator_mock = MagicMock() + ws_validator_mock.validate.side_effect = _TokenInvalid("wrong audience") + + # Core validator accepts the token but the user has NO roles + core_validator_mock = MagicMock() + core_user = _make_user(roles=[]) + core_validator_mock.validate.return_value = core_user + + import asyncio + + async def _run(): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator_mock): + with patch('auth.rbac.get_core_validator', return_value=core_validator_mock): + with pytest.raises(HTTPException) as exc_info: + await dep(credentials=fake_creds, workspace=fake_workspace) + assert exc_info.value.status_code == 403 + + asyncio.get_event_loop().run_until_complete(_run()) class TestAuthenticatedUserHelpers: diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 86a3b83442..1fc692ed18 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -128,6 +128,87 @@ def test_roles_default_to_empty_list(self): assert result.roles == [] + def test_raises_token_signature_invalid_on_invalid_signature(self): + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pytest.importorskip("jwt").InvalidSignatureError("bad sig"), + ): + with pytest.raises(TokenSignatureInvalid): + validator.validate("tampered.jwt.token") + + def test_raises_token_invalid_on_audience_mismatch(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.InvalidAudienceError("wrong audience"), + ): + with pytest.raises(TokenInvalid): + validator.validate("wrong-audience.jwt.token") + + def test_raises_token_invalid_on_issuer_mismatch(self): + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.InvalidIssuerError("wrong issuer"), + ): + with pytest.raises(TokenInvalid): + validator.validate("wrong-issuer.jwt.token") + + def test_raises_token_invalid_on_algorithm_confusion(self): + """Tokens using an unexpected algorithm (e.g. 'none' or HS256) must be rejected.""" + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.InvalidAlgorithmError("algorithm not allowed"), + ): + with pytest.raises(TokenInvalid): + validator.validate("alg-confusion.jwt.token") + + def test_decode_is_called_with_rs256_algorithm_only(self): + """jwt.decode must always be called with algorithms=['RS256'] and no others.""" + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=SAMPLE_CLAIMS) as mock_decode: + validator.validate("valid.jwt.token") + + call_kwargs = mock_decode.call_args + algorithms_arg = call_kwargs[1].get("algorithms") or call_kwargs[0][2] + assert algorithms_arg == ["RS256"], ( + f"Expected algorithms=['RS256'] only, got {algorithms_arg!r}" + ) + + def test_raises_token_invalid_on_missing_oid_claim(self): + """A token whose payload lacks the 'oid' claim must raise TokenInvalid, not KeyError.""" + claims_without_oid = {"name": "User", "email": "u@example.com", "roles": []} + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=claims_without_oid): + with pytest.raises(TokenInvalid, match="oid"): + validator.validate("no-oid.jwt.token") + def test_is_workspace_token_flag_set_from_config(self): signing_key = MagicMock() mock_client = _make_mock_jwks_client(signing_key) From 7a29dafed14e91760ca635ada1a2e38591f728ee Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Mon, 20 Jul 2026 15:10:07 +0000 Subject: [PATCH 09/32] Fix fragile algorithms arg assertion in test_decode_is_called_with_rs256_algorithm_only --- api_app/tests_ma/auth/test_token_validator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 1fc692ed18..4ad1a4e9fe 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -193,7 +193,7 @@ def test_decode_is_called_with_rs256_algorithm_only(self): validator.validate("valid.jwt.token") call_kwargs = mock_decode.call_args - algorithms_arg = call_kwargs[1].get("algorithms") or call_kwargs[0][2] + algorithms_arg = call_kwargs.kwargs.get("algorithms") assert algorithms_arg == ["RS256"], ( f"Expected algorithms=['RS256'] only, got {algorithms_arg!r}" ) From 701fc66e50ad86567cf56fb500ed28f3c72640b0 Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 13:40:02 +0000 Subject: [PATCH 10/32] Address PR review comments: owner-only patch, Field defaults, AsyncMock, role normalization --- api_app/api/routes/workspaces.py | 4 +-- api_app/auth/models.py | 4 +-- api_app/services/aad_authentication.py | 5 ++- api_app/tests_ma/auth/test_rbac.py | 44 ++++++++------------------ api_app/tests_ma/test_api/conftest.py | 2 +- 5 files changed, 22 insertions(+), 37 deletions(-) diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 6a1e21301f..b531275100 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -24,7 +24,7 @@ from models.schemas.resource_template import ResourceTemplateInformationInList from resources import strings from services.aad_authentication import AuthConfigValidationError -from auth.rbac import require_tre_admin, require_workspace_owner, require_workspace_owner_or_researcher, \ +from auth.rbac import require_tre_admin, require_workspace_owner, \ require_tre_user_or_admin, require_workspace_owner_or_researcher_or_airlock_manager, \ require_workspace_owner_or_airlock_manager from services.authentication import get_aad_service, extract_auth_information @@ -278,7 +278,7 @@ async def create_workspace_service(response: Response, workspace_service_input: return OperationInResponse(operation=operation) -@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner_or_researcher), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner), Depends(get_workspace_by_id_from_path)]) async def patch_workspace_service(resource_patch: ResourcePatch, response: Response, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled diff --git a/api_app/auth/models.py b/api_app/auth/models.py index 0a396142d8..056eb4b2e1 100644 --- a/api_app/auth/models.py +++ b/api_app/auth/models.py @@ -1,7 +1,7 @@ from enum import StrEnum from typing import List, Optional, Union -from pydantic import BaseModel +from pydantic import BaseModel, Field class TRERole(StrEnum): @@ -27,7 +27,7 @@ class AuthenticatedUser(BaseModel): id: str name: str email: Optional[str] = None - roles: List[str] = [] + roles: List[str] = Field(default_factory=list) audience: str = "" is_workspace_token: bool = False diff --git a/api_app/services/aad_authentication.py b/api_app/services/aad_authentication.py index c893afa516..cc6ca10d16 100644 --- a/api_app/services/aad_authentication.py +++ b/api_app/services/aad_authentication.py @@ -74,7 +74,10 @@ def __init__(self, auto_error: bool = True, require_one_of_roles: Optional[list] scheme_name="oauth2", auto_error=auto_error, ) - self.require_one_of_roles = require_one_of_roles + # Normalise to an empty list so membership checks in __call__ never + # raise TypeError when the service is instantiated purely for Graph + # API calls (e.g. via get_aad_service()) without role requirements. + self.require_one_of_roles = require_one_of_roles or [] async def __call__(self, request: Request) -> User: token: str = await super().__call__(request) diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index 20793508fd..a7e11be818 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -1,9 +1,6 @@ """Tests for auth.rbac role-checking dependencies.""" import pytest -from unittest.mock import AsyncMock, MagicMock, patch - -from fastapi import FastAPI, Depends -from fastapi.testclient import TestClient +from unittest.mock import MagicMock, patch from auth.models import AuthenticatedUser, TRERole, WorkspaceAccessRole from auth.rbac import require_roles, require_workspace_roles @@ -21,48 +18,33 @@ def _make_user(**kwargs) -> AuthenticatedUser: class TestRequireRoles: - def _app_with_dep(self, dep): - app = FastAPI() - - @app.get("/protected") - async def _route(user=Depends(dep)): - return {"id": user.id} - - return app - def test_allows_user_with_required_role(self): admin = _make_user(roles=["TREAdmin"]) dep = require_roles(TRERole.Admin) - app = self._app_with_dep(dep) - app.dependency_overrides[dep] = lambda: admin # type: ignore[index] - # dependency_overrides must target the inner _check function - # Use the approach below instead + import asyncio + + async def _run(): + result = await dep(user=admin) + assert result.id == "uid" + assert "TREAdmin" in result.roles + + asyncio.get_event_loop().run_until_complete(_run()) def test_raises_403_when_user_lacks_role(self): from fastapi import HTTPException dep = require_roles(TRERole.Admin) - # Extract the inner _check dependency - inner_dep = dep + import asyncio - async def _call_dep(): + async def _run(): user_with_no_roles = _make_user(roles=["TREUser"]) - # Manually invoke the inner check with a user missing the role. - from auth.dependencies import get_authenticated_user - from fastapi import HTTPException - with pytest.raises(HTTPException) as exc_info: - # Simulate what FastAPI would do - from auth.rbac import require_roles as _req - checker = _req(TRERole.Admin) - await checker(user=user_with_no_roles) - + await dep(user=user_with_no_roles) assert exc_info.value.status_code == 403 - import asyncio - asyncio.get_event_loop().run_until_complete(_call_dep()) + asyncio.get_event_loop().run_until_complete(_run()) def test_allows_user_with_any_of_multiple_roles(self): dep = require_roles(TRERole.Admin, TRERole.User) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index 2e0840fa95..f618cd5276 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -27,7 +27,7 @@ def no_auth_token(): fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") - with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', return_value="token"): + with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', new=AsyncMock(return_value="token")): with patch('services.aad_authentication.get_core_validator', return_value=mock_validator): with patch('services.aad_authentication.get_workspace_validator', return_value=mock_validator): with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): From 37cc73178d09169d86b26496a5b94e78cf9735ff Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 13:45:31 +0000 Subject: [PATCH 11/32] Remove unused imports in token validator tests --- api_app/tests_ma/auth/test_token_validator.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 4ad1a4e9fe..7cbc078b1c 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -37,8 +37,6 @@ def _make_mock_jwks_client(signing_key: MagicMock) -> MagicMock: class TestTokenValidatorValidate: def test_returns_authenticated_user_on_valid_token(self): - import jwt as pyjwt - signing_key = MagicMock() mock_client = _make_mock_jwks_client(signing_key) validator = _make_validator(mock_client) @@ -53,8 +51,6 @@ def test_returns_authenticated_user_on_valid_token(self): assert "TREAdmin" in result.roles def test_frozen_user_cannot_be_mutated(self): - import jwt as pyjwt - signing_key = MagicMock() mock_client = _make_mock_jwks_client(signing_key) validator = _make_validator(mock_client) From cb625317d80e254e7e7c5c6cc192bb35d950675e Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 13:56:26 +0000 Subject: [PATCH 12/32] Clean up dead auth code from review - Remove unused get_workspace_authenticated_user (had broken Depends(lambda: None)) - Remove dead AzureADAuthorization.__call__ and helpers; make it a plain Graph service class - Share a single PyJWKClient across token validators via registry --- api_app/auth/dependencies.py | 45 +------ api_app/auth/registry.py | 21 +++- api_app/auth/token_validator.py | 15 ++- api_app/services/aad_authentication.py | 159 +------------------------ api_app/tests_ma/test_api/conftest.py | 12 +- 5 files changed, 44 insertions(+), 208 deletions(-) diff --git a/api_app/auth/dependencies.py b/api_app/auth/dependencies.py index c94c51def2..42c18aae46 100644 --- a/api_app/auth/dependencies.py +++ b/api_app/auth/dependencies.py @@ -1,10 +1,9 @@ from fastapi import Depends, HTTPException, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer -from auth.exceptions import AuthError, TokenExpired, TokenInvalid, TokenSignatureInvalid +from auth.exceptions import AuthError, TokenExpired, TokenSignatureInvalid from auth.models import AuthenticatedUser -from auth.registry import get_core_validator, get_workspace_validator -from models.domain.workspace import Workspace +from auth.registry import get_core_validator from resources import strings from services.logging import logger @@ -40,43 +39,3 @@ async def get_authenticated_user( except AuthError as exc: logger.debug("Core token validation failed: %s", exc) raise _to_http_exception(exc) - - -async def get_workspace_authenticated_user( - credentials: HTTPAuthorizationCredentials = Depends(_bearer), - workspace: Workspace = Depends(lambda: None), -) -> AuthenticatedUser: - """Validate the bearer token for a workspace-scoped request. - - Tries the workspace app registration first. If that audience validation - fails (but not on expiry or signature errors), falls back to the core app - registration so that TREAdmin users can always reach workspace endpoints. - - The *workspace* parameter should be provided via - ``Depends(get_workspace_by_id_from_path)`` when wiring routes. - """ - token = credentials.credentials - - if workspace is not None: - client_id = workspace.properties.get("client_id", "") - if client_id: - try: - return get_workspace_validator(client_id).validate(token) - except (TokenExpired, TokenSignatureInvalid) as exc: - # Hard failures — the token IS for this workspace but is bad. - logger.debug("Workspace token hard failure: %s", exc) - raise _to_http_exception(exc) - except TokenInvalid as exc: - # Wrong audience — fall through to core validator. - logger.debug( - "Workspace token invalid (likely wrong audience), " - "trying core validator: %s", - exc, - ) - - # Fall back to core app registration (allows TREAdmin access). - try: - return get_core_validator().validate(token) - except AuthError as exc: - logger.debug("Core token validation failed: %s", exc) - raise _to_http_exception(exc) diff --git a/api_app/auth/registry.py b/api_app/auth/registry.py index 5764e37871..377adef4f4 100644 --- a/api_app/auth/registry.py +++ b/api_app/auth/registry.py @@ -1,5 +1,7 @@ from functools import lru_cache +from jwt import PyJWKClient + from auth.token_validator import TokenValidator, TokenValidatorConfig from core import config @@ -19,6 +21,17 @@ def _issuer() -> str: ) +@lru_cache(maxsize=1) +def _shared_jwks_client() -> PyJWKClient: + """Single JWKS client shared by all validators. + + Core and every per-workspace token share the same tenant JWKS endpoint, + so a single client (and its key cache) serves all audiences and a single + HTTP fetch is amortised across them. + """ + return PyJWKClient(_jwks_uri(), cache_keys=True, lifespan=300) + + @lru_cache(maxsize=1) def get_core_validator() -> TokenValidator: """Singleton :class:`TokenValidator` for the core TRE app registration.""" @@ -28,7 +41,8 @@ def get_core_validator() -> TokenValidator: audience=config.API_AUDIENCE, issuer=_issuer(), is_workspace_token=False, - ) + ), + jwks_client=_shared_jwks_client(), ) @@ -36,7 +50,7 @@ def get_core_validator() -> TokenValidator: def get_workspace_validator(client_id: str) -> TokenValidator: """Per-workspace :class:`TokenValidator`, cached by *client_id*. - All workspace validators share the same JWKS URI so a single HTTP fetch + All workspace validators share the same JWKS client so a single HTTP fetch serves all audiences; only the audience validation differs. """ return TokenValidator( @@ -45,5 +59,6 @@ def get_workspace_validator(client_id: str) -> TokenValidator: audience=client_id, issuer=_issuer(), is_workspace_token=True, - ) + ), + jwks_client=_shared_jwks_client(), ) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index eedd92e3af..3917ffc58f 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -1,4 +1,5 @@ from dataclasses import dataclass +from typing import Optional import jwt from jwt import PyJWKClient @@ -21,11 +22,21 @@ class TokenValidator: Uses :class:`jwt.PyJWKClient` which handles JWKS caching and key rotation automatically — keys removed from the JWKS endpoint are evicted from the cache, preventing unbounded growth. + + A ``jwks_client`` may be supplied so that multiple validators sharing the + same JWKS endpoint (e.g. all per-workspace validators) reuse a single + client and its key cache, avoiding redundant HTTP fetches. """ - def __init__(self, config: TokenValidatorConfig) -> None: + def __init__( + self, + config: TokenValidatorConfig, + jwks_client: Optional[PyJWKClient] = None, + ) -> None: self._config = config - self._jwks_client = PyJWKClient(config.jwks_uri, cache_keys=True, lifespan=300) + self._jwks_client = jwks_client or PyJWKClient( + config.jwks_uri, cache_keys=True, lifespan=300 + ) def validate(self, token: str) -> AuthenticatedUser: """Validate *token* and return an :class:`AuthenticatedUser`. diff --git a/api_app/services/aad_authentication.py b/api_app/services/aad_authentication.py index cc6ca10d16..6fcf73eaea 100644 --- a/api_app/services/aad_authentication.py +++ b/api_app/services/aad_authentication.py @@ -1,22 +1,16 @@ from collections import defaultdict from enum import Enum -from typing import List, Optional +from typing import List import requests -from fastapi import HTTPException, Request, status -from fastapi.security import OAuth2AuthorizationCodeBearer from msal import ConfidentialClientApplication from semantic_version import Version -from auth.exceptions import AuthError, TokenExpired, TokenInvalid, TokenSignatureInvalid -from auth.registry import get_core_validator, get_workspace_validator from core import config -from db.errors import EntityDoesNotExist from models.domain.authentication import User, RoleAssignment from models.domain.workspace import Workspace, WorkspaceRole from models.domain.workspace_users import AssignableUser, AssignedUser, AssignmentType, Role from resources import strings -from db.repositories.workspaces import WorkspaceRepository from services.logging import logger @@ -33,167 +27,26 @@ class UserRoleAssignmentError(Exception): """Raised when a user role assignment fails.""" -def _authenticated_user_to_user(validated) -> User: - """Convert an :class:`~auth.models.AuthenticatedUser` to the legacy :class:`~models.domain.authentication.User`.""" - return User( - id=validated.id, - name=validated.name, - email=validated.email or "", - roles=list(validated.roles), - ) - - class PrincipalType(Enum): User = "User" Group = "Group" ServicePrincipal = "ServicePrincipal" -class AzureADAuthorization(OAuth2AuthorizationCodeBearer): - """FastAPI security dependency that validates Entra ID JWTs. +class AzureADAuthorization: + """Service wrapper for Microsoft Graph calls related to workspace auth. - Uses :mod:`auth.token_validator` (backed by :class:`jwt.PyJWKClient`) for - JWT validation so key management is handled automatically. + Handles workspace app-registration validation and role-assignment lookups + via the Microsoft Graph API. JWT validation for incoming API requests is + handled separately by the :mod:`auth` package. """ - require_one_of_roles: Optional[list] = None - aad_instance: str = config.AAD_AUTHORITY_URL - - TRE_CORE_ROLES = ['TREAdmin', 'TREUser', 'TREAirlockAutomation'] WORKSPACE_ROLES_DICT = { 'WorkspaceOwner': 'app_role_id_workspace_owner', 'WorkspaceResearcher': 'app_role_id_workspace_researcher', 'AirlockManager': 'app_role_id_workspace_airlock_manager', } - def __init__(self, auto_error: bool = True, require_one_of_roles: Optional[list] = None): - super().__init__( - authorizationUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/authorize", - tokenUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/token", - refreshUrl=f"{self.aad_instance}/{config.AAD_TENANT_ID}/oauth2/v2.0/token", - scheme_name="oauth2", - auto_error=auto_error, - ) - # Normalise to an empty list so membership checks in __call__ never - # raise TypeError when the service is instantiated purely for Graph - # API calls (e.g. via get_aad_service()) without role requirements. - self.require_one_of_roles = require_one_of_roles or [] - - async def __call__(self, request: Request) -> User: - token: str = await super().__call__(request) - - decoded_user = None - - # Try workspace app registration first when a workspace_id is present - # and the route requires workspace-scoped roles. - if 'workspace_id' in request.path_params and any( - role in self.require_one_of_roles for role in self.WORKSPACE_ROLES_DICT - ): - logger.debug("Workspace ID present — attempting workspace app registration") - try: - app_reg_id = await self._fetch_ws_app_reg_id_from_ws_id(request) - if app_reg_id: - try: - validated = get_workspace_validator(app_reg_id).validate(token) - decoded_user = self._get_user_from_token(validated) - except (TokenExpired, TokenSignatureInvalid) as exc: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.EXPIRED_SIGNATURE - if isinstance(exc, TokenExpired) - else strings.INVALID_SIGNATURE, - ) - except TokenInvalid: - logger.debug( - "Workspace token invalid, will try core app registration" - ) - except HTTPException: - raise - except Exception as exc: - logger.debug("Failed to resolve workspace app registration: %s", exc) - - # Try core app registration for TRE core roles. - if decoded_user is None and any( - role in self.require_one_of_roles for role in self.TRE_CORE_ROLES - ): - try: - validated = get_core_validator().validate(token) - decoded_user = self._get_user_from_token(validated) - except TokenExpired: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.EXPIRED_SIGNATURE, - ) - except TokenSignatureInvalid: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.INVALID_SIGNATURE, - ) - except TokenInvalid: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.INVALID_TOKEN, - ) - except AuthError as exc: - logger.debug("Core token validation failed: %s", exc) - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.AUTH_UNABLE_TO_VALIDATE_TOKEN, - ) - - if decoded_user is None: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.AUTH_UNABLE_TO_VALIDATE_TOKEN, - ) - - if not any(role in self.require_one_of_roles for role in decoded_user.roles): - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {self.require_one_of_roles}", - headers={"WWW-Authenticate": "Bearer"}, - ) - - return decoded_user - - @staticmethod - async def _fetch_ws_app_reg_id_from_ws_id(request: Request) -> str: - workspace_id = request.path_params.get('workspace_id') - if not workspace_id: - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, - ) - try: - ws_repo = await WorkspaceRepository.create() - workspace = await ws_repo.get_workspace_by_id(workspace_id) - return workspace.properties.get('client_id', '') - except EntityDoesNotExist: - raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail=strings.WORKSPACE_DOES_NOT_EXIST, - ) - except HTTPException: - raise - except Exception as exc: - logger.exception( - "Failed to get workspace app registration ID for workspace %s", - workspace_id, - ) - raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, - ) from exc - - @staticmethod - def _get_user_from_token(validated) -> User: - """Convert a validated :class:`~auth.models.AuthenticatedUser` to a :class:`User`. - - This method is kept as an instance-patchable hook so tests can inject - specific users without needing real JWTs. - """ - return _authenticated_user_to_user(validated) - @staticmethod def _get_msgraph_token() -> str: scopes = [f"{MICROSOFT_GRAPH_URL}/.default"] diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index f618cd5276..b2f2a0ca42 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -28,13 +28,11 @@ def no_auth_token(): fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', new=AsyncMock(return_value="token")): - with patch('services.aad_authentication.get_core_validator', return_value=mock_validator): - with patch('services.aad_authentication.get_workspace_validator', return_value=mock_validator): - with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): - with patch('auth.dependencies.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): - yield + with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): + with patch('auth.dependencies.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") From 5fa9cd4d2d911f4c050c64315fce807ffa970452 Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 14:17:28 +0000 Subject: [PATCH 13/32] Harden auth model immutability and resolve E231 lint - Store AuthenticatedUser.roles as an immutable tuple so roles cannot be escalated via in-place mutation (frozen only blocked reassignment) - Ignore E231 in flake8 config (Python 3.12 f-string tokenisation false positives) - Fix E306 nested-def blank line in test_migrations.py --- .flake8 | 4 +++- .github/linters/.flake8 | 4 +++- api_app/auth/models.py | 9 +++++---- api_app/tests_ma/auth/test_rbac.py | 6 ++++++ api_app/tests_ma/auth/test_token_validator.py | 7 ++++++- api_app/tests_ma/test_api/test_routes/test_migrations.py | 1 + 6 files changed, 24 insertions(+), 7 deletions(-) diff --git a/.flake8 b/.flake8 index a6f727b8d5..b4d0c76881 100644 --- a/.flake8 +++ b/.flake8 @@ -1,2 +1,4 @@ [flake8] -ignore = E501, W503 +# E231 is ignored because pycodestyle on Python 3.12+ mis-tokenises f-string +# contents (commas/colons inside URL, OData and JSON string literals) as code. +ignore = E501, W503, E231 diff --git a/.github/linters/.flake8 b/.github/linters/.flake8 index a1cce8e7d0..ce0ed63fac 100644 --- a/.github/linters/.flake8 +++ b/.github/linters/.flake8 @@ -1,2 +1,4 @@ [flake8] -ignore = E501,W503 +# E231 is ignored because pycodestyle on Python 3.12+ mis-tokenises f-string +# contents (commas/colons inside URL, OData and JSON string literals) as code. +ignore = E501,W503,E231 diff --git a/api_app/auth/models.py b/api_app/auth/models.py index 056eb4b2e1..209c3126be 100644 --- a/api_app/auth/models.py +++ b/api_app/auth/models.py @@ -1,5 +1,5 @@ from enum import StrEnum -from typing import List, Optional, Union +from typing import Optional, Tuple, Union from pydantic import BaseModel, Field @@ -20,14 +20,15 @@ class AuthenticatedUser(BaseModel): """Immutable, validated user derived from a JWT. Fields map directly to standard JWT claims; ``id`` holds the ``oid`` - claim (the stable object identifier in Entra ID). The model is frozen so - roles cannot be escalated after creation. + claim (the stable object identifier in Entra ID). The model is frozen and + ``roles`` is stored as a tuple so roles cannot be reassigned *or* mutated + in place (e.g. ``roles.append(...)``) after creation. """ id: str name: str email: Optional[str] = None - roles: List[str] = Field(default_factory=list) + roles: Tuple[str, ...] = Field(default_factory=tuple) audience: str = "" is_workspace_token: bool = False diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index a7e11be818..5112f9d770 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -214,3 +214,9 @@ def test_model_is_frozen(self): user = _make_user(roles=["TREAdmin"]) with pytest.raises(TypeError): user.roles = [] # type: ignore[misc] + + def test_roles_cannot_be_mutated_in_place(self): + user = _make_user(roles=["TREUser"]) + assert isinstance(user.roles, tuple) + with pytest.raises(AttributeError): + user.roles.append("TREAdmin") # type: ignore[attr-defined] diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 7cbc078b1c..1e0f0ced24 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -61,6 +61,11 @@ def test_frozen_user_cannot_be_mutated(self): with pytest.raises(TypeError): user.roles = [] # type: ignore[misc] + # roles is a tuple, so in-place escalation is impossible too + assert isinstance(user.roles, tuple) + with pytest.raises(AttributeError): + user.roles.append("TREAdmin") # type: ignore[attr-defined] + def test_raises_token_expired_on_expired_signature(self): import jwt as pyjwt @@ -122,7 +127,7 @@ def test_roles_default_to_empty_list(self): with patch("auth.token_validator.jwt.decode", return_value=claims_no_roles): result = validator.validate("token") - assert result.roles == [] + assert result.roles == () def test_raises_token_signature_invalid_on_invalid_signature(self): signing_key = MagicMock() diff --git a/api_app/tests_ma/test_api/test_routes/test_migrations.py b/api_app/tests_ma/test_api/test_routes/test_migrations.py index 73264574ae..dff2c94d9f 100644 --- a/api_app/tests_ma/test_api/test_routes/test_migrations.py +++ b/api_app/tests_ma/test_api/test_routes/test_migrations.py @@ -13,6 +13,7 @@ class TestMigrationRoutesWithNonAdminRights: @pytest.fixture(autouse=True, scope='class') def log_in_with_non_admin_user(self, app, non_admin_user): from fastapi import HTTPException + def forbidden(): raise HTTPException(status_code=403) app.dependency_overrides[require_tre_admin] = forbidden From c9e7fcfe9ae5b2d62c59bc8841d9f64a45fbb039 Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 15:09:24 +0000 Subject: [PATCH 14/32] Address new PR review comments - Add PR reference to CHANGELOG auth entry - Add WWW-Authenticate: Bearer header to 401 responses in auth dependencies - Guard get_required_roles against endpoint.__defaults__ being None --- CHANGELOG.md | 2 +- api_app/auth/dependencies.py | 17 +++++++---------- api_app/tests_ma/test_api/conftest.py | 5 +++-- 3 files changed, 11 insertions(+), 13 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5f081aa80d..40ba3dd987 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,7 +9,7 @@ ENHANCEMENTS: * Add Windows Server 2025 image support to Guacamole. ([#4890](https://github.com/microsoft/AzureTRE/issues/4890)) * Add support for setting resource processor VMSS SKU via environment variables ([#4936](https://github.com/microsoft/AzureTRE/issues/4936)) * Exclude recovery service vaults from e2e tests ([#4920](https://github.com/microsoft/AzureTRE/issues/4920)) -* Strengthen TRE API authentication: introduce layered `auth/` package with typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, and composable RBAC factories; remove the `AccessService` abstraction that is no longer needed now that Entra ID is the only auth provider. +* Strengthen TRE API authentication: introduce layered `auth/` package with typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, and composable RBAC factories; remove the `AccessService` abstraction that is no longer needed now that Entra ID is the only auth provider. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Update API, CLI, and UI dependencies to address high-severity Dependabot alerts, including `PyJWT`, `Vite`, `lodash`, `fast-uri`, `flatted`, `immutable`, and `minimatch`. * Update dependencies to address Dependabot security alerts: `aiohttp` to 3.14.1, `Pygments` to 2.20.0, `esbuild`, `ws`, `js-yaml`, `@babel/core`, `flatted` (via vitest upgrade), and `react-router-dom`. ([#4950](https://github.com/microsoft/AzureTRE/issues/4950)) * Added support for formatting UI code via `pre-commit` and fixed existing formatting issues. ([#4955](https://github.com/microsoft/AzureTRE/issues/4955)) diff --git a/api_app/auth/dependencies.py b/api_app/auth/dependencies.py index 42c18aae46..e12b463e6e 100644 --- a/api_app/auth/dependencies.py +++ b/api_app/auth/dependencies.py @@ -12,18 +12,15 @@ def _to_http_exception(exc: AuthError) -> HTTPException: if isinstance(exc, TokenExpired): - return HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.EXPIRED_SIGNATURE, - ) - if isinstance(exc, TokenSignatureInvalid): - return HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.INVALID_SIGNATURE, - ) + detail = strings.EXPIRED_SIGNATURE + elif isinstance(exc, TokenSignatureInvalid): + detail = strings.INVALID_SIGNATURE + else: + detail = strings.INVALID_TOKEN return HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, - detail=strings.INVALID_TOKEN, + detail=detail, + headers={"WWW-Authenticate": "Bearer"}, ) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index b2f2a0ca42..f2223856b7 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -91,11 +91,12 @@ def override_get_user(): def get_required_roles(endpoint): - dependencies = list(filter(lambda x: hasattr(x.dependency, 'require_one_of_roles'), endpoint.__defaults__)) + defaults = endpoint.__defaults__ or () + dependencies = list(filter(lambda x: hasattr(x.dependency, 'require_one_of_roles'), defaults)) if dependencies: return dependencies[0].dependency.require_one_of_roles # New-style deps: check for _role_names attribute on the closure - dependencies = list(filter(lambda x: hasattr(x.dependency, '_role_names'), endpoint.__defaults__)) + dependencies = list(filter(lambda x: hasattr(x.dependency, '_role_names'), defaults)) if dependencies: return dependencies[0].dependency._role_names return [] From ed62d1b2e263154447a223997ec184db9141080c Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 15:15:43 +0000 Subject: [PATCH 15/32] Bump api version to 0.26.0 for auth refactor --- api_app/_version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/api_app/_version.py b/api_app/_version.py index 605b3cd20e..7c4a9591e1 100644 --- a/api_app/_version.py +++ b/api_app/_version.py @@ -1 +1 @@ -__version__ = "0.25.29" +__version__ = "0.26.0" From a2ade5515b50a23c4ad3ae85714e95632c528b1e Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 15:33:57 +0000 Subject: [PATCH 16/32] Address new review comments: harden workspace fallback + 401 on missing creds - require_workspace_roles core fallback now only grants access to TREAdmin; any other valid core token is rejected with 401 (no cross-audience elevation) - HTTPBearer(auto_error=False) + require_bearer_credentials returns 401 + WWW-Authenticate for missing/malformed Authorization headers (was FastAPI 403) - Update workspace-roles tests to exercise the workspace-validator path and the hardened fallback; add missing-credentials tests --- api_app/auth/dependencies.py | 22 ++++++++- api_app/auth/rbac.py | 19 ++++---- api_app/tests_ma/auth/test_rbac.py | 77 +++++++++++++++++++++++------- 3 files changed, 90 insertions(+), 28 deletions(-) diff --git a/api_app/auth/dependencies.py b/api_app/auth/dependencies.py index e12b463e6e..c34f7e9665 100644 --- a/api_app/auth/dependencies.py +++ b/api_app/auth/dependencies.py @@ -1,3 +1,5 @@ +from typing import Optional + from fastapi import Depends, HTTPException, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer @@ -7,7 +9,10 @@ from resources import strings from services.logging import logger -_bearer = HTTPBearer(auto_error=True) +# auto_error=False so a missing/malformed Authorization header is mapped to a +# consistent 401 + WWW-Authenticate response by require_bearer_credentials +# (FastAPI's built-in auto_error would raise a 403 "Not authenticated"). +_bearer = HTTPBearer(auto_error=False) def _to_http_exception(exc: AuthError) -> HTTPException: @@ -24,8 +29,21 @@ def _to_http_exception(exc: AuthError) -> HTTPException: ) +async def require_bearer_credentials( + credentials: Optional[HTTPAuthorizationCredentials] = Depends(_bearer), +) -> HTTPAuthorizationCredentials: + """Return the bearer credentials, or raise 401 if the header is missing/malformed.""" + if credentials is None: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.AUTH_COULD_NOT_VALIDATE_CREDENTIALS, + headers={"WWW-Authenticate": "Bearer"}, + ) + return credentials + + async def get_authenticated_user( - credentials: HTTPAuthorizationCredentials = Depends(_bearer), + credentials: HTTPAuthorizationCredentials = Depends(require_bearer_credentials), ) -> AuthenticatedUser: """Validate the bearer token against the core TRE app registration. diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py index d07ffd200d..f998bb4534 100644 --- a/api_app/auth/rbac.py +++ b/api_app/auth/rbac.py @@ -3,7 +3,7 @@ from fastapi import Depends, HTTPException, status from fastapi.security import HTTPAuthorizationCredentials -from auth.dependencies import _bearer, _to_http_exception, get_authenticated_user +from auth.dependencies import require_bearer_credentials, _to_http_exception, get_authenticated_user from auth.exceptions import AuthError, TokenExpired, TokenSignatureInvalid, TokenInvalid from auth.models import AuthenticatedUser, TRERole, WorkspaceAccessRole from auth.registry import get_core_validator, get_workspace_validator @@ -45,8 +45,9 @@ def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Call Validates the bearer token against the workspace app registration first. If that fails with a wrong-audience error, falls back to the core app - registration so that TREAdmin users can always reach workspace endpoints. - TREAdmin is always allowed regardless of workspace role. + registration **only** so that TREAdmin users can always reach workspace + endpoints. A valid core token that is not TREAdmin is rejected (401) — a + non-admin core token must never satisfy a workspace-scoped check. The workspace is resolved from the URL path (``workspace_id`` path parameter) so this factory should only be used on routes whose path @@ -58,7 +59,7 @@ def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Call role_names = [r.value for r in roles] async def _check( - credentials: HTTPAuthorizationCredentials = Depends(_bearer), + credentials: HTTPAuthorizationCredentials = Depends(require_bearer_credentials), workspace: Workspace = Depends(get_workspace_by_id_from_path), ) -> AuthenticatedUser: token = credentials.credentials @@ -84,16 +85,18 @@ async def _check( "Workspace token invalid (likely wrong audience), trying core validator" ) - # Fall back to core app registration (allows TREAdmin access). + # Fall back to core app registration. A core token is only accepted + # here for TREAdmin (cross-audience access to any workspace); any other + # valid core token is treated as invalid for this workspace resource. try: user = get_core_validator().validate(token) except AuthError as exc: raise _to_http_exception(exc) - if not (set(user.roles) & allowed_values): + if not user.is_tre_admin(): raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail=f"{strings.ACCESS_USER_DOES_NOT_HAVE_REQUIRED_ROLE}: {role_names}", + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_TOKEN, headers={"WWW-Authenticate": "Bearer"}, ) return user diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index 5112f9d770..c2555a3037 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -65,20 +65,20 @@ async def _run(): class TestRequireWorkspaceRoles: - def _make_fake_deps(self, user_roles): - """Return (fake_credentials, fake_workspace, mock_validator) for testing _check directly.""" + def _make_fake_deps(self, user_roles, with_client_id=True): + """Return (fake_credentials, fake_workspace, mock_validator, user) for testing _check directly.""" from fastapi.security import HTTPAuthorizationCredentials from models.domain.workspace import Workspace fake_creds = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") - # Workspace with no client_id → falls straight through to core validator + properties = {"client_id": "ws-client-id"} if with_client_id else {} fake_workspace = Workspace( id="ws-id", templateName="test", templateVersion="0.1.0", etag="", resourcePath="/workspaces/ws-id", - properties={}, + properties=properties, ) validated_user = _make_user(roles=user_roles) mock_validator = MagicMock() @@ -86,26 +86,36 @@ def _make_fake_deps(self, user_roles): return fake_creds, fake_workspace, mock_validator, validated_user def test_admin_always_passes_without_workspace_role(self): + """A TREAdmin using a core token reaches a workspace endpoint via the fallback.""" + from auth.exceptions import TokenInvalid as _TokenInvalid + dep = require_workspace_roles(WorkspaceAccessRole.Owner) - fake_creds, fake_workspace, mock_validator, admin = self._make_fake_deps(["TREAdmin"]) + fake_creds, fake_workspace, _, admin = self._make_fake_deps(["TREAdmin"]) + + # Workspace validator rejects the core token (wrong audience); core validator accepts it. + ws_validator = MagicMock() + ws_validator.validate.side_effect = _TokenInvalid("wrong audience") + core_validator = MagicMock() + core_validator.validate.return_value = admin import asyncio async def _run(): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): - result = await dep(credentials=fake_creds, workspace=fake_workspace) + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator): + with patch('auth.rbac.get_core_validator', return_value=core_validator): + result = await dep(credentials=fake_creds, workspace=fake_workspace) assert result.id == "uid" asyncio.get_event_loop().run_until_complete(_run()) def test_workspace_owner_passes(self): dep = require_workspace_roles(WorkspaceAccessRole.Owner) - fake_creds, fake_workspace, mock_validator, owner = self._make_fake_deps(["WorkspaceOwner"]) + fake_creds, fake_workspace, ws_validator, owner = self._make_fake_deps(["WorkspaceOwner"]) import asyncio async def _run(): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator): result = await dep(credentials=fake_creds, workspace=fake_workspace) assert result.id == "uid" @@ -115,12 +125,12 @@ def test_raises_403_for_user_without_workspace_role(self): from fastapi import HTTPException dep = require_workspace_roles(WorkspaceAccessRole.Owner) - fake_creds, fake_workspace, mock_validator, researcher = self._make_fake_deps(["WorkspaceResearcher"]) + fake_creds, fake_workspace, ws_validator, researcher = self._make_fake_deps(["WorkspaceResearcher"]) import asyncio async def _run(): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator): with pytest.raises(HTTPException) as exc_info: await dep(credentials=fake_creds, workspace=fake_workspace) assert exc_info.value.status_code == 403 @@ -160,11 +170,11 @@ async def _run(): asyncio.get_event_loop().run_until_complete(_run()) def test_non_admin_workspace_user_cannot_elevate_via_core_fallback(self): - """A workspace user (non-admin) whose token is rejected by the workspace validator - must not be able to use the core-validator fallback to gain access. + """A non-admin core token must never satisfy a workspace-scoped check. - Even if the core validator 'accepts' the token, the user must hold at - least one of the required workspace roles (or TREAdmin) to proceed. + Even if the core validator accepts the token *and* its claims contain a + workspace role, the fallback path only grants access to TREAdmin; any + other core token is rejected with 401 (wrong audience for this resource). """ from fastapi import HTTPException from auth.exceptions import TokenInvalid as _TokenInvalid @@ -176,9 +186,10 @@ def test_non_admin_workspace_user_cannot_elevate_via_core_fallback(self): ws_validator_mock = MagicMock() ws_validator_mock.validate.side_effect = _TokenInvalid("wrong audience") - # Core validator accepts the token but the user has NO roles + # Core validator accepts the token and it even carries a workspace role, + # but the user is NOT TREAdmin. core_validator_mock = MagicMock() - core_user = _make_user(roles=[]) + core_user = _make_user(roles=["WorkspaceOwner"]) core_validator_mock.validate.return_value = core_user import asyncio @@ -188,7 +199,37 @@ async def _run(): with patch('auth.rbac.get_core_validator', return_value=core_validator_mock): with pytest.raises(HTTPException) as exc_info: await dep(credentials=fake_creds, workspace=fake_workspace) - assert exc_info.value.status_code == 403 + assert exc_info.value.status_code == 401 + + asyncio.get_event_loop().run_until_complete(_run()) + + +class TestRequireBearerCredentials: + def test_missing_credentials_raises_401_with_www_authenticate(self): + from fastapi import HTTPException + from auth.dependencies import require_bearer_credentials + + import asyncio + + async def _run(): + with pytest.raises(HTTPException) as exc_info: + await require_bearer_credentials(credentials=None) + assert exc_info.value.status_code == 401 + assert exc_info.value.headers.get("WWW-Authenticate") == "Bearer" + + asyncio.get_event_loop().run_until_complete(_run()) + + def test_present_credentials_are_returned(self): + from fastapi.security import HTTPAuthorizationCredentials + from auth.dependencies import require_bearer_credentials + + creds = HTTPAuthorizationCredentials(scheme="Bearer", credentials="tok") + + import asyncio + + async def _run(): + result = await require_bearer_credentials(credentials=creds) + assert result is creds asyncio.get_event_loop().run_until_complete(_run()) From 298bc0b45c397dc02b7edc6f9763fa6bd6aa683d Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 15:56:14 +0000 Subject: [PATCH 17/32] Restore pre-PR TREAdmin authorization semantics on workspace endpoints The migration inadvertently granted TREAdmin access to ALL workspace-scoped endpoints (require_workspace_roles always allowed TREAdmin). Pre-PR, only endpoints explicitly using the *_or_tre_admin dependencies allowed TREAdmin; owner/researcher/airlock-manager-only endpoints did not. - Make TREAdmin access opt-in via require_workspace_roles(..., allow_tre_admin) - When allow_tre_admin=False, no core-token fallback occurs (wrong-audience token -> 401), preserving separation of platform admin vs workspace access - Add require_workspace_owner_or_tre_admin and require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin, and re-map the exact endpoints that used *_or_tre_admin in main (costs, workspace_users router, workspaces shared router + operations/history + template listing) - Update route tests' dependency overrides and rbac tests accordingly --- api_app/api/routes/costs.py | 6 +- api_app/api/routes/workspace_users.py | 4 +- api_app/api/routes/workspaces.py | 15 ++--- api_app/auth/rbac.py | 58 +++++++++++++++---- api_app/tests_ma/auth/test_rbac.py | 33 +++++++++-- .../test_routes/test_workspace_users.py | 4 +- .../test_api/test_routes/test_workspaces.py | 13 ++++- 7 files changed, 101 insertions(+), 32 deletions(-) diff --git a/api_app/api/routes/costs.py b/api_app/api/routes/costs.py index d76264ab5f..a248fa846a 100644 --- a/api_app/api/routes/costs.py +++ b/api_app/api/routes/costs.py @@ -15,13 +15,13 @@ from db.repositories.workspaces import WorkspaceRepository from models.domain.costs import CostReport, GranularityEnum, WorkspaceCostReport from resources import strings -from auth.rbac import require_tre_admin, require_workspace_owner +from auth.rbac import require_tre_admin, require_workspace_owner_or_tre_admin from services.cost_service import CostService, ServiceUnavailable, SubscriptionNotSupported, TooManyRequests, WorkspaceDoesNotExist, cost_service_factory from services.logging import logger costs_core_router = APIRouter(dependencies=[Depends(require_tre_admin)]) -costs_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner)]) +costs_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_tre_admin)]) def validate_report_period(from_date: Optional[datetime], to_date: Optional[datetime]): @@ -86,7 +86,7 @@ async def costs( @costs_workspace_router.get("/workspaces/{workspace_id}/costs", response_model=WorkspaceCostReport, name=strings.API_GET_WORKSPACE_COSTS, - dependencies=[Depends(require_workspace_owner)], + dependencies=[Depends(require_workspace_owner_or_tre_admin)], responses=get_workspace_cost_report_responses()) async def workspace_costs(workspace_id: UUID4, params: CostsQueryParams = Depends(), cost_service: CostService = Depends(cost_service_factory), diff --git a/api_app/api/routes/workspace_users.py b/api_app/api/routes/workspace_users.py index 8f7469e233..896d5f6512 100644 --- a/api_app/api/routes/workspace_users.py +++ b/api_app/api/routes/workspace_users.py @@ -5,10 +5,10 @@ from services.authentication import get_aad_service from models.schemas.users import UsersInResponse, AssignableUsersInResponse, WorkspaceUserOperationResponse from models.schemas.roles import RolesInResponse -from auth.rbac import require_tre_admin, require_workspace_owner_or_researcher_or_airlock_manager +from auth.rbac import require_tre_admin, require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin workspaces_users_admin_router = APIRouter(dependencies=[Depends(require_tre_admin)]) -workspaces_users_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) +workspaces_users_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin)]) @workspaces_users_shared_router.get("/workspaces/{workspace_id}/users", response_model=UsersInResponse, name=strings.API_GET_WORKSPACE_USERS) diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 11674fdb22..04ea2d6513 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -26,7 +26,8 @@ from services.aad_authentication import AuthConfigValidationError from auth.rbac import require_tre_admin, require_workspace_owner, \ require_tre_user_or_admin, require_workspace_owner_or_researcher_or_airlock_manager, \ - require_workspace_owner_or_airlock_manager + require_workspace_owner_or_airlock_manager, require_workspace_owner_or_tre_admin, \ + require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin from services.authentication import get_aad_service, extract_auth_information from services.azure_resource_status import get_azure_resource_status from azure.cosmos.exceptions import CosmosAccessConditionFailedError @@ -37,7 +38,7 @@ from services.logging import logger workspaces_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -workspaces_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) +workspaces_shared_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin)]) workspace_services_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) user_resources_workspace_router = APIRouter(dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) @@ -195,7 +196,7 @@ async def invoke_action_on_workspace(response: Response, action: str, user=Depen async def get_workspace_service_templates( workspace=Depends(get_workspace_by_id_from_path), template_repo=Depends(get_repository(ResourceTemplateRepository)), - user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> ResourceTemplateInformationInList: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin)) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.WorkspaceService, user.roles) return ResourceTemplateInformationInList(templates=template_infos) @@ -206,22 +207,22 @@ async def get_user_resource_templates( service_template_name: str, workspace=Depends(get_workspace_by_id_from_path), template_repo=Depends(get_repository(ResourceTemplateRepository)), - user=Depends(require_workspace_owner_or_researcher_or_airlock_manager)) -> ResourceTemplateInformationInList: + user=Depends(require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin)) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.UserResource, user.roles, service_template_name) return ResourceTemplateInformationInList(templates=template_infos) -@workspaces_shared_router.get("/workspaces/{workspace_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(require_workspace_owner)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/operations", response_model=OperationInList, name=strings.API_GET_RESOURCE_OPERATIONS, dependencies=[Depends(require_workspace_owner_or_tre_admin)]) async def retrieve_workspace_operations_by_workspace_id(workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository))) -> OperationInList: return OperationInList(operations=await operations_repo.get_operations_by_resource_id(resource_id=workspace.id)) -@workspaces_shared_router.get("/workspaces/{workspace_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(require_workspace_owner)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/operations/{operation_id}", response_model=OperationInResponse, name=strings.API_GET_RESOURCE_OPERATION_BY_ID, dependencies=[Depends(require_workspace_owner_or_tre_admin)]) async def retrieve_workspace_operation_by_workspace_id_and_operation_id(workspace=Depends(get_workspace_by_id_from_path), operation=Depends(get_operation_by_id_from_path)) -> OperationInList: return OperationInResponse(operation=operation) -@workspaces_shared_router.get("/workspaces/{workspace_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(require_workspace_owner)]) +@workspaces_shared_router.get("/workspaces/{workspace_id}/history", response_model=ResourceHistoryInList, name=strings.API_GET_RESOURCE_HISTORY, dependencies=[Depends(require_workspace_owner_or_tre_admin)]) async def retrieve_workspace_history_by_workspace_id(workspace=Depends(get_workspace_by_id_from_path), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> ResourceHistoryInList: return ResourceHistoryInList(resource_history=await resource_history_repo.get_resource_history_by_resource_id(resource_id=workspace.id)) diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py index f998bb4534..5329d66ba9 100644 --- a/api_app/auth/rbac.py +++ b/api_app/auth/rbac.py @@ -40,22 +40,31 @@ async def _check( return _check -def require_workspace_roles(*roles: Union[TRERole, WorkspaceAccessRole]) -> Callable: +def require_workspace_roles( + *roles: Union[TRERole, WorkspaceAccessRole], + allow_tre_admin: bool = False, +) -> Callable: """Factory that returns a dependency enforcing workspace-scoped *roles*. - Validates the bearer token against the workspace app registration first. - If that fails with a wrong-audience error, falls back to the core app - registration **only** so that TREAdmin users can always reach workspace - endpoints. A valid core token that is not TREAdmin is rejected (401) — a - non-admin core token must never satisfy a workspace-scoped check. + Validates the bearer token against the workspace app registration first + (audience-aware). A token carrying one of the required workspace *roles* is + accepted. + + ``allow_tre_admin`` controls whether a TREAdmin — who authenticates with a + *core* token (wrong audience for the workspace app registration) — may also + reach the endpoint. When ``True``, a wrong-audience token falls back to the + core app registration and is accepted **only** if it is a valid TREAdmin + token. When ``False`` (the default), no core fallback occurs and any token + that is not valid for the workspace audience is rejected with 401 — this + preserves the separation between platform administration and workspace + access. The workspace is resolved from the URL path (``workspace_id`` path parameter) so this factory should only be used on routes whose path includes ``{workspace_id}``. """ role_values = frozenset(r.value for r in roles) - # TREAdmin can access any workspace endpoint - allowed_values = role_values | {TRERole.Admin.value} + allowed_values = role_values | ({TRERole.Admin.value} if allow_tre_admin else frozenset()) role_names = [r.value for r in roles] async def _check( @@ -80,14 +89,24 @@ async def _check( except (TokenExpired, TokenSignatureInvalid) as exc: raise _to_http_exception(exc) except TokenInvalid: - # Wrong audience — fall through to core validator. + # Wrong audience — only a TREAdmin core token may proceed, and + # only when this endpoint opts in via allow_tre_admin. logger.debug( "Workspace token invalid (likely wrong audience), trying core validator" ) + # Endpoints that do not permit TREAdmin get no cross-audience fallback: + # a token that is not valid for the workspace audience is rejected. + if not allow_tre_admin: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_TOKEN, + headers={"WWW-Authenticate": "Bearer"}, + ) + # Fall back to core app registration. A core token is only accepted - # here for TREAdmin (cross-audience access to any workspace); any other - # valid core token is treated as invalid for this workspace resource. + # here for TREAdmin; any other valid core token is treated as invalid + # for this workspace resource. try: user = get_core_validator().validate(token) except AuthError as exc: @@ -113,6 +132,8 @@ async def _check( require_tre_user = require_roles(TRERole.User) require_tre_admin = require_roles(TRERole.Admin) require_tre_user_or_admin = require_roles(TRERole.User, TRERole.Admin) + +# Workspace-scoped checks WITHOUT TREAdmin access (workspace roles only). require_workspace_owner = require_workspace_roles(WorkspaceAccessRole.Owner) require_workspace_researcher = require_workspace_roles(WorkspaceAccessRole.Researcher) require_airlock_manager = require_workspace_roles(WorkspaceAccessRole.AirlockManager) @@ -127,3 +148,18 @@ async def _check( WorkspaceAccessRole.Researcher, WorkspaceAccessRole.AirlockManager, ) + +# Workspace-scoped checks that ALSO permit TREAdmin (mirror the old +# ``..._or_tre_admin`` dependencies). +require_workspace_owner_or_tre_admin = require_workspace_roles( + WorkspaceAccessRole.Owner, allow_tre_admin=True +) +require_workspace_owner_or_researcher_or_tre_admin = require_workspace_roles( + WorkspaceAccessRole.Owner, WorkspaceAccessRole.Researcher, allow_tre_admin=True +) +require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin = require_workspace_roles( + WorkspaceAccessRole.Owner, + WorkspaceAccessRole.Researcher, + WorkspaceAccessRole.AirlockManager, + allow_tre_admin=True, +) diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index c2555a3037..1d238f65dd 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -86,10 +86,10 @@ def _make_fake_deps(self, user_roles, with_client_id=True): return fake_creds, fake_workspace, mock_validator, validated_user def test_admin_always_passes_without_workspace_role(self): - """A TREAdmin using a core token reaches a workspace endpoint via the fallback.""" + """A TREAdmin using a core token reaches an admin-permitted workspace endpoint via the fallback.""" from auth.exceptions import TokenInvalid as _TokenInvalid - dep = require_workspace_roles(WorkspaceAccessRole.Owner) + dep = require_workspace_roles(WorkspaceAccessRole.Owner, allow_tre_admin=True) fake_creds, fake_workspace, _, admin = self._make_fake_deps(["TREAdmin"]) # Workspace validator rejects the core token (wrong audience); core validator accepts it. @@ -147,7 +147,7 @@ def test_wrong_workspace_token_rejected_with_401(self): from fastapi import HTTPException from auth.exceptions import TokenInvalid as _TokenInvalid - dep = require_workspace_roles(WorkspaceAccessRole.Owner) + dep = require_workspace_roles(WorkspaceAccessRole.Owner, allow_tre_admin=True) fake_creds, fake_workspace, _, _ = self._make_fake_deps(["WorkspaceOwner"]) # Workspace validator: wrong audience (workspace B rejects a workspace A token) @@ -179,7 +179,7 @@ def test_non_admin_workspace_user_cannot_elevate_via_core_fallback(self): from fastapi import HTTPException from auth.exceptions import TokenInvalid as _TokenInvalid - dep = require_workspace_roles(WorkspaceAccessRole.Owner) + dep = require_workspace_roles(WorkspaceAccessRole.Owner, allow_tre_admin=True) fake_creds, fake_workspace, _, _ = self._make_fake_deps([]) # Workspace validator rejects the token (wrong audience) @@ -203,6 +203,31 @@ async def _run(): asyncio.get_event_loop().run_until_complete(_run()) + def test_non_admin_endpoint_does_not_fall_back_to_core(self): + """When allow_tre_admin is False, a wrong-audience token is rejected with + 401 and the core validator is never consulted (no cross-audience path).""" + from fastapi import HTTPException + from auth.exceptions import TokenInvalid as _TokenInvalid + + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, _, _ = self._make_fake_deps([]) + + ws_validator_mock = MagicMock() + ws_validator_mock.validate.side_effect = _TokenInvalid("wrong audience") + core_validator_mock = MagicMock() # must NOT be called + + import asyncio + + async def _run(): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator_mock): + with patch('auth.rbac.get_core_validator', return_value=core_validator_mock): + with pytest.raises(HTTPException) as exc_info: + await dep(credentials=fake_creds, workspace=fake_workspace) + assert exc_info.value.status_code == 401 + core_validator_mock.validate.assert_not_called() + + asyncio.get_event_loop().run_until_complete(_run()) + class TestRequireBearerCredentials: def test_missing_credentials_raises_401_with_www_authenticate(self): diff --git a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py index f6f1fe8ebd..bcad94e401 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspace_users.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspace_users.py @@ -8,7 +8,7 @@ from tests_ma.test_api.conftest import create_admin_user from auth.rbac import require_tre_admin, \ require_tre_user_or_admin, \ - require_workspace_owner_or_researcher_or_airlock_manager + require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin from models.domain.workspace import Workspace from resources import strings @@ -46,7 +46,7 @@ def sample_workspace(workspace_id=WORKSPACE_ID, auth_info: dict = {}) -> Workspa class TestWorkspaceUserRoutesWithTreAdmin: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): - app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin] = admin_user app.dependency_overrides[require_tre_user_or_admin] = admin_user app.dependency_overrides[require_tre_admin] = admin_user yield diff --git a/api_app/tests_ma/test_api/test_routes/test_workspaces.py b/api_app/tests_ma/test_api/test_routes/test_workspaces.py index 5e060d0123..8eb8fd8a53 100644 --- a/api_app/tests_ma/test_api/test_routes/test_workspaces.py +++ b/api_app/tests_ma/test_api/test_routes/test_workspaces.py @@ -28,7 +28,9 @@ require_workspace_owner_or_researcher, \ require_workspace_owner_or_researcher_or_airlock_manager, \ require_workspace_owner_or_airlock_manager, \ - require_airlock_manager + require_airlock_manager, \ + require_workspace_owner_or_tre_admin, \ + require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin from azure.cosmos.exceptions import CosmosAccessConditionFailedError @@ -298,12 +300,12 @@ async def test_get_workspace_by_id_get_as_tre_user_returns_403(self, access_serv def forbidden(): raise HTTPException(status_code=status.HTTP_403_FORBIDDEN) - app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = forbidden + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin] = forbidden try: response = await client.get(app.url_path_for(strings.API_GET_WORKSPACE_BY_ID, workspace_id=WORKSPACE_ID)) assert response.status_code == status.HTTP_403_FORBIDDEN finally: - app.dependency_overrides.pop(require_workspace_owner_or_researcher_or_airlock_manager, None) + app.dependency_overrides.pop(require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin, None) # [GET] /workspaces/{workspace_id} @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", side_effect=EntityDoesNotExist) @@ -351,9 +353,11 @@ class TestWorkspaceRoutesThatRequireAdminRights: @pytest.fixture(autouse=True, scope='class') def _prepare(self, app, admin_user): app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = admin_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin] = admin_user app.dependency_overrides[require_workspace_owner_or_researcher] = admin_user app.dependency_overrides[require_workspace_owner_or_airlock_manager] = admin_user app.dependency_overrides[require_workspace_owner] = admin_user + app.dependency_overrides[require_workspace_owner_or_tre_admin] = admin_user app.dependency_overrides[require_airlock_manager] = admin_user app.dependency_overrides[require_tre_user_or_admin] = admin_user app.dependency_overrides[require_tre_admin] = admin_user @@ -735,7 +739,9 @@ class TestWorkspaceServiceRoutesThatRequireOwnerRights: def log_in_with_owner_user(self, app, owner_user): # The following ws services requires the WS app registration app.dependency_overrides[require_workspace_owner] = owner_user + app.dependency_overrides[require_workspace_owner_or_tre_admin] = owner_user app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = owner_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin] = owner_user app.dependency_overrides[require_workspace_owner_or_researcher] = owner_user app.dependency_overrides[require_workspace_owner_or_airlock_manager] = owner_user yield @@ -1383,6 +1389,7 @@ class TestWorkspaceServiceRoutesThatRequireOwnerOrResearcherRights: def log_in_with_researcher_user(self, app, researcher_user): # The following ws services requires the WS app registration app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin] = researcher_user app.dependency_overrides[require_workspace_owner_or_researcher] = researcher_user yield app.dependency_overrides = {} From 67bc00b79bb1379096483ae7af7aa8157c6b7adf Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 22 Jul 2026 16:18:04 +0000 Subject: [PATCH 18/32] Address review comments: require exp claim + asyncio.run in tests - token_validator: add require=[exp,iss,aud] so tokens missing these claims are rejected (PyJWT only validates a claim when present; a token without exp was otherwise accepted indefinitely) - test_rbac: use asyncio.run() instead of get_event_loop().run_until_complete() - add tests for the required-claim behavior --- api_app/auth/token_validator.py | 4 +++ api_app/tests_ma/auth/test_rbac.py | 22 ++++++------- api_app/tests_ma/auth/test_token_validator.py | 32 +++++++++++++++++++ 3 files changed, 47 insertions(+), 11 deletions(-) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index 3917ffc58f..40ede85c3d 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -65,6 +65,10 @@ def validate(self, token: str) -> AuthenticatedUser: "verify_exp": True, "verify_aud": True, "verify_iss": True, + # Reject tokens that omit these claims entirely — PyJWT only + # validates a claim's value when present, so without this a + # token lacking `exp` would never be considered expired. + "require": ["exp", "iss", "aud"], }, ) except jwt.ExpiredSignatureError as exc: diff --git a/api_app/tests_ma/auth/test_rbac.py b/api_app/tests_ma/auth/test_rbac.py index 1d238f65dd..1b69b2b1a0 100644 --- a/api_app/tests_ma/auth/test_rbac.py +++ b/api_app/tests_ma/auth/test_rbac.py @@ -29,7 +29,7 @@ async def _run(): assert result.id == "uid" assert "TREAdmin" in result.roles - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_raises_403_when_user_lacks_role(self): from fastapi import HTTPException @@ -44,7 +44,7 @@ async def _run(): await dep(user=user_with_no_roles) assert exc_info.value.status_code == 403 - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_allows_user_with_any_of_multiple_roles(self): dep = require_roles(TRERole.Admin, TRERole.User) @@ -56,7 +56,7 @@ async def _run(): result = await dep(user=tre_user) assert result.id == "uid" - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) # --------------------------------------------------------------------------- @@ -106,7 +106,7 @@ async def _run(): result = await dep(credentials=fake_creds, workspace=fake_workspace) assert result.id == "uid" - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_workspace_owner_passes(self): dep = require_workspace_roles(WorkspaceAccessRole.Owner) @@ -119,7 +119,7 @@ async def _run(): result = await dep(credentials=fake_creds, workspace=fake_workspace) assert result.id == "uid" - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_raises_403_for_user_without_workspace_role(self): from fastapi import HTTPException @@ -135,7 +135,7 @@ async def _run(): await dep(credentials=fake_creds, workspace=fake_workspace) assert exc_info.value.status_code == 403 - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_wrong_workspace_token_rejected_with_401(self): """A token issued for workspace A must not grant access to workspace B. @@ -167,7 +167,7 @@ async def _run(): await dep(credentials=fake_creds, workspace=fake_workspace) assert exc_info.value.status_code == 401 - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_non_admin_workspace_user_cannot_elevate_via_core_fallback(self): """A non-admin core token must never satisfy a workspace-scoped check. @@ -201,7 +201,7 @@ async def _run(): await dep(credentials=fake_creds, workspace=fake_workspace) assert exc_info.value.status_code == 401 - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_non_admin_endpoint_does_not_fall_back_to_core(self): """When allow_tre_admin is False, a wrong-audience token is rejected with @@ -226,7 +226,7 @@ async def _run(): assert exc_info.value.status_code == 401 core_validator_mock.validate.assert_not_called() - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) class TestRequireBearerCredentials: @@ -242,7 +242,7 @@ async def _run(): assert exc_info.value.status_code == 401 assert exc_info.value.headers.get("WWW-Authenticate") == "Bearer" - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) def test_present_credentials_are_returned(self): from fastapi.security import HTTPAuthorizationCredentials @@ -256,7 +256,7 @@ async def _run(): result = await require_bearer_credentials(credentials=creds) assert result is creds - asyncio.get_event_loop().run_until_complete(_run()) + asyncio.run(_run()) class TestAuthenticatedUserHelpers: diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 1e0f0ced24..5776fc7f76 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -102,6 +102,38 @@ def test_raises_token_invalid_when_signing_key_unavailable(self): with pytest.raises(TokenInvalid, match="Cannot obtain signing key"): validator.validate("any.jwt.token") + def test_token_missing_required_claim_is_rejected(self): + """A token lacking a required claim (e.g. exp) must be rejected, not accepted.""" + import jwt as pyjwt + + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch( + "auth.token_validator.jwt.decode", + side_effect=pyjwt.MissingRequiredClaimError("exp"), + ): + with pytest.raises(TokenInvalid): + validator.validate("token.without.exp") + + def test_exp_is_a_required_claim(self): + """The validator must ask PyJWT to require exp/iss/aud presence.""" + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + captured = {} + + def fake_decode(token, key, **kwargs): + captured.update(kwargs.get("options", {})) + return SAMPLE_CLAIMS + + with patch("auth.token_validator.jwt.decode", side_effect=fake_decode): + validator.validate("token") + + assert "exp" in captured.get("require", []) + def test_email_falls_back_to_preferred_username(self): claims_no_email = { "oid": "uid", From ff0e8923b57d5ebad2ca43fc16169b39915fd640 Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Thu, 23 Jul 2026 16:11:17 +0100 Subject: [PATCH 19/32] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- api_app/tests_ma/test_api/conftest.py | 13 +++++-------- 1 file changed, 5 insertions(+), 8 deletions(-) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index f2223856b7..fa400d0ff5 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -25,14 +25,11 @@ def no_auth_token(): mock_validator = MagicMock() mock_validator.validate.return_value = default_validated - fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") - - with patch('fastapi.security.OAuth2AuthorizationCodeBearer.__call__', new=AsyncMock(return_value="token")): - with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): - with patch('auth.dependencies.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): - yield + with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): + with patch('auth.dependencies.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") From c327b11fa34376838464b29f23a2429f9ec692ea Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Thu, 23 Jul 2026 16:22:36 +0100 Subject: [PATCH 20/32] no_auth_token references fake_credentials Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- api_app/tests_ma/test_api/conftest.py | 28 ++++++++++++++------------- 1 file changed, 15 insertions(+), 13 deletions(-) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index fa400d0ff5..faac83a2d5 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -17,19 +17,21 @@ def no_lifespan_events(): @pytest.fixture(autouse=True) def no_auth_token(): """ overrides validating and decoding tokens for all tests""" - from auth.models import AuthenticatedUser - from fastapi.security import HTTPAuthorizationCredentials - from mock import AsyncMock, MagicMock - - default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) - mock_validator = MagicMock() - mock_validator.validate.return_value = default_validated - - with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): - with patch('auth.dependencies.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): - yield +from auth.models import AuthenticatedUser +from fastapi.security import HTTPAuthorizationCredentials +from mock import AsyncMock, MagicMock + +fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") + +default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) +mock_validator = MagicMock() +mock_validator.validate.return_value = default_validated + +with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): + with patch('auth.dependencies.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") From c79c65123cce8f2ca64f8c0a3126b2223c75fc46 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 23 Jul 2026 16:24:54 +0000 Subject: [PATCH 21/32] Fix no_auth_token fixture indentation: move body inside function scope --- api_app/tests_ma/test_api/conftest.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index faac83a2d5..7ad45f341c 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -17,21 +17,21 @@ def no_lifespan_events(): @pytest.fixture(autouse=True) def no_auth_token(): """ overrides validating and decoding tokens for all tests""" -from auth.models import AuthenticatedUser -from fastapi.security import HTTPAuthorizationCredentials -from mock import AsyncMock, MagicMock + from auth.models import AuthenticatedUser + from fastapi.security import HTTPAuthorizationCredentials + from mock import AsyncMock, MagicMock -fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") + fake_credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials="test-token") -default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) -mock_validator = MagicMock() -mock_validator.validate.return_value = default_validated + default_validated = AuthenticatedUser(id="test-user", name="Test User", roles=["TREAdmin"]) + mock_validator = MagicMock() + mock_validator.validate.return_value = default_validated -with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): - with patch('auth.dependencies.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_core_validator', return_value=mock_validator): - with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): - yield + with patch('fastapi.security.HTTPBearer.__call__', new=AsyncMock(return_value=fake_credentials)): + with patch('auth.dependencies.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_core_validator', return_value=mock_validator): + with patch('auth.rbac.get_workspace_validator', return_value=mock_validator): + yield @pytest.fixture(autouse=True, scope="session") From 4b637447a15db4fbc57f4037d50b06c524303650 Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Thu, 23 Jul 2026 17:40:21 +0100 Subject: [PATCH 22/32] Add KeyError Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- api_app/auth/token_validator.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index 40ede85c3d..d317b0119c 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -78,14 +78,18 @@ def validate(self, token: str) -> AuthenticatedUser: except jwt.InvalidTokenError as exc: raise TokenInvalid(f"Token invalid: {exc}") from exc + from pydantic import ValidationError + try: return AuthenticatedUser( id=claims["oid"], name=claims.get("name", ""), email=claims.get("email") or claims.get("preferred_username"), - roles=claims.get("roles", []), + roles=claims.get("roles") or [], audience=self._config.audience, is_workspace_token=self._config.is_workspace_token, ) except KeyError as exc: raise TokenInvalid("Token is missing required claim: oid") from exc + except (ValidationError, TypeError) as exc: + raise TokenInvalid("Token claims are invalid") from exc From 4fcfb976a8a76cea62a4f3091fa1a03add4dd955 Mon Sep 17 00:00:00 2001 From: James Chapman Date: Fri, 24 Jul 2026 10:16:15 +0000 Subject: [PATCH 23/32] Remove redundant route-level auth dependencies where APIRouter already enforces them (#4797) --- api_app/api/routes/airlock.py | 3 +-- api_app/api/routes/costs.py | 1 - api_app/api/routes/shared_service_templates.py | 2 +- api_app/api/routes/shared_services.py | 4 ++-- api_app/api/routes/user_resource_templates.py | 4 ++-- api_app/api/routes/workspace_service_templates.py | 4 ++-- api_app/api/routes/workspaces.py | 4 ++-- 7 files changed, 10 insertions(+), 12 deletions(-) diff --git a/api_app/api/routes/airlock.py b/api_app/api/routes/airlock.py index dbf96cc117..7e3e5aa7b0 100644 --- a/api_app/api/routes/airlock.py +++ b/api_app/api/routes/airlock.py @@ -54,8 +54,7 @@ async def create_draft_request(airlock_request_input: AirlockRequestInCreate, us status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActionsInList, name=strings.API_LIST_AIRLOCK_REQUESTS, - dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), - Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(get_workspace_by_id_from_path)]) async def get_all_airlock_requests_by_workspace( airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path), diff --git a/api_app/api/routes/costs.py b/api_app/api/routes/costs.py index a248fa846a..1fc4c9bd79 100644 --- a/api_app/api/routes/costs.py +++ b/api_app/api/routes/costs.py @@ -86,7 +86,6 @@ async def costs( @costs_workspace_router.get("/workspaces/{workspace_id}/costs", response_model=WorkspaceCostReport, name=strings.API_GET_WORKSPACE_COSTS, - dependencies=[Depends(require_workspace_owner_or_tre_admin)], responses=get_workspace_cost_report_responses()) async def workspace_costs(workspace_id: UUID4, params: CostsQueryParams = Depends(), cost_service: CostService = Depends(cost_service_factory), diff --git a/api_app/api/routes/shared_service_templates.py b/api_app/api/routes/shared_service_templates.py index 7a487dc1b2..c4d1bf6de8 100644 --- a/api_app/api/routes/shared_service_templates.py +++ b/api_app/api/routes/shared_service_templates.py @@ -22,7 +22,7 @@ async def get_shared_service_templates(authorized_only: bool = False, template_r return ResourceTemplateInformationInList(templates=templates_infos) -@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) +@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME) async def get_shared_service_template(shared_service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServiceTemplateInResponse: try: template = await get_template(shared_service_template_name, template_repo, ResourceType.SharedService, is_update=is_update, version=version) diff --git a/api_app/api/routes/shared_services.py b/api_app/api/routes/shared_services.py index 9ffa89d9b4..0cdc5bc4ea 100644 --- a/api_app/api/routes/shared_services.py +++ b/api_app/api/routes/shared_services.py @@ -32,7 +32,7 @@ def user_is_tre_admin(user): return False -@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES, dependencies=[Depends(require_tre_user_or_admin)]) +@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES) async def retrieve_shared_services(shared_services_repo=Depends(get_repository(SharedServiceRepository)), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServicesInList: shared_services = await shared_services_repo.get_active_shared_services() await asyncio.gather(*[enrich_resource_with_available_upgrades(shared_service, resource_template_repo) for shared_service in shared_services]) @@ -42,7 +42,7 @@ async def retrieve_shared_services(shared_services_repo=Depends(get_repository(S return RestrictedSharedServicesInList(sharedServices=shared_services) -@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(require_tre_user_or_admin), Depends(get_shared_service_by_id_from_path)]) +@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(get_shared_service_by_id_from_path)]) async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_service_by_id_from_path), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))): await enrich_resource_with_available_upgrades(shared_service, resource_template_repo) if user_is_tre_admin(user): diff --git a/api_app/api/routes/user_resource_templates.py b/api_app/api/routes/user_resource_templates.py index 2ae18bfabf..c93703765c 100644 --- a/api_app/api/routes/user_resource_templates.py +++ b/api_app/api/routes/user_resource_templates.py @@ -18,13 +18,13 @@ user_resource_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES) async def get_user_resource_templates_for_service_template(service_template_name: str, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.UserResource, parent_service_name=service_template_name) return ResourceTemplateInformationInList(templates=template_infos) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME) async def get_user_resource_template(service_template_name: str, user_resource_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> UserResourceTemplateInResponse: template = await get_template(user_resource_template_name, template_repo, ResourceType.UserResource, service_template_name, is_update=is_update, version=version) return parse_obj_as(UserResourceTemplateInResponse, template) diff --git a/api_app/api/routes/workspace_service_templates.py b/api_app/api/routes/workspace_service_templates.py index 48acb4aad7..6b374b721f 100644 --- a/api_app/api/routes/workspace_service_templates.py +++ b/api_app/api/routes/workspace_service_templates.py @@ -16,13 +16,13 @@ workspace_service_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) +@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES) async def get_workspace_service_templates(template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: templates_infos = await template_repo.get_templates_information(ResourceType.WorkspaceService) return ResourceTemplateInformationInList(templates=templates_infos) -@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) +@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME) async def get_workspace_service_template(service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceTemplateInResponse: template = await get_template(service_template_name, template_repo, ResourceType.WorkspaceService, is_update=is_update, version=version) return parse_obj_as(WorkspaceServiceTemplateInResponse, template) diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 04ea2d6513..ebf3c56f58 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -228,14 +228,14 @@ async def retrieve_workspace_history_by_workspace_id(workspace=Depends(get_works # WORKSPACE SERVICES ROUTES -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES) async def retrieve_users_active_workspace_services(workspace=Depends(get_workspace_by_id_from_path), workspace_services_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServicesInList: workspace_services = await workspace_services_repo.get_active_workspace_services_for_workspace(workspace.id) await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) for workspace_service in workspace_services]) return WorkspaceServicesInList(workspaceServices=workspace_services) -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_by_id(workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceInResponse: await enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) return WorkspaceServiceInResponse(workspaceService=workspace_service) From 07e77ce45e86c68061496afb5aa179deaac778eb Mon Sep 17 00:00:00 2001 From: James Chapman Date: Fri, 24 Jul 2026 10:38:44 +0000 Subject: [PATCH 24/32] Address Copilot review feedback for #4797: preserve APIRouter level auth and remove duplicate route decorator auth --- CHANGELOG.md | 1 + api_app/api/routes/shared_services.py | 8 ++++---- api_app/api/routes/workspaces.py | 16 ++++++++-------- api_app/tests_ma/test_api/conftest.py | 4 ++-- 4 files changed, 15 insertions(+), 14 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 40ba3dd987..e04d2a2b0e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,6 +16,7 @@ ENHANCEMENTS: * Updated the version of `super-linter` used in the `build_validation_develop` workflow ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) BUG FIXES: +* Remove duplicate API route dependencies to prevent authentication checks from running twice per request. ([#4797](https://github.com/microsoft/AzureTRE/issues/4797)) * Fix UI TypeScript deprecation warning by updating `moduleResolution` to `bundler` in `tsconfig.json`. ([#4968](https://github.com/microsoft/AzureTRE/issues/4968)) * Fix API timeout and name collision failures on workspace creation by checking storage account name availability and improved logging. ([#4946](https://github.com/microsoft/AzureTRE/pull/4946)) * Fix error handling in airlock processor ([#4929](https://github.com/microsoft/AzureTRE/pull/4929)) diff --git a/api_app/api/routes/shared_services.py b/api_app/api/routes/shared_services.py index 0cdc5bc4ea..293e9fb817 100644 --- a/api_app/api/routes/shared_services.py +++ b/api_app/api/routes/shared_services.py @@ -51,7 +51,7 @@ async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_servic return RestrictedSharedServiceInResponse(sharedService=shared_service) -@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE) async def create_shared_service(response: Response, shared_service_input: SharedServiceInCreate, user=Depends(require_tre_admin), shared_services_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: shared_service, resource_template = await shared_services_repo.create_shared_service_item(shared_service_input, user.roles) @@ -82,7 +82,7 @@ async def create_shared_service(response: Response, shared_service_input: Shared status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_SHARED_SERVICE, - dependencies=[Depends(require_tre_admin), Depends(get_shared_service_by_id_from_path)]) + dependencies=[Depends(get_shared_service_by_id_from_path)]) async def patch_shared_service(shared_service_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), etag: str = Header(...), force_version_update: bool = False) -> SharedServiceInResponse: try: patched_shared_service, _ = await shared_service_repo.patch_shared_service(shared_service, shared_service_patch, etag, resource_template_repo, resource_history_repo, user, force_version_update) @@ -105,7 +105,7 @@ async def patch_shared_service(shared_service_patch: ResourcePatch, response: Re raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE) async def delete_shared_service(response: Response, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if shared_service.isEnabled: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=strings.SHARED_SERVICE_NEEDS_TO_BE_DISABLED_BEFORE_DELETION) @@ -124,7 +124,7 @@ async def delete_shared_service(response: Response, user=Depends(require_tre_adm return OperationInResponse(operation=operation) -@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) +@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE) async def invoke_action_on_shared_service(response: Response, action: str, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=shared_service, diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index ebf3c56f58..86ea0df801 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -93,7 +93,7 @@ async def retrieve_workspace_scope_id_by_workspace_id(workspace=Depends(get_work return WorkspaceAuthInResponse(workspaceAuth=wsAuth) -@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE) async def create_workspace(workspace_create: WorkspaceInCreate, response: Response, user=Depends(require_tre_admin), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: # TODO: This requires Directory.ReadAll ( Application.Read.All ) to be enabled in the Azure AD application to enable a users workspaces to be listed. This should be made optional. @@ -127,7 +127,7 @@ async def create_workspace(workspace_create: WorkspaceInCreate, response: Respon return OperationInResponse(operation=operation) -@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE) async def patch_workspace(resource_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), workspace_repo: WorkspaceRepository = Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled @@ -155,7 +155,7 @@ async def patch_workspace(resource_patch: ResourcePatch, response: Response, use raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE) async def delete_workspace(response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace, workspace_repo): operation = await send_uninstall_message( @@ -173,7 +173,7 @@ async def delete_workspace(response: Response, user=Depends(require_tre_admin), return OperationInResponse(operation=operation) -@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE, dependencies=[Depends(require_tre_admin)]) +@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE) async def invoke_action_on_workspace(response: Response, action: str, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace, @@ -241,7 +241,7 @@ async def retrieve_workspace_service_by_id(workspace_service=Depends(get_workspa return WorkspaceServiceInResponse(workspaceService=workspace_service) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE) async def create_workspace_service(response: Response, workspace_service_input: WorkspaceServiceInCreate, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path)) -> OperationInResponse: try: @@ -286,7 +286,7 @@ async def create_workspace_service(response: Response, workspace_service_input: return OperationInResponse(operation=operation) -@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner), Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(get_workspace_by_id_from_path)]) async def patch_workspace_service(resource_patch: ResourcePatch, response: Response, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled @@ -312,7 +312,7 @@ async def patch_workspace_service(resource_patch: ResourcePatch, response: Respo raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE) async def delete_workspace_service(response: Response, user=Depends(require_workspace_owner), workspace=Depends(get_workspace_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace_service, workspace_service_repo): operation = await send_uninstall_message( @@ -330,7 +330,7 @@ async def delete_workspace_service(response: Response, user=Depends(require_work return OperationInResponse(operation=operation) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE) async def invoke_action_on_workspace_service(response: Response, action: str, user=Depends(require_workspace_owner), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace_service, diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index 7ad45f341c..3168330c8c 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -91,11 +91,11 @@ def override_get_user(): def get_required_roles(endpoint): defaults = endpoint.__defaults__ or () - dependencies = list(filter(lambda x: hasattr(x.dependency, 'require_one_of_roles'), defaults)) + dependencies = list(filter(lambda x: hasattr(x, 'dependency') and hasattr(x.dependency, 'require_one_of_roles'), defaults)) if dependencies: return dependencies[0].dependency.require_one_of_roles # New-style deps: check for _role_names attribute on the closure - dependencies = list(filter(lambda x: hasattr(x.dependency, '_role_names'), defaults)) + dependencies = list(filter(lambda x: hasattr(x, 'dependency') and hasattr(x.dependency, '_role_names'), defaults)) if dependencies: return dependencies[0].dependency._role_names return [] From 9c681cc9b0465bccbeceae8cf96aa122610cfe43 Mon Sep 17 00:00:00 2001 From: James Chapman Date: Tue, 28 Jul 2026 08:12:29 +0000 Subject: [PATCH 25/32] Revert "Address Copilot review feedback for #4797: preserve APIRouter level auth and remove duplicate route decorator auth" This reverts commit 07e77ce45e86c68061496afb5aa179deaac778eb. --- CHANGELOG.md | 1 - api_app/api/routes/shared_services.py | 8 ++++---- api_app/api/routes/workspaces.py | 16 ++++++++-------- api_app/tests_ma/test_api/conftest.py | 4 ++-- 4 files changed, 14 insertions(+), 15 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e04d2a2b0e..40ba3dd987 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,7 +16,6 @@ ENHANCEMENTS: * Updated the version of `super-linter` used in the `build_validation_develop` workflow ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) BUG FIXES: -* Remove duplicate API route dependencies to prevent authentication checks from running twice per request. ([#4797](https://github.com/microsoft/AzureTRE/issues/4797)) * Fix UI TypeScript deprecation warning by updating `moduleResolution` to `bundler` in `tsconfig.json`. ([#4968](https://github.com/microsoft/AzureTRE/issues/4968)) * Fix API timeout and name collision failures on workspace creation by checking storage account name availability and improved logging. ([#4946](https://github.com/microsoft/AzureTRE/pull/4946)) * Fix error handling in airlock processor ([#4929](https://github.com/microsoft/AzureTRE/pull/4929)) diff --git a/api_app/api/routes/shared_services.py b/api_app/api/routes/shared_services.py index 293e9fb817..0cdc5bc4ea 100644 --- a/api_app/api/routes/shared_services.py +++ b/api_app/api/routes/shared_services.py @@ -51,7 +51,7 @@ async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_servic return RestrictedSharedServiceInResponse(sharedService=shared_service) -@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE) +@shared_services_router.post("/shared-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) async def create_shared_service(response: Response, shared_service_input: SharedServiceInCreate, user=Depends(require_tre_admin), shared_services_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: shared_service, resource_template = await shared_services_repo.create_shared_service_item(shared_service_input, user.roles) @@ -82,7 +82,7 @@ async def create_shared_service(response: Response, shared_service_input: Shared status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_SHARED_SERVICE, - dependencies=[Depends(get_shared_service_by_id_from_path)]) + dependencies=[Depends(require_tre_admin), Depends(get_shared_service_by_id_from_path)]) async def patch_shared_service(shared_service_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), etag: str = Header(...), force_version_update: bool = False) -> SharedServiceInResponse: try: patched_shared_service, _ = await shared_service_repo.patch_shared_service(shared_service, shared_service_patch, etag, resource_template_repo, resource_history_repo, user, force_version_update) @@ -105,7 +105,7 @@ async def patch_shared_service(shared_service_patch: ResourcePatch, response: Re raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE) +@shared_services_router.delete("/shared-services/{shared_service_id}", response_model=OperationInResponse, name=strings.API_DELETE_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) async def delete_shared_service(response: Response, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if shared_service.isEnabled: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=strings.SHARED_SERVICE_NEEDS_TO_BE_DISABLED_BEFORE_DELETION) @@ -124,7 +124,7 @@ async def delete_shared_service(response: Response, user=Depends(require_tre_adm return OperationInResponse(operation=operation) -@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE) +@shared_services_router.post("/shared-services/{shared_service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_SHARED_SERVICE, dependencies=[Depends(require_tre_admin)]) async def invoke_action_on_shared_service(response: Response, action: str, user=Depends(require_tre_admin), shared_service=Depends(get_shared_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), shared_service_repo=Depends(get_repository(SharedServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=shared_service, diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index 86ea0df801..ebf3c56f58 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -93,7 +93,7 @@ async def retrieve_workspace_scope_id_by_workspace_id(workspace=Depends(get_work return WorkspaceAuthInResponse(workspaceAuth=wsAuth) -@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE) +@workspaces_core_router.post("/workspaces", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) async def create_workspace(workspace_create: WorkspaceInCreate, response: Response, user=Depends(require_tre_admin), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: try: # TODO: This requires Directory.ReadAll ( Application.Read.All ) to be enabled in the Azure AD application to enable a users workspaces to be listed. This should be made optional. @@ -127,7 +127,7 @@ async def create_workspace(workspace_create: WorkspaceInCreate, response: Respon return OperationInResponse(operation=operation) -@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE) +@workspaces_core_router.patch("/workspaces/{workspace_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) async def patch_workspace(resource_patch: ResourcePatch, response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), workspace_repo: WorkspaceRepository = Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled @@ -155,7 +155,7 @@ async def patch_workspace(resource_patch: ResourcePatch, response: Response, use raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE) +@workspaces_core_router.delete("/workspaces/{workspace_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE, dependencies=[Depends(require_tre_admin)]) async def delete_workspace(response: Response, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace, workspace_repo): operation = await send_uninstall_message( @@ -173,7 +173,7 @@ async def delete_workspace(response: Response, user=Depends(require_tre_admin), return OperationInResponse(operation=operation) -@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE) +@workspaces_core_router.post("/workspaces/{workspace_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE, dependencies=[Depends(require_tre_admin)]) async def invoke_action_on_workspace(response: Response, action: str, user=Depends(require_tre_admin), workspace=Depends(get_workspace_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace, @@ -241,7 +241,7 @@ async def retrieve_workspace_service_by_id(workspace_service=Depends(get_workspa return WorkspaceServiceInResponse(workspaceService=workspace_service) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE) +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_CREATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) async def create_workspace_service(response: Response, workspace_service_input: WorkspaceServiceInCreate, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_repo=Depends(get_repository(WorkspaceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path)) -> OperationInResponse: try: @@ -286,7 +286,7 @@ async def create_workspace_service(response: Response, workspace_service_input: return OperationInResponse(operation=operation) -@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.patch("/workspaces/{workspace_id}/workspace-services/{service_id}", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_UPDATE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner), Depends(get_workspace_by_id_from_path)]) async def patch_workspace_service(resource_patch: ResourcePatch, response: Response, user=Depends(require_workspace_owner), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository)), etag: str = Header(...), force_version_update: bool = False) -> OperationInResponse: try: is_disablement = resource_patch.isEnabled is not None and not resource_patch.isEnabled @@ -312,7 +312,7 @@ async def patch_workspace_service(resource_patch: ResourcePatch, response: Respo raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) -@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE) +@workspace_services_workspace_router.delete("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=OperationInResponse, name=strings.API_DELETE_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) async def delete_workspace_service(response: Response, user=Depends(require_workspace_owner), workspace=Depends(get_workspace_by_id_from_path), workspace_service=Depends(get_workspace_service_by_id_from_path), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), user_resource_repo=Depends(get_repository(UserResourceRepository)), operations_repo=Depends(get_repository(OperationRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: if await delete_validation(workspace_service, workspace_service_repo): operation = await send_uninstall_message( @@ -330,7 +330,7 @@ async def delete_workspace_service(response: Response, user=Depends(require_work return OperationInResponse(operation=operation) -@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE) +@workspace_services_workspace_router.post("/workspaces/{workspace_id}/workspace-services/{service_id}/invoke-action", status_code=status.HTTP_202_ACCEPTED, response_model=OperationInResponse, name=strings.API_INVOKE_ACTION_ON_WORKSPACE_SERVICE, dependencies=[Depends(require_workspace_owner)]) async def invoke_action_on_workspace_service(response: Response, action: str, user=Depends(require_workspace_owner), workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository)), operations_repo=Depends(get_repository(OperationRepository)), workspace_service_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_history_repo=Depends(get_repository(ResourceHistoryRepository))) -> OperationInResponse: operation = await send_custom_action_message( resource=workspace_service, diff --git a/api_app/tests_ma/test_api/conftest.py b/api_app/tests_ma/test_api/conftest.py index 3168330c8c..7ad45f341c 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -91,11 +91,11 @@ def override_get_user(): def get_required_roles(endpoint): defaults = endpoint.__defaults__ or () - dependencies = list(filter(lambda x: hasattr(x, 'dependency') and hasattr(x.dependency, 'require_one_of_roles'), defaults)) + dependencies = list(filter(lambda x: hasattr(x.dependency, 'require_one_of_roles'), defaults)) if dependencies: return dependencies[0].dependency.require_one_of_roles # New-style deps: check for _role_names attribute on the closure - dependencies = list(filter(lambda x: hasattr(x, 'dependency') and hasattr(x.dependency, '_role_names'), defaults)) + dependencies = list(filter(lambda x: hasattr(x.dependency, '_role_names'), defaults)) if dependencies: return dependencies[0].dependency._role_names return [] From 5b8a4ad68881e9558db744f1b068a44a29f98fa0 Mon Sep 17 00:00:00 2001 From: James Chapman Date: Tue, 28 Jul 2026 08:12:29 +0000 Subject: [PATCH 26/32] Revert "Remove redundant route-level auth dependencies where APIRouter already enforces them (#4797)" This reverts commit 4fcfb976a8a76cea62a4f3091fa1a03add4dd955. --- api_app/api/routes/airlock.py | 3 ++- api_app/api/routes/costs.py | 1 + api_app/api/routes/shared_service_templates.py | 2 +- api_app/api/routes/shared_services.py | 4 ++-- api_app/api/routes/user_resource_templates.py | 4 ++-- api_app/api/routes/workspace_service_templates.py | 4 ++-- api_app/api/routes/workspaces.py | 4 ++-- 7 files changed, 12 insertions(+), 10 deletions(-) diff --git a/api_app/api/routes/airlock.py b/api_app/api/routes/airlock.py index 7e3e5aa7b0..dbf96cc117 100644 --- a/api_app/api/routes/airlock.py +++ b/api_app/api/routes/airlock.py @@ -54,7 +54,8 @@ async def create_draft_request(airlock_request_input: AirlockRequestInCreate, us status_code=status_code.HTTP_200_OK, response_model=AirlockRequestWithAllowedUserActionsInList, name=strings.API_LIST_AIRLOCK_REQUESTS, - dependencies=[Depends(get_workspace_by_id_from_path)]) + dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), + Depends(get_workspace_by_id_from_path)]) async def get_all_airlock_requests_by_workspace( airlock_request_repo=Depends(get_repository(AirlockRequestRepository)), workspace=Depends(get_deployed_workspace_by_id_from_path), diff --git a/api_app/api/routes/costs.py b/api_app/api/routes/costs.py index 1fc4c9bd79..a248fa846a 100644 --- a/api_app/api/routes/costs.py +++ b/api_app/api/routes/costs.py @@ -86,6 +86,7 @@ async def costs( @costs_workspace_router.get("/workspaces/{workspace_id}/costs", response_model=WorkspaceCostReport, name=strings.API_GET_WORKSPACE_COSTS, + dependencies=[Depends(require_workspace_owner_or_tre_admin)], responses=get_workspace_cost_report_responses()) async def workspace_costs(workspace_id: UUID4, params: CostsQueryParams = Depends(), cost_service: CostService = Depends(cost_service_factory), diff --git a/api_app/api/routes/shared_service_templates.py b/api_app/api/routes/shared_service_templates.py index c4d1bf6de8..7a487dc1b2 100644 --- a/api_app/api/routes/shared_service_templates.py +++ b/api_app/api/routes/shared_service_templates.py @@ -22,7 +22,7 @@ async def get_shared_service_templates(authorized_only: bool = False, template_r return ResourceTemplateInformationInList(templates=templates_infos) -@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME) +@shared_service_templates_core_router.get("/shared-service-templates/{shared_service_template_name}", response_model=SharedServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_SHARED_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_shared_service_template(shared_service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServiceTemplateInResponse: try: template = await get_template(shared_service_template_name, template_repo, ResourceType.SharedService, is_update=is_update, version=version) diff --git a/api_app/api/routes/shared_services.py b/api_app/api/routes/shared_services.py index 0cdc5bc4ea..9ffa89d9b4 100644 --- a/api_app/api/routes/shared_services.py +++ b/api_app/api/routes/shared_services.py @@ -32,7 +32,7 @@ def user_is_tre_admin(user): return False -@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES) +@shared_services_router.get("/shared-services", response_model=SharedServicesInList, name=strings.API_GET_ALL_SHARED_SERVICES, dependencies=[Depends(require_tre_user_or_admin)]) async def retrieve_shared_services(shared_services_repo=Depends(get_repository(SharedServiceRepository)), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> SharedServicesInList: shared_services = await shared_services_repo.get_active_shared_services() await asyncio.gather(*[enrich_resource_with_available_upgrades(shared_service, resource_template_repo) for shared_service in shared_services]) @@ -42,7 +42,7 @@ async def retrieve_shared_services(shared_services_repo=Depends(get_repository(S return RestrictedSharedServicesInList(sharedServices=shared_services) -@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(get_shared_service_by_id_from_path)]) +@shared_services_router.get("/shared-services/{shared_service_id}", response_model=SharedServiceInResponse, name=strings.API_GET_SHARED_SERVICE_BY_ID, dependencies=[Depends(require_tre_user_or_admin), Depends(get_shared_service_by_id_from_path)]) async def retrieve_shared_service_by_id(shared_service=Depends(get_shared_service_by_id_from_path), user=Depends(require_tre_user_or_admin), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))): await enrich_resource_with_available_upgrades(shared_service, resource_template_repo) if user_is_tre_admin(user): diff --git a/api_app/api/routes/user_resource_templates.py b/api_app/api/routes/user_resource_templates.py index c93703765c..2ae18bfabf 100644 --- a/api_app/api/routes/user_resource_templates.py +++ b/api_app/api/routes/user_resource_templates.py @@ -18,13 +18,13 @@ user_resource_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_USER_RESOURCE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) async def get_user_resource_templates_for_service_template(service_template_name: str, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: template_infos = await template_repo.get_templates_information(ResourceType.UserResource, parent_service_name=service_template_name) return ResourceTemplateInformationInList(templates=template_infos) -@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME) +@user_resource_templates_core_router.get("/workspace-service-templates/{service_template_name}/user-resource-templates/{user_resource_template_name}", response_model=UserResourceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_USER_RESOURCE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_user_resource_template(service_template_name: str, user_resource_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> UserResourceTemplateInResponse: template = await get_template(user_resource_template_name, template_repo, ResourceType.UserResource, service_template_name, is_update=is_update, version=version) return parse_obj_as(UserResourceTemplateInResponse, template) diff --git a/api_app/api/routes/workspace_service_templates.py b/api_app/api/routes/workspace_service_templates.py index 6b374b721f..48acb4aad7 100644 --- a/api_app/api/routes/workspace_service_templates.py +++ b/api_app/api/routes/workspace_service_templates.py @@ -16,13 +16,13 @@ workspace_service_templates_core_router = APIRouter(dependencies=[Depends(require_tre_user_or_admin)]) -@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES) +@workspace_service_templates_core_router.get("/workspace-service-templates", response_model=ResourceTemplateInformationInList, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATES, dependencies=[Depends(require_tre_user_or_admin)]) async def get_workspace_service_templates(template_repo=Depends(get_repository(ResourceTemplateRepository))) -> ResourceTemplateInformationInList: templates_infos = await template_repo.get_templates_information(ResourceType.WorkspaceService) return ResourceTemplateInformationInList(templates=templates_infos) -@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME) +@workspace_service_templates_core_router.get("/workspace-service-templates/{service_template_name}", response_model=WorkspaceServiceTemplateInResponse, response_model_exclude_none=True, name=strings.API_GET_WORKSPACE_SERVICE_TEMPLATE_BY_NAME, dependencies=[Depends(require_tre_user_or_admin)]) async def get_workspace_service_template(service_template_name: str, is_update: bool = False, version: Optional[str] = None, template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceTemplateInResponse: template = await get_template(service_template_name, template_repo, ResourceType.WorkspaceService, is_update=is_update, version=version) return parse_obj_as(WorkspaceServiceTemplateInResponse, template) diff --git a/api_app/api/routes/workspaces.py b/api_app/api/routes/workspaces.py index ebf3c56f58..04ea2d6513 100644 --- a/api_app/api/routes/workspaces.py +++ b/api_app/api/routes/workspaces.py @@ -228,14 +228,14 @@ async def retrieve_workspace_history_by_workspace_id(workspace=Depends(get_works # WORKSPACE SERVICES ROUTES -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services", response_model=WorkspaceServicesInList, name=strings.API_GET_ALL_WORKSPACE_SERVICES, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager)]) async def retrieve_users_active_workspace_services(workspace=Depends(get_workspace_by_id_from_path), workspace_services_repo=Depends(get_repository(WorkspaceServiceRepository)), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServicesInList: workspace_services = await workspace_services_repo.get_active_workspace_services_for_workspace(workspace.id) await asyncio.gather(*[enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) for workspace_service in workspace_services]) return WorkspaceServicesInList(workspaceServices=workspace_services) -@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(get_workspace_by_id_from_path)]) +@workspace_services_workspace_router.get("/workspaces/{workspace_id}/workspace-services/{service_id}", response_model=WorkspaceServiceInResponse, name=strings.API_GET_WORKSPACE_SERVICE_BY_ID, dependencies=[Depends(require_workspace_owner_or_researcher_or_airlock_manager), Depends(get_workspace_by_id_from_path)]) async def retrieve_workspace_service_by_id(workspace_service=Depends(get_workspace_service_by_id_from_path), resource_template_repo=Depends(get_repository(ResourceTemplateRepository))) -> WorkspaceServiceInResponse: await enrich_resource_with_available_upgrades(workspace_service, resource_template_repo) return WorkspaceServiceInResponse(workspaceService=workspace_service) From c6df04921db02e8d067653e0646de64301a35d89 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 29 Jul 2026 13:55:31 +0000 Subject: [PATCH 27/32] Add Event Grid retry, split Graph/EG errors in airlock, bump version to 0.26.1 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - event_grid/helpers.py: exponential-backoff retry for HttpResponseError 429/5xx and ServiceRequestError; re-raises after 3 attempts - services/airlock.py: separate Graph role-assignment lookup into its own try/except so auth/Graph failures are reported distinctly; include underlying exception detail in Event Grid 503 responses - resources/strings.py: add EVENT_GRID_PUBLISH_FAILED and GRAPH_ROLE_ASSIGNMENT_ERROR - _version.py: 0.26.0 → 0.26.1 - tests: 11 new tests (7 helpers retry, 4 airlock error separation); all 722 pass - CHANGELOG.md: updated ENHANCEMENTS --- CHANGELOG.md | 1 + api_app/_version.py | 2 +- api_app/event_grid/helpers.py | 46 ++++- api_app/resources/strings.py | 4 + api_app/services/airlock.py | 26 ++- api_app/tests_ma/test_event_grid/__init__.py | 0 .../tests_ma/test_event_grid/test_helpers.py | 168 ++++++++++++++++++ .../tests_ma/test_services/test_airlock.py | 96 ++++++++++ 8 files changed, 329 insertions(+), 14 deletions(-) create mode 100644 api_app/tests_ma/test_event_grid/__init__.py create mode 100644 api_app/tests_ma/test_event_grid/test_helpers.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 8e0c80ae69..75e078ce71 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,7 @@ ENHANCEMENTS: * Add support for setting resource processor VMSS SKU via environment variables ([#4936](https://github.com/microsoft/AzureTRE/issues/4936)) * Exclude recovery service vaults from e2e tests ([#4920](https://github.com/microsoft/AzureTRE/issues/4920)) * Strengthen TRE API authentication: introduce layered `auth/` package with typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, and composable RBAC factories; remove the `AccessService` abstraction that is no longer needed now that Entra ID is the only auth provider. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) +* Add exponential-backoff retry in Event Grid publisher for transient failures (HTTP 429/5xx and network errors); split Graph role-assignment lookup from Event Grid publish in airlock service so auth failures are reported distinctly and underlying exception detail is included in 503 responses. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Add support for formatting UI code via `pre-commit` and fix existing formatting issues. ([#4955](https://github.com/microsoft/AzureTRE/issues/4955)) * Update the version of `super-linter` used in the `build_validation_develop` workflow to 8.7.0 ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) diff --git a/api_app/_version.py b/api_app/_version.py index 7c4a9591e1..025f4c5d0b 100644 --- a/api_app/_version.py +++ b/api_app/_version.py @@ -1 +1 @@ -__version__ = "0.26.0" +__version__ = "0.26.1" diff --git a/api_app/event_grid/helpers.py b/api_app/event_grid/helpers.py index bcad3e65e1..61b21dc997 100644 --- a/api_app/event_grid/helpers.py +++ b/api_app/event_grid/helpers.py @@ -1,10 +1,46 @@ +import asyncio + +from azure.core.exceptions import HttpResponseError, ServiceRequestError from azure.eventgrid import EventGridEvent from azure.eventgrid.aio import EventGridPublisherClient from core import credentials +from services.logging import logger + +_MAX_RETRIES = 3 +_BASE_DELAY_SECONDS = 1.0 + + +def _is_retryable(exc: HttpResponseError) -> bool: + """Return True for 429 (rate-limited) and any 5xx (server-side) errors.""" + return exc.status_code == 429 or (exc.status_code is not None and exc.status_code >= 500) + + +async def publish_event(event: EventGridEvent, topic_endpoint: str) -> None: + last_exc: Exception | None = None + for attempt in range(_MAX_RETRIES): + try: + async with credentials.get_credential_async_context() as credential: + client = EventGridPublisherClient(topic_endpoint, credential) + async with client: + await client.send([event]) + return + except HttpResponseError as exc: + if not _is_retryable(exc): + raise + last_exc = exc + logger.warning( + f"Event Grid publish failed with HTTP {exc.status_code} " + f"(attempt {attempt + 1}/{_MAX_RETRIES}): {exc}" + ) + except ServiceRequestError as exc: + last_exc = exc + logger.warning( + f"Event Grid publish failed with a transient network error " + f"(attempt {attempt + 1}/{_MAX_RETRIES}): {exc}" + ) + if attempt < _MAX_RETRIES - 1: + delay = _BASE_DELAY_SECONDS * (2 ** attempt) + await asyncio.sleep(delay) -async def publish_event(event: EventGridEvent, topic_endpoint: str): - async with credentials.get_credential_async_context() as credential: - client = EventGridPublisherClient(topic_endpoint, credential) - async with client: - await client.send([event]) + raise last_exc # type: ignore[misc] diff --git a/api_app/resources/strings.py b/api_app/resources/strings.py index c54a40ba52..7cda05e95a 100644 --- a/api_app/resources/strings.py +++ b/api_app/resources/strings.py @@ -262,6 +262,10 @@ # Event grid EVENT_GRID_GENERAL_ERROR_MESSAGE = "Event grid failure" +EVENT_GRID_PUBLISH_FAILED = "Failed to publish Event Grid event: {}" + +# Graph / role assignments +GRAPH_ROLE_ASSIGNMENT_ERROR = "Failed to fetch workspace role assignments from Microsoft Graph: {}" # Workspace creation validation MISSING_REQUIRED_PARAMETERS = "Missing required parameters" diff --git a/api_app/services/airlock.py b/api_app/services/airlock.py index eced3d5bfa..9a34089b46 100644 --- a/api_app/services/airlock.py +++ b/api_app/services/airlock.py @@ -272,8 +272,13 @@ async def _handle_existing_review_resource(existing_resource: AirlockReviewUserR async def save_and_publish_event_airlock_request(airlock_request: AirlockRequest, airlock_request_repo: AirlockRequestRepository, user: User, workspace: Workspace): - access_service = get_aad_service() - role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) + try: + access_service = get_aad_service() + role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) + except Exception as e: + logger.exception("Failed to retrieve workspace role assignments from Microsoft Graph") + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR.format(e)) + if config.ENABLE_AIRLOCK_EMAIL_CHECK: check_email_exists(role_assignment_details) @@ -290,10 +295,10 @@ async def save_and_publish_event_airlock_request(airlock_request: AirlockRequest logger.debug(f"Sending status changed event for airlock request item: {airlock_request.id}") await send_status_changed_event(airlock_request=airlock_request, previous_status=None) await send_airlock_notification_event(airlock_request, workspace, role_assignment_details) - except Exception: + except Exception as e: await airlock_request_repo.delete_item(airlock_request.id) logger.exception("Failed sending status_changed message") - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.EVENT_GRID_GENERAL_ERROR_MESSAGE) + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.EVENT_GRID_PUBLISH_FAILED.format(e)) async def update_and_publish_event_airlock_request( @@ -329,15 +334,20 @@ async def update_and_publish_event_airlock_request( return updated_airlock_request try: - logger.debug(f"Sending status changed event for airlock request item: {airlock_request.id}") - await send_status_changed_event(airlock_request=updated_airlock_request, previous_status=airlock_request.status) access_service = get_aad_service() role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) + except Exception as e: + logger.exception("Failed to retrieve workspace role assignments from Microsoft Graph") + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR.format(e)) + + try: + logger.debug(f"Sending status changed event for airlock request item: {airlock_request.id}") + await send_status_changed_event(airlock_request=updated_airlock_request, previous_status=airlock_request.status) await send_airlock_notification_event(updated_airlock_request, workspace, role_assignment_details) return updated_airlock_request - except Exception: + except Exception as e: logger.exception("Failed sending status_changed message") - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.EVENT_GRID_GENERAL_ERROR_MESSAGE) + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.EVENT_GRID_PUBLISH_FAILED.format(e)) def get_timestamp() -> float: diff --git a/api_app/tests_ma/test_event_grid/__init__.py b/api_app/tests_ma/test_event_grid/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/api_app/tests_ma/test_event_grid/test_helpers.py b/api_app/tests_ma/test_event_grid/test_helpers.py new file mode 100644 index 0000000000..042c152ffa --- /dev/null +++ b/api_app/tests_ma/test_event_grid/test_helpers.py @@ -0,0 +1,168 @@ +import pytest +from unittest.mock import AsyncMock, MagicMock, patch +from azure.core.exceptions import HttpResponseError, ServiceRequestError +from azure.eventgrid import EventGridEvent + +from event_grid.helpers import publish_event + + +def _make_event(): + return EventGridEvent( + event_type="test", + data={"key": "value"}, + subject="test/subject", + data_version="1.0", + ) + + +def _http_error(status_code: int) -> HttpResponseError: + err = HttpResponseError() + err.status_code = status_code + return err + + +@pytest.mark.asyncio +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_succeeds_on_first_attempt(mock_cred_ctx): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock() + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + await publish_event(_make_event(), "https://topic.endpoint") + + mock_client.send.assert_called_once() + + +@pytest.mark.asyncio +@patch("event_grid.helpers.asyncio.sleep", new_callable=AsyncMock) +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_retries_on_429_and_succeeds(mock_cred_ctx, mock_sleep): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=[_http_error(429), None]) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + await publish_event(_make_event(), "https://topic.endpoint") + + assert mock_client.send.call_count == 2 + mock_sleep.assert_called_once() + + +@pytest.mark.asyncio +@patch("event_grid.helpers.asyncio.sleep", new_callable=AsyncMock) +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_retries_on_503_and_succeeds(mock_cred_ctx, mock_sleep): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=[_http_error(503), None]) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + await publish_event(_make_event(), "https://topic.endpoint") + + assert mock_client.send.call_count == 2 + mock_sleep.assert_called_once() + + +@pytest.mark.asyncio +@patch("event_grid.helpers.asyncio.sleep", new_callable=AsyncMock) +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_retries_on_service_request_error(mock_cred_ctx, mock_sleep): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=[ServiceRequestError("network error"), None]) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + await publish_event(_make_event(), "https://topic.endpoint") + + assert mock_client.send.call_count == 2 + mock_sleep.assert_called_once() + + +@pytest.mark.asyncio +@patch("event_grid.helpers.asyncio.sleep", new_callable=AsyncMock) +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_raises_after_exhausting_retries(mock_cred_ctx, mock_sleep): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=_http_error(429)) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + with pytest.raises(HttpResponseError): + await publish_event(_make_event(), "https://topic.endpoint") + + assert mock_client.send.call_count == 3 # _MAX_RETRIES + assert mock_sleep.call_count == 2 # no sleep after last attempt + + +@pytest.mark.asyncio +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_does_not_retry_on_non_retryable_http_error(mock_cred_ctx): + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=_http_error(401)) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + with pytest.raises(HttpResponseError): + await publish_event(_make_event(), "https://topic.endpoint") + + assert mock_client.send.call_count == 1 # no retries for 401 + + +@pytest.mark.asyncio +@patch("event_grid.helpers.asyncio.sleep", new_callable=AsyncMock) +@patch("event_grid.helpers.credentials.get_credential_async_context") +async def test_publish_event_exponential_backoff_delays(mock_cred_ctx, mock_sleep): + """Verify delays grow as 1s, 2s (base * 2^attempt).""" + mock_client = AsyncMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.send = AsyncMock(side_effect=_http_error(429)) + + mock_cred = MagicMock() + mock_cred.__aenter__ = AsyncMock(return_value=mock_cred) + mock_cred.__aexit__ = AsyncMock(return_value=None) + mock_cred_ctx.return_value = mock_cred + + with patch("event_grid.helpers.EventGridPublisherClient", return_value=mock_client): + with pytest.raises(HttpResponseError): + await publish_event(_make_event(), "https://topic.endpoint") + + delays = [call.args[0] for call in mock_sleep.call_args_list] + assert delays == [1.0, 2.0] diff --git a/api_app/tests_ma/test_services/test_airlock.py b/api_app/tests_ma/test_services/test_airlock.py index 31cb6a0068..04e7b7ac45 100644 --- a/api_app/tests_ma/test_services/test_airlock.py +++ b/api_app/tests_ma/test_services/test_airlock.py @@ -586,3 +586,99 @@ async def test_delete_review_user_resource_disables_the_resource_before_deletion resource_history_repo=AsyncMock(), user=create_test_user()) disable_user_resource.assert_called_once() + + +# --- Graph / role-assignment error separation tests --- + +@pytest.mark.asyncio +@patch("services.aad_authentication.AzureADAuthorization.get_workspace_user_emails_by_role_assignment", + side_effect=Exception("Graph call failed")) +async def test_save_and_publish_event_airlock_request_raises_503_if_graph_call_fails(_, airlock_request_repo_mock): + """A Graph/auth failure during role-assignment lookup returns a distinct 503 before DB or Event Grid are touched.""" + airlock_request_mock = sample_airlock_request() + airlock_request_repo_mock.save_item = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as ex: + await save_and_publish_event_airlock_request( + airlock_request=airlock_request_mock, + airlock_request_repo=airlock_request_repo_mock, + user=create_test_user(), + workspace=sample_workspace()) + + assert ex.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert "Graph" in ex.value.detail + # DB should never have been touched + airlock_request_repo_mock.save_item.assert_not_called() + + +@pytest.mark.asyncio +@patch("event_grid.helpers.EventGridPublisherClient", return_value=AsyncMock()) +@patch("services.aad_authentication.AzureADAuthorization.get_workspace_user_emails_by_role_assignment", + return_value={"WorkspaceResearcher": ["r@example.com"], "WorkspaceOwner": ["o@example.com"], "AirlockManager": ["m@example.com"]}) +async def test_save_and_publish_event_airlock_request_503_includes_underlying_error(_, event_grid_publisher_client_mock, + airlock_request_repo_mock): + """The 503 detail from an Event Grid failure includes the underlying exception message.""" + underlying_msg = "connection reset by peer" + airlock_request_mock = sample_airlock_request() + airlock_request_repo_mock.save_item = AsyncMock(return_value=None) + airlock_request_repo_mock.delete_item = AsyncMock(return_value=None) + event_grid_sender_client_mock = event_grid_publisher_client_mock.return_value + event_grid_sender_client_mock.send = AsyncMock(side_effect=Exception(underlying_msg)) + + with pytest.raises(HTTPException) as ex: + await save_and_publish_event_airlock_request( + airlock_request=airlock_request_mock, + airlock_request_repo=airlock_request_repo_mock, + user=create_test_user(), + workspace=sample_workspace()) + + assert ex.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert underlying_msg in ex.value.detail + + +@pytest.mark.asyncio +@patch("services.aad_authentication.AzureADAuthorization.get_workspace_user_emails_by_role_assignment", + side_effect=Exception("Graph call failed")) +async def test_update_and_publish_event_airlock_request_raises_503_if_graph_call_fails(_, airlock_request_repo_mock): + """A Graph failure during role-assignment lookup in update_and_publish raises a distinct 503.""" + airlock_request_mock = sample_airlock_request() + updated_airlock_request_mock = sample_airlock_request(status=AirlockRequestStatus.Submitted) + airlock_request_repo_mock.update_airlock_request = AsyncMock(return_value=updated_airlock_request_mock) + + with patch("services.airlock.send_status_changed_event", new_callable=AsyncMock): + with pytest.raises(HTTPException) as ex: + await update_and_publish_event_airlock_request( + airlock_request=airlock_request_mock, + airlock_request_repo=airlock_request_repo_mock, + updated_by=create_test_user(), + new_status=AirlockRequestStatus.Submitted, + workspace=sample_workspace()) + + assert ex.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert "Graph" in ex.value.detail + + +@pytest.mark.asyncio +@patch("event_grid.helpers.EventGridPublisherClient", return_value=AsyncMock()) +@patch("services.aad_authentication.AzureADAuthorization.get_workspace_user_emails_by_role_assignment", + return_value={"WorkspaceResearcher": ["r@example.com"], "WorkspaceOwner": ["o@example.com"], "AirlockManager": ["m@example.com"]}) +async def test_update_and_publish_event_airlock_request_503_includes_underlying_error(_, event_grid_publisher_client_mock, + airlock_request_repo_mock): + """The 503 detail from an Event Grid failure in update_and_publish includes the underlying exception message.""" + underlying_msg = "topic endpoint unreachable" + airlock_request_mock = sample_airlock_request() + updated_airlock_request_mock = sample_airlock_request(status=AirlockRequestStatus.Submitted) + airlock_request_repo_mock.update_airlock_request = AsyncMock(return_value=updated_airlock_request_mock) + event_grid_sender_client_mock = event_grid_publisher_client_mock.return_value + event_grid_sender_client_mock.send = AsyncMock(side_effect=Exception(underlying_msg)) + + with pytest.raises(HTTPException) as ex: + await update_and_publish_event_airlock_request( + airlock_request=airlock_request_mock, + airlock_request_repo=airlock_request_repo_mock, + updated_by=create_test_user(), + new_status=AirlockRequestStatus.Submitted, + workspace=sample_workspace()) + + assert ex.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert underlying_msg in ex.value.detail From fde350fa42ccebdb28b09a92fd8462d07f4f317f Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 29 Jul 2026 16:15:14 +0000 Subject: [PATCH 28/32] Default AuthenticatedUser email to empty string to fix airlock Event Grid 503 Co-authored-by: marrobi <17089773+marrobi@users.noreply.github.com> --- CHANGELOG.md | 1 + api_app/_version.py | 2 +- api_app/auth/token_validator.py | 2 +- api_app/tests_ma/auth/test_token_validator.py | 15 ++++++ .../tests_ma/test_services/test_airlock.py | 50 +++++++++++++++++++ 5 files changed, 68 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 75e078ce71..8890da5513 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,6 +18,7 @@ ENHANCEMENTS: * Update the version of `super-linter` used in the `build_validation_develop` workflow to 8.7.0 ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) BUG FIXES: +* Fix airlock Event Grid 503 when a user token omits both `email` and `preferred_username` claims: the derived `AuthenticatedUser` email now defaults to an empty string so the airlock notification payload (which requires a string email) no longer fails validation. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Fix Nexus shared service security: fetch admin password from Key Vault at runtime via managed identity (IMDS) instead of embedding it in the VM Run Command script content. Fix `deploy_nexus_container.sh` short-circuit path to fail loudly if the container does not start. (`sonatype-nexus` 3.10.0) ([#4983](https://github.com/microsoft/AzureTRE/pull/4983)) * Fix UI TypeScript deprecation warning by updating `moduleResolution` to `bundler` in `tsconfig.json`. ([#4968](https://github.com/microsoft/AzureTRE/issues/4968)) * Fix API timeout and name collision failures on workspace creation by checking storage account name availability and improved logging. ([#4946](https://github.com/microsoft/AzureTRE/pull/4946)) diff --git a/api_app/_version.py b/api_app/_version.py index 025f4c5d0b..40ab52fff9 100644 --- a/api_app/_version.py +++ b/api_app/_version.py @@ -1 +1 @@ -__version__ = "0.26.1" +__version__ = "0.26.2" diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py index d317b0119c..32a4e31b60 100644 --- a/api_app/auth/token_validator.py +++ b/api_app/auth/token_validator.py @@ -84,7 +84,7 @@ def validate(self, token: str) -> AuthenticatedUser: return AuthenticatedUser( id=claims["oid"], name=claims.get("name", ""), - email=claims.get("email") or claims.get("preferred_username"), + email=claims.get("email") or claims.get("preferred_username") or "", roles=claims.get("roles") or [], audience=self._config.audience, is_workspace_token=self._config.is_workspace_token, diff --git a/api_app/tests_ma/auth/test_token_validator.py b/api_app/tests_ma/auth/test_token_validator.py index 5776fc7f76..f0833772e4 100644 --- a/api_app/tests_ma/auth/test_token_validator.py +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -150,6 +150,21 @@ def test_email_falls_back_to_preferred_username(self): assert result.email == "user@tenant.com" + def test_email_defaults_to_empty_string_when_no_email_or_preferred_username(self): + # Regression: a token with neither `email` nor `preferred_username` + # must yield an empty-string email (not None), so that downstream + # consumers requiring a string email (e.g. the Event Grid airlock + # notification payload) do not fail validation and return a 503. + claims_no_email = {"oid": "uid", "name": "User", "roles": []} + signing_key = MagicMock() + mock_client = _make_mock_jwks_client(signing_key) + validator = _make_validator(mock_client) + + with patch("auth.token_validator.jwt.decode", return_value=claims_no_email): + result = validator.validate("token") + + assert result.email == "" + def test_roles_default_to_empty_list(self): claims_no_roles = {"oid": "uid", "name": "User"} signing_key = MagicMock() diff --git a/api_app/tests_ma/test_services/test_airlock.py b/api_app/tests_ma/test_services/test_airlock.py index 04e7b7ac45..caa0182f65 100644 --- a/api_app/tests_ma/test_services/test_airlock.py +++ b/api_app/tests_ma/test_services/test_airlock.py @@ -682,3 +682,53 @@ async def test_update_and_publish_event_airlock_request_503_includes_underlying_ assert ex.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE assert underlying_msg in ex.value.detail + + +def test_authenticated_user_with_empty_email_persists_and_builds_notification(): + """Regression: an AuthenticatedUser with an empty email and tuple roles is a + valid createdBy/updatedBy for an airlock request and does not break the + Event Grid airlock notification payload (which requires a string email). + + A missing `email`/`preferred_username` claim previously yielded a None email + that failed AirlockNotificationUserData validation, surfacing as a 503. + """ + from auth.models import AuthenticatedUser + + user = AuthenticatedUser( + id="user-oid", + name="No Email User", + email="", + roles=("WorkspaceResearcher",), + audience="api://workspace", + is_workspace_token=True, + ) + + airlock_request = AirlockRequest( + id=AIRLOCK_REQUEST_ID, + workspaceId=WORKSPACE_ID, + type=AirlockRequestType.Import, + createdBy=user, + updatedBy=user, + createdWhen=CURRENT_TIME, + updatedWhen=CURRENT_TIME, + ) + + # createdBy/updatedBy are coerced to plain dicts, roles remain a tuple. + assert isinstance(airlock_request.createdBy, dict) + assert airlock_request.createdBy["email"] == "" + assert airlock_request.createdBy["roles"] == ("WorkspaceResearcher",) + + # Building the notification payload from the persisted user must not raise. + notification = AirlockNotificationRequestData( + id=airlock_request.id, + created_when=airlock_request.createdWhen, + created_by=airlock_request.createdBy, + updated_when=airlock_request.updatedWhen, + updated_by=airlock_request.updatedBy, + request_type=airlock_request.type, + files=airlock_request.files, + status=airlock_request.status, + business_justification=airlock_request.businessJustification) + + assert notification.created_by.email == "" + assert notification.updated_by.email == "" From 764f69abbe322aec000010e3a522c5965da39caa Mon Sep 17 00:00:00 2001 From: Marcus Robinson Date: Wed, 29 Jul 2026 17:23:56 +0100 Subject: [PATCH 29/32] Update CHANGELOG with recent bug fixes Fix airlock Event Grid 503 error by defaulting AuthenticatedUser email to an empty string. Update Nexus security to fetch admin password from Key Vault and fix UI TypeScript deprecation warning. --- CHANGELOG.md | 1 - 1 file changed, 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 8890da5513..75e078ce71 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,7 +18,6 @@ ENHANCEMENTS: * Update the version of `super-linter` used in the `build_validation_develop` workflow to 8.7.0 ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) BUG FIXES: -* Fix airlock Event Grid 503 when a user token omits both `email` and `preferred_username` claims: the derived `AuthenticatedUser` email now defaults to an empty string so the airlock notification payload (which requires a string email) no longer fails validation. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Fix Nexus shared service security: fetch admin password from Key Vault at runtime via managed identity (IMDS) instead of embedding it in the VM Run Command script content. Fix `deploy_nexus_container.sh` short-circuit path to fail loudly if the container does not start. (`sonatype-nexus` 3.10.0) ([#4983](https://github.com/microsoft/AzureTRE/pull/4983)) * Fix UI TypeScript deprecation warning by updating `moduleResolution` to `bundler` in `tsconfig.json`. ([#4968](https://github.com/microsoft/AzureTRE/issues/4968)) * Fix API timeout and name collision failures on workspace creation by checking storage account name availability and improved logging. ([#4946](https://github.com/microsoft/AzureTRE/pull/4946)) From 106343583c040f806a190a83868513ed23464446 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 29 Jul 2026 16:35:54 +0000 Subject: [PATCH 30/32] Consolidate PR CHANGELOG into a single concise entry --- CHANGELOG.md | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 75e078ce71..08fe143a13 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,8 +12,7 @@ ENHANCEMENTS: * Add Windows Server 2025 image support to Guacamole. ([#4890](https://github.com/microsoft/AzureTRE/issues/4890)) * Add support for setting resource processor VMSS SKU via environment variables ([#4936](https://github.com/microsoft/AzureTRE/issues/4936)) * Exclude recovery service vaults from e2e tests ([#4920](https://github.com/microsoft/AzureTRE/issues/4920)) -* Strengthen TRE API authentication: introduce layered `auth/` package with typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, and composable RBAC factories; remove the `AccessService` abstraction that is no longer needed now that Entra ID is the only auth provider. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) -* Add exponential-backoff retry in Event Grid publisher for transient failures (HTTP 429/5xx and network errors); split Graph role-assignment lookup from Event Grid publish in airlock service so auth failures are reported distinctly and underlying exception detail is included in 503 responses. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) +* Strengthen TRE API authentication with a layered `auth/` package (typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, composable RBAC factories), removing the redundant `AccessService` abstraction, and add Event Grid publish resilience (retry on transient 429/5xx and network errors, with Graph and publish failures reported distinctly). ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Add support for formatting UI code via `pre-commit` and fix existing formatting issues. ([#4955](https://github.com/microsoft/AzureTRE/issues/4955)) * Update the version of `super-linter` used in the `build_validation_develop` workflow to 8.7.0 ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) From eeaa8d2cf59d99a279856aa92ed50a55d0de095a Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:20:54 +0000 Subject: [PATCH 31/32] Fix CHANGELOG markdown line-length lint error (MD013) --- CHANGELOG.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 08fe143a13..fa7b5bf932 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,7 +12,7 @@ ENHANCEMENTS: * Add Windows Server 2025 image support to Guacamole. ([#4890](https://github.com/microsoft/AzureTRE/issues/4890)) * Add support for setting resource processor VMSS SKU via environment variables ([#4936](https://github.com/microsoft/AzureTRE/issues/4936)) * Exclude recovery service vaults from e2e tests ([#4920](https://github.com/microsoft/AzureTRE/issues/4920)) -* Strengthen TRE API authentication with a layered `auth/` package (typed exceptions, `PyJWKClient`-backed token validation with issuer checking, immutable `AuthenticatedUser` model, composable RBAC factories), removing the redundant `AccessService` abstraction, and add Event Grid publish resilience (retry on transient 429/5xx and network errors, with Graph and publish failures reported distinctly). ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) +* Strengthen TRE API authentication with a layered `auth/` package (typed exceptions, `PyJWKClient`-backed token validation, immutable `AuthenticatedUser` model, composable RBAC factories), remove the redundant `AccessService` abstraction, and add Event Grid publish resilience with distinct Graph/publish failure reporting. ([#4989](https://github.com/microsoft/AzureTRE/pull/4989)) * Add support for formatting UI code via `pre-commit` and fix existing formatting issues. ([#4955](https://github.com/microsoft/AzureTRE/issues/4955)) * Update the version of `super-linter` used in the `build_validation_develop` workflow to 8.7.0 ([#4957](https://github.com/microsoft/AzureTRE/issues/4957)) From 04b1ea02b69712d7f9a5e824b199e67024d8c359 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:38:37 +0000 Subject: [PATCH 32/32] Return generic Graph 503 message; keep detail in logs; align API version to 0.26.0 --- api_app/_version.py | 2 +- api_app/resources/strings.py | 2 +- api_app/services/airlock.py | 8 ++++---- 3 files changed, 6 insertions(+), 6 deletions(-) diff --git a/api_app/_version.py b/api_app/_version.py index 40ab52fff9..7c4a9591e1 100644 --- a/api_app/_version.py +++ b/api_app/_version.py @@ -1 +1 @@ -__version__ = "0.26.2" +__version__ = "0.26.0" diff --git a/api_app/resources/strings.py b/api_app/resources/strings.py index 7cda05e95a..09432e64a3 100644 --- a/api_app/resources/strings.py +++ b/api_app/resources/strings.py @@ -265,7 +265,7 @@ EVENT_GRID_PUBLISH_FAILED = "Failed to publish Event Grid event: {}" # Graph / role assignments -GRAPH_ROLE_ASSIGNMENT_ERROR = "Failed to fetch workspace role assignments from Microsoft Graph: {}" +GRAPH_ROLE_ASSIGNMENT_ERROR = "Failed to fetch workspace role assignments from Microsoft Graph. See API logs for details." # Workspace creation validation MISSING_REQUIRED_PARAMETERS = "Missing required parameters" diff --git a/api_app/services/airlock.py b/api_app/services/airlock.py index 9a34089b46..36d7a158e1 100644 --- a/api_app/services/airlock.py +++ b/api_app/services/airlock.py @@ -275,9 +275,9 @@ async def save_and_publish_event_airlock_request(airlock_request: AirlockRequest try: access_service = get_aad_service() role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) - except Exception as e: + except Exception: logger.exception("Failed to retrieve workspace role assignments from Microsoft Graph") - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR.format(e)) + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR) if config.ENABLE_AIRLOCK_EMAIL_CHECK: check_email_exists(role_assignment_details) @@ -336,9 +336,9 @@ async def update_and_publish_event_airlock_request( try: access_service = get_aad_service() role_assignment_details = access_service.get_workspace_user_emails_by_role_assignment(workspace) - except Exception as e: + except Exception: logger.exception("Failed to retrieve workspace role assignments from Microsoft Graph") - raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR.format(e)) + raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail=strings.GRAPH_ROLE_ASSIGNMENT_ERROR) try: logger.debug(f"Sending status changed event for airlock request item: {airlock_request.id}")