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/CHANGELOG.md b/CHANGELOG.md index d14373da30..40ba3dd987 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. ([#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/_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" 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..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 services.authentication import get_current_admin_user, get_current_workspace_owner_or_tre_admin +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(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_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(get_current_workspace_owner_or_tre_admin)], + 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/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/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/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 b90a6b07ea..896d5f6512 100644 --- a/api_app/api/routes/workspace_users.py +++ b/api_app/api/routes/workspace_users.py @@ -2,35 +2,35 @@ 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 +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(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_or_tre_admin)]) @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 df912744ab..04ea2d6513 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 @@ -23,14 +23,12 @@ 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.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_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 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_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 @@ -39,10 +37,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_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)]) def validate_user_has_valid_role_for_user_resource(user, user_resource): @@ -57,30 +55,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_access_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) @@ -97,8 +93,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) @@ -131,8 +127,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: @@ -159,8 +155,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, @@ -177,8 +173,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, @@ -200,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(get_current_workspace_owner_or_researcher_user_or_airlock_manager_or_tre_admin)) -> 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) @@ -211,42 +207,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_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(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_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(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_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(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_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)) # 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) @@ -290,8 +286,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), 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: @@ -316,8 +312,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, @@ -334,8 +330,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, @@ -352,17 +348,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)) @@ -372,7 +368,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) @@ -394,7 +390,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: @@ -412,7 +408,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: @@ -464,7 +460,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)), @@ -494,7 +490,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)), @@ -528,7 +524,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, @@ -550,7 +546,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)) @@ -559,13 +555,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/__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..c34f7e9665 --- /dev/null +++ b/api_app/auth/dependencies.py @@ -0,0 +1,56 @@ +from typing import Optional + +from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer + +from auth.exceptions import AuthError, TokenExpired, TokenSignatureInvalid +from auth.models import AuthenticatedUser +from auth.registry import get_core_validator +from resources import strings +from services.logging import logger + +# 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: + if isinstance(exc, TokenExpired): + 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=detail, + headers={"WWW-Authenticate": "Bearer"}, + ) + + +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(require_bearer_credentials), +) -> 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) 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..209c3126be --- /dev/null +++ b/api_app/auth/models.py @@ -0,0 +1,44 @@ +from enum import StrEnum +from typing import Optional, Tuple, Union + +from pydantic import BaseModel, Field + + +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 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: Tuple[str, ...] = Field(default_factory=tuple) + 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.value in self.roles diff --git a/api_app/auth/rbac.py b/api_app/auth/rbac.py new file mode 100644 index 0000000000..5329d66ba9 --- /dev/null +++ b/api_app/auth/rbac.py @@ -0,0 +1,165 @@ +from typing import Callable, Union + +from fastapi import Depends, HTTPException, status +from fastapi.security import HTTPAuthorizationCredentials + +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 +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: + """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 + + _check._role_names = role_names + return _check + + +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 + (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) + allowed_values = role_values | ({TRERole.Admin.value} if allow_tre_admin else frozenset()) + role_names = [r.value for r in roles] + + async def _check( + credentials: HTTPAuthorizationCredentials = Depends(require_bearer_credentials), + 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 — 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; 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 user.is_tre_admin(): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=strings.INVALID_TOKEN, + headers={"WWW-Authenticate": "Bearer"}, + ) + return user + + _check._role_names = role_names + 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) + +# 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) +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, +) + +# 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/auth/registry.py b/api_app/auth/registry.py new file mode 100644 index 0000000000..377adef4f4 --- /dev/null +++ b/api_app/auth/registry.py @@ -0,0 +1,64 @@ +from functools import lru_cache + +from jwt import PyJWKClient + +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 _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.""" + return TokenValidator( + TokenValidatorConfig( + jwks_uri=_jwks_uri(), + audience=config.API_AUDIENCE, + issuer=_issuer(), + is_workspace_token=False, + ), + jwks_client=_shared_jwks_client(), + ) + + +@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 client 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, + ), + jwks_client=_shared_jwks_client(), + ) diff --git a/api_app/auth/token_validator.py b/api_app/auth/token_validator.py new file mode 100644 index 0000000000..d317b0119c --- /dev/null +++ b/api_app/auth/token_validator.py @@ -0,0 +1,95 @@ +from dataclasses import dataclass +from typing import Optional + +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. + + 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, + jwks_client: Optional[PyJWKClient] = None, + ) -> None: + self._config = config + 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`. + + 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, + # 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: + 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 + + 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") 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 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..6fcf73eaea 100644 --- a/api_app/services/aad_authentication.py +++ b/api_app/services/aad_authentication.py @@ -1,203 +1,52 @@ -import base64 from collections import defaultdict from enum import Enum -from typing import List, Optional -import jwt -import requests +from typing import List -from fastapi import Request, HTTPException, status +import requests from msal import ConfidentialClientApplication +from semantic_version import Version -from services.access_service import AccessService, AuthConfigValidationError, UserRoleAssignmentError 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 PrincipalType(Enum): - User = "User" - Group = "Group" - ServicePrincipal = "ServicePrincipal" - +class AuthConfigValidationError(Exception): + """Raised when the input auth information is invalid.""" -class AzureADAuthorization(AccessService): - _jwt_keys: dict = {} - - require_one_of_roles = None - aad_instance = 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(AzureADAuthorization, self).__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 - ) - self.require_one_of_roles = require_one_of_roles - - async def __call__(self, request: Request) -> User: - - token: str = await super(AzureADAuthorization, self).__call__(request) - - 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: - # 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): - 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"}) - - 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"}) +class UserRoleAssignmentError(Exception): + """Raised when a user role assignment fails.""" - return 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) - 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 - 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'] +class PrincipalType(Enum): + User = "User" + Group = "Group" + ServicePrincipal = "ServicePrincipal" - 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) +class AzureADAuthorization: + """Service wrapper for Microsoft Graph calls related to workspace auth. - 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) + 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. + """ - @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 + WORKSPACE_ROLES_DICT = { + 'WorkspaceOwner': 'app_role_id_workspace_owner', + 'WorkspaceResearcher': 'app_role_id_workspace_researcher', + 'AirlockManager': 'app_role_id_workspace_airlock_manager', + } - @staticmethod - def _ensure_b64padding(key: str) -> str: - """ - 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] - - # 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..9dd2378cbe 100644 --- a/api_app/services/authentication.py +++ b/api_app/services/authentication.py @@ -1,57 +1,16 @@ -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 + 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) - - -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"]) +def get_aad_service() -> AzureADAuthorization: + """Return an :class:`AzureADAuthorization` instance for Graph API calls.""" + return AzureADAuthorization() 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..1b69b2b1a0 --- /dev/null +++ b/api_app/tests_ma/auth/test_rbac.py @@ -0,0 +1,288 @@ +"""Tests for auth.rbac role-checking dependencies.""" +import pytest +from unittest.mock import MagicMock, patch + +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 test_allows_user_with_required_role(self): + admin = _make_user(roles=["TREAdmin"]) + dep = require_roles(TRERole.Admin) + + import asyncio + + async def _run(): + result = await dep(user=admin) + assert result.id == "uid" + assert "TREAdmin" in result.roles + + asyncio.run(_run()) + + def test_raises_403_when_user_lacks_role(self): + from fastapi import HTTPException + + dep = require_roles(TRERole.Admin) + + import asyncio + + async def _run(): + user_with_no_roles = _make_user(roles=["TREUser"]) + with pytest.raises(HTTPException) as exc_info: + await dep(user=user_with_no_roles) + assert exc_info.value.status_code == 403 + + asyncio.run(_run()) + + 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.run(_run()) + + +# --------------------------------------------------------------------------- +# require_workspace_roles +# --------------------------------------------------------------------------- + + +class TestRequireWorkspaceRoles: + 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") + 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, + ) + 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): + """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, 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. + 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_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.run(_run()) + + def test_workspace_owner_passes(self): + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, ws_validator, owner = self._make_fake_deps(["WorkspaceOwner"]) + + import asyncio + + async def _run(): + with patch('auth.rbac.get_workspace_validator', return_value=ws_validator): + result = await dep(credentials=fake_creds, workspace=fake_workspace) + assert result.id == "uid" + + asyncio.run(_run()) + + def test_raises_403_for_user_without_workspace_role(self): + from fastapi import HTTPException + + dep = require_workspace_roles(WorkspaceAccessRole.Owner) + fake_creds, fake_workspace, ws_validator, researcher = self._make_fake_deps(["WorkspaceResearcher"]) + + import asyncio + + async def _run(): + 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 + + 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. + + 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, allow_tre_admin=True) + 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.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. + + 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 + + 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) + ws_validator_mock = MagicMock() + ws_validator_mock.validate.side_effect = _TokenInvalid("wrong audience") + + # 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=["WorkspaceOwner"]) + 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 == 401 + + 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 + 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.run(_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.run(_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.run(_run()) + + +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] + + 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 new file mode 100644 index 0000000000..5776fc7f76 --- /dev/null +++ b/api_app/tests_ma/auth/test_token_validator.py @@ -0,0 +1,260 @@ +"""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): + 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): + 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] + + # 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 + + 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_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", + "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_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.kwargs.get("algorithms") + 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) + 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..7ad45f341c 100644 --- a/api_app/tests_ma/test_api/conftest.py +++ b/api_app/tests_ma/test_api/conftest.py @@ -17,9 +17,21 @@ 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 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") @@ -78,9 +90,15 @@ 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 + 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'), 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..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 @@ -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,13 @@ 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 +30,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..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 @@ -6,10 +6,9 @@ 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_or_tre_admin from models.domain.workspace import Workspace from resources import strings @@ -47,13 +46,11 @@ 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_or_tre_admin] = admin_user + app.dependency_overrides[require_tre_user_or_admin] = 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 9c12371022..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 @@ -23,12 +23,14 @@ 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, \ + require_workspace_owner_or_tre_admin, \ + require_workspace_owner_or_researcher_or_airlock_manager_or_tre_admin from azure.cosmos.exceptions import CosmosAccessConditionFailedError @@ -249,8 +251,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 +292,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_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_or_tre_admin, None) # [GET] /workspaces/{workspace_id} @patch("api.dependencies.workspaces.WorkspaceRepository.get_workspace_by_id", side_effect=EntityDoesNotExist) @@ -341,13 +352,17 @@ 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_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 + yield + app.dependency_overrides = {} # [GET] /workspaces @patch("api.routes.workspaces.WorkspaceRepository.get_active_workspaces") @@ -723,10 +738,12 @@ 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_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 app.dependency_overrides = {} @@ -1371,9 +1388,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_or_airlock_manager_or_tre_admin] = researcher_user + app.dependency_overrides[require_workspace_owner_or_researcher] = researcher_user yield app.dependency_overrides = {} 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"