diff --git a/app/main/validators.py b/app/main/validators.py index 081556986..9d09fbe0c 100644 --- a/app/main/validators.py +++ b/app/main/validators.py @@ -12,7 +12,7 @@ from wtforms import ValidationError from app.main._commonly_used_passwords import commonly_used_passwords from app.models.spreadsheet import Spreadsheet -from app.utils import is_gov_user +from app.utils.user import is_gov_user class CommonlyUsedPassword: diff --git a/app/main/views/add_service.py b/app/main/views/add_service.py index 9802370c2..4b968c46d 100644 --- a/app/main/views/add_service.py +++ b/app/main/views/add_service.py @@ -6,7 +6,7 @@ from app import service_api_client from app.formatters import email_safe from app.main import main from app.main.forms import CreateNhsServiceForm, CreateServiceForm -from app.utils import user_is_gov_user, user_is_logged_in +from app.utils.user import user_is_gov_user, user_is_logged_in def _create_service(service_name, organisation_type, email_from, form): diff --git a/app/main/views/agreement.py b/app/main/views/agreement.py index 6fccac17b..930185924 100644 --- a/app/main/views/agreement.py +++ b/app/main/views/agreement.py @@ -8,7 +8,7 @@ from app.main import main from app.main.forms import AcceptAgreementForm from app.models.organisation import Organisation from app.s3_client.s3_mou_client import get_mou -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions @main.route('/services//agreement') diff --git a/app/main/views/api_keys.py b/app/main/views/api_keys.py index 6df188b1e..840a753b4 100644 --- a/app/main/views/api_keys.py +++ b/app/main/views/api_keys.py @@ -23,7 +23,7 @@ from app.notify_client.api_key_api_client import ( KEY_TYPE_TEAM, KEY_TYPE_TEST, ) -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions dummy_bearer_token = 'bearer_token_set' diff --git a/app/main/views/broadcast.py b/app/main/views/broadcast.py index cae93e45f..1ad578a58 100644 --- a/app/main/views/broadcast.py +++ b/app/main/views/broadcast.py @@ -19,7 +19,8 @@ from app.main.forms import ( SearchByNameForm, ) from app.models.broadcast_message import BroadcastMessage, BroadcastMessages -from app.utils import service_has_permission, user_has_permissions +from app.utils import service_has_permission +from app.utils.user import user_has_permissions def _get_back_link_from_view_broadcast_endpoint(): diff --git a/app/main/views/choose_account.py b/app/main/views/choose_account.py index 3626a66e6..1c26a0b8f 100644 --- a/app/main/views/choose_account.py +++ b/app/main/views/choose_account.py @@ -4,7 +4,8 @@ from flask_login import current_user from app import status_api_client from app.main import main from app.models.organisation import Organisations -from app.utils import PermanentRedirect, user_is_logged_in +from app.utils import PermanentRedirect +from app.utils.user import user_is_logged_in @main.route("/services") diff --git a/app/main/views/conversation.py b/app/main/views/conversation.py index a858688e7..6e7fd257f 100644 --- a/app/main/views/conversation.py +++ b/app/main/views/conversation.py @@ -8,7 +8,7 @@ from app import current_service, notification_api_client, service_api_client from app.main import main from app.main.forms import SearchByNameForm from app.models.template_list import TemplateList -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions @main.route("/services//conversation/") diff --git a/app/main/views/dashboard.py b/app/main/views/dashboard.py index ab940050f..33101ec5a 100644 --- a/app/main/views/dashboard.py +++ b/app/main/views/dashboard.py @@ -35,8 +35,8 @@ from app.utils import ( generate_previous_dict, get_current_financial_year, service_has_permission, - user_has_permissions, ) +from app.utils.user import user_has_permissions @main.route("/services//dashboard") diff --git a/app/main/views/email_branding.py b/app/main/views/email_branding.py index c32227242..c3b5b10ec 100644 --- a/app/main/views/email_branding.py +++ b/app/main/views/email_branding.py @@ -11,7 +11,8 @@ from app.s3_client.s3_logo_client import ( persist_logo, upload_email_logo, ) -from app.utils import get_logo_cdn_domain, user_is_platform_admin +from app.utils import get_logo_cdn_domain +from app.utils.user import user_is_platform_admin @main.route("/email-branding", methods=['GET', 'POST']) diff --git a/app/main/views/find_services.py b/app/main/views/find_services.py index ff841e488..8156c7f3b 100644 --- a/app/main/views/find_services.py +++ b/app/main/views/find_services.py @@ -6,7 +6,7 @@ from flask import redirect, render_template, url_for from app import service_api_client from app.main import main from app.main.forms import SearchByNameForm -from app.utils import user_is_platform_admin +from app.utils.user import user_is_platform_admin @main.route("/find-services-by-name", methods=['GET', 'POST']) diff --git a/app/main/views/find_users.py b/app/main/views/find_users.py index 10850cf9f..00c8bdac2 100644 --- a/app/main/views/find_users.py +++ b/app/main/views/find_users.py @@ -7,7 +7,7 @@ from app.event_handlers import create_archive_user_event from app.main import main from app.main.forms import SearchUsersByEmailForm from app.models.user import User -from app.utils import user_is_platform_admin +from app.utils.user import user_is_platform_admin @main.route("/find-users-by-email", methods=['GET', 'POST']) diff --git a/app/main/views/history.py b/app/main/views/history.py index 1be4f1386..6c0603a7c 100644 --- a/app/main/views/history.py +++ b/app/main/views/history.py @@ -6,7 +6,7 @@ from flask import render_template, request from app import current_service, format_date_numeric from app.main import main from app.models.event import APIKeyEvent, APIKeyEvents, ServiceEvents -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions @main.route("/services//history") diff --git a/app/main/views/inbound_number.py b/app/main/views/inbound_number.py index 3007db9e5..c56620b33 100644 --- a/app/main/views/inbound_number.py +++ b/app/main/views/inbound_number.py @@ -2,7 +2,7 @@ from flask import render_template from app import inbound_number_client from app.main import main -from app.utils import user_is_platform_admin +from app.utils.user import user_is_platform_admin @main.route('/inbound-sms-admin', methods=['GET', 'POST']) diff --git a/app/main/views/jobs.py b/app/main/views/jobs.py index d6009c8b8..8a63828e4 100644 --- a/app/main/views/jobs.py +++ b/app/main/views/jobs.py @@ -42,8 +42,8 @@ from app.utils import ( parse_filter_args, printing_today_or_tomorrow, set_status_filters, - user_has_permissions, ) +from app.utils.user import user_has_permissions @main.route("/services//jobs") diff --git a/app/main/views/letter_branding.py b/app/main/views/letter_branding.py index 76c694e30..571581ec9 100644 --- a/app/main/views/letter_branding.py +++ b/app/main/views/letter_branding.py @@ -25,7 +25,8 @@ from app.s3_client.s3_logo_client import ( persist_logo, upload_letter_temp_logo, ) -from app.utils import get_logo_cdn_domain, user_is_platform_admin +from app.utils import get_logo_cdn_domain +from app.utils.user import user_is_platform_admin @main.route("/letter-branding", methods=['GET']) diff --git a/app/main/views/manage_users.py b/app/main/views/manage_users.py index 8140288bc..3a75f1359 100644 --- a/app/main/views/manage_users.py +++ b/app/main/views/manage_users.py @@ -30,7 +30,7 @@ from app.main.forms import ( ) from app.models.roles_and_permissions import broadcast_permissions, permissions from app.models.user import InvitedUser, User -from app.utils import is_gov_user, user_has_permissions +from app.utils.user import is_gov_user, user_has_permissions @main.route("/services//users") diff --git a/app/main/views/notifications.py b/app/main/views/notifications.py index aac38a026..83f7221e9 100644 --- a/app/main/views/notifications.py +++ b/app/main/views/notifications.py @@ -44,8 +44,8 @@ from app.utils import ( get_template, parse_filter_args, set_status_filters, - user_has_permissions, ) +from app.utils.user import user_has_permissions @main.route("/services//notification/") diff --git a/app/main/views/organisations.py b/app/main/views/organisations.py index 0d8ae2a5c..12405cbd8 100644 --- a/app/main/views/organisations.py +++ b/app/main/views/organisations.py @@ -43,7 +43,7 @@ from app.main.views.dashboard import ( from app.main.views.service_settings import get_branding_as_value_and_label from app.models.organisation import Organisation, Organisations from app.models.user import InvitedOrgUser, User -from app.utils import user_has_permissions, user_is_platform_admin +from app.utils.user import user_has_permissions, user_is_platform_admin @main.route("/organisations", methods=['GET']) diff --git a/app/main/views/platform_admin.py b/app/main/views/platform_admin.py index 0c7a319a0..544da3a58 100644 --- a/app/main/views/platform_admin.py +++ b/app/main/views/platform_admin.py @@ -32,8 +32,8 @@ from app.utils import ( generate_next_dict, generate_previous_dict, get_page_from_request, - user_is_platform_admin, ) +from app.utils.user import user_is_platform_admin COMPLAINT_THRESHOLD = 0.02 FAILURE_THRESHOLD = 3 diff --git a/app/main/views/providers.py b/app/main/views/providers.py index 0484f4463..977936cc9 100644 --- a/app/main/views/providers.py +++ b/app/main/views/providers.py @@ -8,7 +8,7 @@ from werkzeug.utils import redirect from app import format_date_numeric, provider_client from app.main import main from app.main.forms import ProviderForm, ProviderRatioForm -from app.utils import user_is_platform_admin +from app.utils.user import user_is_platform_admin PROVIDER_PRIORITY_MEANING_SWITCHOVER = datetime(2019, 11, 29, 11, 0).isoformat() diff --git a/app/main/views/returned_letters.py b/app/main/views/returned_letters.py index d4861733a..ba222a066 100644 --- a/app/main/views/returned_letters.py +++ b/app/main/views/returned_letters.py @@ -5,7 +5,7 @@ from flask import render_template from app import current_service, service_api_client from app.main import main from app.models.spreadsheet import Spreadsheet -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions @main.route("/services//returned-letters") diff --git a/app/main/views/send.py b/app/main/views/send.py index 50ab717a7..ddd8ce73a 100644 --- a/app/main/views/send.py +++ b/app/main/views/send.py @@ -57,8 +57,8 @@ from app.utils import ( get_template, should_skip_template_page, unicode_truncate, - user_has_permissions, ) +from app.utils.user import user_has_permissions letter_address_columns = [ column.replace('_', ' ') diff --git a/app/main/views/service_settings.py b/app/main/views/service_settings.py index 8006e8492..187f7aaf1 100644 --- a/app/main/views/service_settings.py +++ b/app/main/views/service_settings.py @@ -62,10 +62,8 @@ from app.main.forms import ( SetLetterBranding, SMSPrefixForm, ) -from app.utils import ( - DELIVERED_STATUSES, - FAILURE_STATUSES, - SENDING_STATUSES, +from app.utils import DELIVERED_STATUSES, FAILURE_STATUSES, SENDING_STATUSES +from app.utils.user import ( user_has_permissions, user_is_gov_user, user_is_platform_admin, diff --git a/app/main/views/templates.py b/app/main/views/templates.py index 2e69e29bb..bb01de36b 100644 --- a/app/main/views/templates.py +++ b/app/main/views/templates.py @@ -43,9 +43,8 @@ from app.utils import ( NOTIFICATION_TYPES, get_template, should_skip_template_page, - user_has_permissions, - user_is_platform_admin, ) +from app.utils.user import user_has_permissions, user_is_platform_admin form_objects = { 'email': EmailTemplateForm, diff --git a/app/main/views/tour.py b/app/main/views/tour.py index d6223a0e8..9317367e3 100644 --- a/app/main/views/tour.py +++ b/app/main/views/tour.py @@ -9,7 +9,8 @@ from app.main.views.send import ( get_placeholder_form_instance, get_recipient_and_placeholders_from_session, ) -from app.utils import get_template, user_has_permissions +from app.utils import get_template +from app.utils.user import user_has_permissions @main.route("/services//tour/") diff --git a/app/main/views/uploads.py b/app/main/views/uploads.py index b2609bde6..4f5d49015 100644 --- a/app/main/views/uploads.py +++ b/app/main/views/uploads.py @@ -56,8 +56,8 @@ from app.utils import ( get_sample_template, get_template, unicode_truncate, - user_has_permissions, ) +from app.utils.user import user_has_permissions MAX_FILE_UPLOAD_SIZE = 2 * 1024 * 1024 # 2MB diff --git a/app/main/views/user_profile.py b/app/main/views/user_profile.py index f6c701874..abae91c43 100644 --- a/app/main/views/user_profile.py +++ b/app/main/views/user_profile.py @@ -27,7 +27,7 @@ from app.main.forms import ( TwoFactorForm, ) from app.models.user import User -from app.utils import ( +from app.utils.user import ( user_is_gov_user, user_is_logged_in, user_is_platform_admin, diff --git a/app/main/views/webauthn_credentials.py b/app/main/views/webauthn_credentials.py index a0e9ddb53..1f891ead9 100644 --- a/app/main/views/webauthn_credentials.py +++ b/app/main/views/webauthn_credentials.py @@ -10,11 +10,8 @@ from app.main.views.two_factor import log_in_user from app.models.user import User from app.models.webauthn_credential import RegistrationError, WebAuthnCredential from app.notify_client.user_api_client import user_api_client -from app.utils import ( - is_less_than_days_ago, - redirect_to_sign_in, - user_is_platform_admin, -) +from app.utils import is_less_than_days_ago, redirect_to_sign_in +from app.utils.user import user_is_platform_admin @main.route('/webauthn/register') diff --git a/app/models/user.py b/app/models/user.py index 584938513..f0ff97c02 100644 --- a/app/models/user.py +++ b/app/models/user.py @@ -15,7 +15,7 @@ from app.notify_client import InviteTokenError from app.notify_client.invite_api_client import invite_api_client from app.notify_client.org_invite_api_client import org_invite_api_client from app.notify_client.user_api_client import user_api_client -from app.utils import is_gov_user +from app.utils.user import is_gov_user def _get_service_id_from_view_args(): diff --git a/app/utils/__init__.py b/app/utils/__init__.py index 59f52d585..ddaebac17 100644 --- a/app/utils/__init__.py +++ b/app/utils/__init__.py @@ -1,4 +1,3 @@ -import os from datetime import datetime, timedelta from functools import wraps from itertools import chain @@ -16,7 +15,7 @@ from flask import ( session, url_for, ) -from flask_login import current_user, login_required +from flask_login import current_user from notifications_utils.field import Field from notifications_utils.formatters import unescaped_formatted_list from notifications_utils.letter_timings import letter_can_be_cancelled @@ -39,7 +38,6 @@ from werkzeug.datastructures import MultiDict from werkzeug.routing import RequestRedirect from app.models.spreadsheet import Spreadsheet -from app.notify_client.organisations_api_client import organisations_client SENDING_STATUSES = ['created', 'pending', 'sending', 'pending-virus-check'] DELIVERED_STATUSES = ['delivered', 'sent', 'returned-letter'] @@ -50,28 +48,6 @@ REQUESTED_STATUSES = SENDING_STATUSES + DELIVERED_STATUSES + FAILURE_STATUSES NOTIFICATION_TYPES = ["sms", "email", "letter", "broadcast"] -with open('{}/email_domains.txt'.format( - os.path.dirname(os.path.realpath(__file__)) -)) as email_domains: - GOVERNMENT_EMAIL_DOMAIN_NAMES = [line.strip() for line in email_domains] - - -user_is_logged_in = login_required - - -def user_has_permissions(*permissions, **permission_kwargs): - def wrap(func): - @wraps(func) - def wrap_func(*args, **kwargs): - if not current_user.is_authenticated: - return current_app.login_manager.unauthorized() - if not current_user.has_permissions(*permissions, **permission_kwargs): - abort(403) - return func(*args, **kwargs) - return wrap_func - return wrap - - def service_has_permission(permission): from app import current_service @@ -86,28 +62,6 @@ def service_has_permission(permission): return wrap -def user_is_gov_user(f): - @wraps(f) - def wrapped(*args, **kwargs): - if not current_user.is_authenticated: - return current_app.login_manager.unauthorized() - if not current_user.is_gov_user: - abort(403) - return f(*args, **kwargs) - return wrapped - - -def user_is_platform_admin(f): - @wraps(f) - def wrapped(*args, **kwargs): - if not current_user.is_authenticated: - return current_app.login_manager.unauthorized() - if not current_user.platform_admin: - abort(403) - return f(*args, **kwargs) - return wrapped - - def redirect_to_sign_in(f): @wraps(f) def wrapped(*args, **kwargs): @@ -264,24 +218,6 @@ def get_help_argument(): return request.args.get('help') if request.args.get('help') in ('1', '2', '3') else None -def email_address_ends_with(email_address, known_domains): - return any( - email_address.lower().endswith(( - "@{}".format(known), - ".{}".format(known), - )) - for known in known_domains - ) - - -def is_gov_user(email_address): - return email_address_ends_with( - email_address, GOVERNMENT_EMAIL_DOMAIN_NAMES - ) or email_address_ends_with( - email_address, organisations_client.get_domains() - ) - - def get_template( template, service, diff --git a/app/utils/user.py b/app/utils/user.py new file mode 100644 index 000000000..577abb305 --- /dev/null +++ b/app/utils/user.py @@ -0,0 +1,68 @@ +import os +from functools import wraps + +from flask import abort, current_app +from flask_login import current_user, login_required + +from app.notify_client.organisations_api_client import organisations_client + +user_is_logged_in = login_required + + +with open('{}/email_domains.txt'.format( + os.path.dirname(os.path.realpath(__file__)) +)) as email_domains: + GOVERNMENT_EMAIL_DOMAIN_NAMES = [line.strip() for line in email_domains] + + +def user_has_permissions(*permissions, **permission_kwargs): + def wrap(func): + @wraps(func) + def wrap_func(*args, **kwargs): + if not current_user.is_authenticated: + return current_app.login_manager.unauthorized() + if not current_user.has_permissions(*permissions, **permission_kwargs): + abort(403) + return func(*args, **kwargs) + return wrap_func + return wrap + + +def user_is_gov_user(f): + @wraps(f) + def wrapped(*args, **kwargs): + if not current_user.is_authenticated: + return current_app.login_manager.unauthorized() + if not current_user.is_gov_user: + abort(403) + return f(*args, **kwargs) + return wrapped + + +def user_is_platform_admin(f): + @wraps(f) + def wrapped(*args, **kwargs): + if not current_user.is_authenticated: + return current_app.login_manager.unauthorized() + if not current_user.platform_admin: + abort(403) + return f(*args, **kwargs) + return wrapped + + +def is_gov_user(email_address): + return _email_address_ends_with( + email_address, GOVERNMENT_EMAIL_DOMAIN_NAMES + ) or _email_address_ends_with( + email_address, organisations_client.get_domains() + ) + + +def _email_address_ends_with(email_address, known_domains): + return any( + email_address.lower().endswith(( + "@{}".format(known), + ".{}".format(known), + )) + for known in known_domains + ) diff --git a/tests/app/main/test_permissions.py b/tests/app/main/test_permissions.py index 727dcdcd3..f8af16135 100644 --- a/tests/app/main/test_permissions.py +++ b/tests/app/main/test_permissions.py @@ -11,7 +11,7 @@ from app.models.roles_and_permissions import ( translate_permissions_from_admin_roles_to_db, translate_permissions_from_db_to_admin_roles, ) -from app.utils import user_has_permissions +from app.utils.user import user_has_permissions from tests import service_json from tests.conftest import ( ORGANISATION_ID, diff --git a/tests/app/main/views/test_add_service.py b/tests/app/main/views/test_add_service.py index ced1d3270..82b943702 100644 --- a/tests/app/main/views/test_add_service.py +++ b/tests/app/main/views/test_add_service.py @@ -3,7 +3,7 @@ from flask import session, url_for from freezegun import freeze_time from notifications_python_client.errors import HTTPError -from app.utils import is_gov_user +from app.utils.user import is_gov_user from tests import organisation_json from tests.conftest import normalize_spaces diff --git a/tests/app/main/views/test_manage_users.py b/tests/app/main/views/test_manage_users.py index 44fbddb0f..513b26780 100644 --- a/tests/app/main/views/test_manage_users.py +++ b/tests/app/main/views/test_manage_users.py @@ -5,7 +5,7 @@ import pytest from flask import url_for import app -from app.utils import is_gov_user +from app.utils.user import is_gov_user from tests.conftest import ( ORGANISATION_ID, ORGANISATION_TWO_ID,