Fix for security hole with setting session['user_id'] before second factor of authentication has been authorised.

This commit is contained in:
Nicholas Staples
2016-01-07 12:43:10 +00:00
parent 10c2978f85
commit 7001d8261d
17 changed files with 162 additions and 119 deletions

View File

@@ -16,7 +16,9 @@ def add_code(user_id, code, code_type):
return code.id return code.id
def get_codes(user_id, code_type): def get_codes(user_id, code_type=None):
if not code_type:
return VerifyCodes.query.filter_by(user_id=user_id, code_used=False).all()
return VerifyCodes.query.filter_by(user_id=user_id, code_type=code_type, code_used=False).all() return VerifyCodes.query.filter_by(user_id=user_id, code_type=code_type, code_used=False).all()

View File

@@ -1,12 +1,11 @@
from datetime import datetime from datetime import datetime
from flask import session
from flask_wtf import Form from flask_wtf import Form
from wtforms import StringField, PasswordField, ValidationError from wtforms import StringField, PasswordField, ValidationError
from wtforms.validators import DataRequired, Email, Length, Regexp from wtforms.validators import DataRequired, Email, Length, Regexp
from app.main.dao import verify_codes_dao from app.main.dao import verify_codes_dao
from app.main.encryption import check_hash from app.main.encryption import check_hash
from app.main.validators import Blacklist from app.main.validators import Blacklist, ValidateUserCodes
class LoginForm(Form): class LoginForm(Form):
@@ -62,26 +61,40 @@ class RegisterUserForm(Form):
class TwoFactorForm(Form): class TwoFactorForm(Form):
sms_code = StringField('sms code', validators=[DataRequired(message='Enter verification code'),
Regexp(regex=verify_code, message='Code must be 5 digits')])
def validate_sms_code(self, a): def __init__(self, user_codes, *args, **kwargs):
return validate_codes(self.sms_code, 'sms') '''
Keyword arguments:
user_codes -- List of user code objects which have the fields
(code_type, expiry_datetime, code)
'''
self.user_codes = user_codes
super(TwoFactorForm, self).__init__(*args, **kwargs)
sms_code = StringField('sms code', validators=[DataRequired(message='Enter verification code'),
Regexp(regex=verify_code, message='Code must be 5 digits'),
ValidateUserCodes(code_type='sms')])
class VerifyForm(Form): class VerifyForm(Form):
def __init__(self, user_codes, *args, **kwargs):
'''
Keyword arguments:
user_codes -- List of user code objects which have the fields
(code_type, expiry_datetime, code)
'''
self.user_codes = user_codes
super(VerifyForm, self).__init__(*args, **kwargs)
sms_code = StringField("Text message confirmation code", sms_code = StringField("Text message confirmation code",
validators=[DataRequired(message='SMS code can not be empty'), validators=[DataRequired(message='SMS code can not be empty'),
Regexp(regex=verify_code, message='Code must be 5 digits')]) Regexp(regex=verify_code, message='Code must be 5 digits'),
ValidateUserCodes(code_type='sms')])
email_code = StringField("Email confirmation code", email_code = StringField("Email confirmation code",
validators=[DataRequired(message='Email code can not be empty'), validators=[DataRequired(message='Email code can not be empty'),
Regexp(regex=verify_code, message='Code must be 5 digits')]) Regexp(regex=verify_code, message='Code must be 5 digits'),
ValidateUserCodes(code_type='email')])
def validate_email_code(self, a):
return validate_codes(self.email_code, 'email')
def validate_sms_code(self, a):
return validate_codes(self.sms_code, 'sms')
class EmailNotReceivedForm(Form): class EmailNotReceivedForm(Form):
@@ -110,19 +123,3 @@ class AddServiceForm(Form):
def validate_service_name(self, a): def validate_service_name(self, a):
if self.service_name.data in self.service_names: if self.service_name.data in self.service_names:
raise ValidationError('Service name already exists') raise ValidationError('Service name already exists')
def validate_codes(field, code_type):
codes = verify_codes_dao.get_codes(user_id=session['user_id'], code_type=code_type)
# TODO need to remove this manual logging.
print('validate_codes for user_id: {} are {}'.format(session['user_id'], codes))
if not [code for code in codes if validate_code(field, code)]:
raise ValidationError('Code does not match')
def validate_code(field, code):
if field.data and check_hash(field.data, code.code):
if code.expiry_datetime <= datetime.now():
raise ValidationError('Code has expired')
else:
return code.code

View File

@@ -1,4 +1,6 @@
from wtforms import ValidationError from wtforms import ValidationError
from datetime import datetime
from app.main.encryption import check_hash
class Blacklist(object): class Blacklist(object):
@@ -10,3 +12,29 @@ class Blacklist(object):
def __call__(self, form, field): def __call__(self, form, field):
if field.data in ['password1234', 'passw0rd1234']: if field.data in ['password1234', 'passw0rd1234']:
raise ValidationError(self.message) raise ValidationError(self.message)
class ValidateUserCodes(object):
def __init__(self,
expiry_msg='Code has expired',
invalid_msg='Code does not match',
code_type=None):
self.expiry_msg = expiry_msg
self.invalid_msg = invalid_msg
self.code_type = code_type
def __call__(self, form, field):
# TODO would be great to do this sql query but
# not couple those parts of the code.
user_codes = getattr(form, 'user_codes', [])
valid_code = False
for code in user_codes:
if check_hash(field.data, code.code) and self.code_type == code.code_type:
if code.expiry_datetime <= datetime.now():
raise ValidationError(self.expiry_msg)
else:
# Valid code
valid_code = True
break
if not valid_code:
raise ValidationError(self.invalid_msg)

View File

@@ -10,7 +10,7 @@ from app.main.views import send_sms_code, send_email_code
@main.route('/email-not-received', methods=['GET', 'POST']) @main.route('/email-not-received', methods=['GET', 'POST'])
def check_and_resend_email_code(): def check_and_resend_email_code():
# TODO there needs to be a way to regenerate a session id # TODO there needs to be a way to regenerate a session id
user = users_dao.get_user_by_id(session['user_id']) user = users_dao.get_user_by_email(session['user_email'])
form = EmailNotReceivedForm(email_address=user.email_address) form = EmailNotReceivedForm(email_address=user.email_address)
if form.validate_on_submit(): if form.validate_on_submit():
users_dao.update_email_address(id=user.id, email_address=form.email_address.data) users_dao.update_email_address(id=user.id, email_address=form.email_address.data)
@@ -22,7 +22,7 @@ def check_and_resend_email_code():
@main.route('/text-not-received', methods=['GET', 'POST']) @main.route('/text-not-received', methods=['GET', 'POST'])
def check_and_resend_text_code(): def check_and_resend_text_code():
# TODO there needs to be a way to regenerate a session id # TODO there needs to be a way to regenerate a session id
user = users_dao.get_user_by_id(session['user_id']) user = users_dao.get_user_by_email(session['user_email'])
form = TextNotReceivedForm(mobile_number=user.mobile_number) form = TextNotReceivedForm(mobile_number=user.mobile_number)
if form.validate_on_submit(): if form.validate_on_submit():
users_dao.update_mobile_number(id=user.id, mobile_number=form.mobile_number.data) users_dao.update_mobile_number(id=user.id, mobile_number=form.mobile_number.data)
@@ -39,6 +39,6 @@ def verification_code_not_received():
@main.route('/send-new-code', methods=['GET']) @main.route('/send-new-code', methods=['GET'])
def check_and_resend_verification_code(): def check_and_resend_verification_code():
# TODO there needs to be a way to generate a new session id # TODO there needs to be a way to generate a new session id
user = users_dao.get_user_by_id(session['user_id']) user = users_dao.get_user_by_email(session['user_email'])
send_sms_code(user.id, user.mobile_number) send_sms_code(user.id, user.mobile_number)
return redirect(url_for('main.two_factor')) return redirect(url_for('main.two_factor'))

View File

@@ -42,7 +42,7 @@ def process_register():
send_sms_code(user_id=user.id, mobile_number=form.mobile_number.data) send_sms_code(user_id=user.id, mobile_number=form.mobile_number.data)
send_email_code(user_id=user.id, email=form.email_address.data) send_email_code(user_id=user.id, email=form.email_address.data)
session['expiry_date'] = str(datetime.now() + timedelta(hours=1)) session['expiry_date'] = str(datetime.now() + timedelta(hours=1))
session['user_id'] = user.id session['user_email'] = user.email_address
return redirect('/verify') return redirect('/verify')
return render_template('views/register.html', form=form) return render_template('views/register.html', form=form)

View File

@@ -18,7 +18,7 @@ def sign_in():
if user: if user:
if not user.is_locked() and user.is_active() and check_hash(form.password.data, user.password): if not user.is_locked() and user.is_active() and check_hash(form.password.data, user.password):
send_sms_code(user.id, user.mobile_number) send_sms_code(user.id, user.mobile_number)
session['user_id'] = user.id session['user_email'] = user.email_address
return redirect(url_for('.two_factor')) return redirect(url_for('.two_factor'))
else: else:
users_dao.increment_failed_login_count(user.id) users_dao.increment_failed_login_count(user.id)

View File

@@ -11,10 +11,12 @@ from app.main.forms import TwoFactorForm
@main.route('/two-factor', methods=['GET', 'POST']) @main.route('/two-factor', methods=['GET', 'POST'])
def two_factor(): def two_factor():
form = TwoFactorForm() # TODO handle user_email not in session
user = users_dao.get_user_by_email(session['user_email'])
codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(codes)
if form.validate_on_submit(): if form.validate_on_submit():
user = users_dao.get_user_by_id(session['user_id'])
verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='sms') verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='sms')
login_user(user) login_user(user)
return redirect(url_for('.dashboard')) return redirect(url_for('.dashboard'))

View File

@@ -12,8 +12,9 @@ from app.main.forms import VerifyForm
def verify(): def verify():
# TODO there needs to be a way to regenerate a session id # TODO there needs to be a way to regenerate a session id
# or handle gracefully. # or handle gracefully.
user = users_dao.get_user_by_id(session['user_id']) user = users_dao.get_user_by_email(session['user_email'])
form = VerifyForm() codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(codes)
if form.validate_on_submit(): if form.validate_on_submit():
verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='email') verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='email')
verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='sms') verify_codes_dao.use_code_for_user_and_type(user_id=user.id, code_type='sms')

View File

@@ -26,16 +26,24 @@
{% block inside_header %} {% block inside_header %}
<div class="phase-banner-beta">
<strong class="phase-tag">BETA</strong>
</div>
{% if current_user.is_authenticated %}
<div class="">
<a class="" href="{{ url_for('main.sign_out') }}">Sign out</a>
</div>
{% endif %}
{% endblock %} {% endblock %}
{% block header_class %}with-proposition{% endblock %}
{% block proposition_header %}
<div class="header-proposition">
<div class='content'>
<div class="phase-banner-beta">
<strong class="phase-tag">BETA</strong>
</div>
{% if current_user.is_authenticated() %}
<nav id='proposition-menu'>
<p id='proposition-link'>
<a href="{{ url_for('main.sign_out')}}">Sign out</a>
</p>
</nav>
{% endif %}
</div>
</div>
{% endblock %}
{% set global_header_text = "GOV.UK Notify" %} {% set global_header_text = "GOV.UK Notify" %}

View File

@@ -9,8 +9,8 @@ def test_form_is_valid_returns_no_errors(notifications_admin, notifications_admi
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '12345'}) as req: data={'sms_code': '12345'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(req.request.form) form = TwoFactorForm(codes)
assert form.validate() is True assert form.validate() is True
assert len(form.errors) == 0 assert len(form.errors) == 0
@@ -19,8 +19,8 @@ def test_returns_errors_when_code_is_too_short(notifications_admin, notification
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '145'}) as req: data={'sms_code': '145'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(req.request.form) form = TwoFactorForm(codes)
assert form.validate() is False assert form.validate() is False
assert len(form.errors) == 1 assert len(form.errors) == 1
assert set(form.errors) == set({'sms_code': ['Code must be 5 digits', 'Code does not match']}) assert set(form.errors) == set({'sms_code': ['Code must be 5 digits', 'Code does not match']})
@@ -30,8 +30,8 @@ def test_returns_errors_when_code_is_missing(notifications_admin, notifications_
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={}) as req: data={}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(req.request.form) form = TwoFactorForm(codes)
assert form.validate() is False assert form.validate() is False
assert len(form.errors) == 1 assert len(form.errors) == 1
assert set(form.errors) == set({'sms_code': ['Code must not be empty']}) assert set(form.errors) == set({'sms_code': ['Code must not be empty']})
@@ -41,8 +41,8 @@ def test_returns_errors_when_code_contains_letters(notifications_admin, notifica
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': 'asdfg'}) as req: data={'sms_code': 'asdfg'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(req.request.form) form = TwoFactorForm(codes)
assert form.validate() is False assert form.validate() is False
assert len(form.errors) == 1 assert len(form.errors) == 1
assert set(form.errors) == set({'sms_code': ['Code must be 5 digits', 'Code does not match']}) assert set(form.errors) == set({'sms_code': ['Code must be 5 digits', 'Code does not match']})
@@ -52,12 +52,12 @@ def test_should_return_errors_when_code_is_expired(notifications_admin, notifica
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '23456'}) as req: data={'sms_code': '23456'}) as req:
user = create_test_user('active') user = create_test_user('active')
req.session['user_id'] = user.id
verify_codes_dao.add_code_with_expiry(user_id=user.id, verify_codes_dao.add_code_with_expiry(user_id=user.id,
code='23456', code='23456',
code_type='sms', code_type='sms',
expiry=datetime.now() + timedelta(hours=-2)) expiry=datetime.now() + timedelta(hours=-2))
form = TwoFactorForm(req.request.form) codes = verify_codes_dao.get_codes(user.id)
form = TwoFactorForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
assert len(errors) == 1 assert len(errors) == 1

View File

@@ -8,8 +8,8 @@ def test_form_should_have_error_when_code_is_not_valid(notifications_admin, noti
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '12345aa', 'email_code': 'abcde'}) as req: data={'sms_code': '12345aa', 'email_code': 'abcde'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
assert len(errors) == 2 assert len(errors) == 2
@@ -23,8 +23,8 @@ def test_should_return_errors_when_code_missing(notifications_admin, notificatio
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={}) as req: data={}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
expected = {'sms_code': ['SMS code can not be empty'], expected = {'sms_code': ['SMS code can not be empty'],
@@ -37,8 +37,8 @@ def test_should_return_errors_when_code_is_too_short(notifications_admin, notifi
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '123', 'email_code': '123'}) as req: data={'sms_code': '123', 'email_code': '123'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
expected = {'sms_code': ['Code must be 5 digits', 'Code does not match'], expected = {'sms_code': ['Code must be 5 digits', 'Code does not match'],
@@ -51,8 +51,8 @@ def test_should_return_errors_when_code_does_not_match(notifications_admin, noti
with notifications_admin.test_request_context(method='POST', with notifications_admin.test_request_context(method='POST',
data={'sms_code': '34567', 'email_code': '34567'}) as req: data={'sms_code': '34567', 'email_code': '34567'}) as req:
user = set_up_test_data() user = set_up_test_data()
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
expected = {'sms_code': ['Code does not match'], expected = {'sms_code': ['Code does not match'],
@@ -66,7 +66,6 @@ def test_should_return_errors_when_code_is_expired(notifications_admin, notifica
data={'sms_code': '23456', data={'sms_code': '23456',
'email_code': '23456'}) as req: 'email_code': '23456'}) as req:
user = create_test_user('pending') user = create_test_user('pending')
req.session['user_id'] = user.id
verify_codes_dao.add_code_with_expiry(user_id=user.id, verify_codes_dao.add_code_with_expiry(user_id=user.id,
code='23456', code='23456',
code_type='sms', code_type='sms',
@@ -76,8 +75,8 @@ def test_should_return_errors_when_code_is_expired(notifications_admin, notifica
code='23456', code='23456',
code_type='email', code_type='email',
expiry=datetime.now() + timedelta(hours=-2)) expiry=datetime.now() + timedelta(hours=-2))
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is False assert form.validate() is False
errors = form.errors errors = form.errors
expected = {'sms_code': ['Code has expired'], expected = {'sms_code': ['Code has expired'],
@@ -97,8 +96,8 @@ def test_should_return_valid_form_when_many_codes_exist(notifications_admin,
verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='60456', code_type='email') verify_codes_dao.add_code(user_id=user.id, code='60456', code_type='email')
verify_codes_dao.add_code(user_id=user.id, code='27856', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='27856', code_type='sms')
req.session['user_id'] = user.id codes = verify_codes_dao.get_codes(user.id)
form = VerifyForm(req.request.form) form = VerifyForm(codes)
assert form.validate() is True assert form.validate() is True

View File

@@ -3,40 +3,40 @@ from tests.app.main import create_test_user
def test_get_should_render_add_service_template(notifications_admin, notifications_admin_db, notify_db_session): def test_get_should_render_add_service_template(notifications_admin, notifications_admin_db, notify_db_session):
with notifications_admin.test_client() as client: with notifications_admin.test_request_context():
with client.session_transaction() as session: with notifications_admin.test_client() as client:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id client.login(user)
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
client.post('/two-factor', data={'sms_code': '12345'}) client.post('/two-factor', data={'sms_code': '12345'})
response = client.get('/add-service') response = client.get('/add-service')
assert response.status_code == 200 assert response.status_code == 200
assert 'Set up notifications for your service' in response.get_data(as_text=True) assert 'Set up notifications for your service' in response.get_data(as_text=True)
def test_should_add_service_and_redirect_to_next_page(notifications_admin, notifications_admin_db, notify_db_session): def test_should_add_service_and_redirect_to_next_page(notifications_admin, notifications_admin_db, notify_db_session):
with notifications_admin.test_client() as client: with notifications_admin.test_request_context():
with client.session_transaction() as session: with notifications_admin.test_client() as client:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id client.login(user)
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
client.post('/two-factor', data={'sms_code': '12345'}) client.post('/two-factor', data={'sms_code': '12345'})
response = client.post('/add-service', data={'service_name': 'testing the post'}) response = client.post('/add-service', data={'service_name': 'testing the post'})
assert response.status_code == 302 assert response.status_code == 302
assert response.location == 'http://localhost/dashboard' assert response.location == 'http://localhost/dashboard'
saved_service = services_dao.find_service_by_service_name('testing the post') saved_service = services_dao.find_service_by_service_name('testing the post')
assert saved_service is not None assert saved_service is not None
def test_should_return_form_errors_when_service_name_is_empty(notifications_admin, def test_should_return_form_errors_when_service_name_is_empty(notifications_admin,
notifications_admin_db, notifications_admin_db,
notify_db_session): notify_db_session):
with notifications_admin.test_client() as client: with notifications_admin.test_request_context():
with client.session_transaction() as session: with notifications_admin.test_client() as client:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id client.login(user)
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
client.post('/two-factor', data={'sms_code': '12345'}) client.post('/two-factor', data={'sms_code': '12345'})
response = client.post('/add-service', data={}) response = client.post('/add-service', data={})
assert response.status_code == 200 assert response.status_code == 200
assert 'Please enter your service name' in response.get_data(as_text=True) assert 'Please enter your service name' in response.get_data(as_text=True)

View File

@@ -12,7 +12,7 @@ def test_should_render_email_code_not_received_template_and_populate_email_addre
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
response = client.get(url_for('main.check_and_resend_email_code')) response = client.get(url_for('main.check_and_resend_email_code'))
assert response.status_code == 200 assert response.status_code == 200
assert 'Check your email address is correct and then resend the confirmation code' \ assert 'Check your email address is correct and then resend the confirmation code' \
@@ -29,7 +29,7 @@ def test_should_check_and_resend_email_code_redirect_to_verify(notifications_adm
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user.id, code='12345', code_type='email') verify_codes_dao.add_code(user.id, code='12345', code_type='email')
response = client.post(url_for('main.check_and_resend_email_code'), response = client.post(url_for('main.check_and_resend_email_code'),
data={'email_address': 'test@user.gov.uk'}) data={'email_address': 'test@user.gov.uk'})
@@ -46,7 +46,7 @@ def test_should_render_text_code_not_received_template(notifications_admin,
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user.id, code='12345', code_type='sms')
response = client.get(url_for('main.check_and_resend_text_code')) response = client.get(url_for('main.check_and_resend_text_code'))
assert response.status_code == 200 assert response.status_code == 200
@@ -64,7 +64,7 @@ def test_should_check_and_redirect_to_verify(notifications_admin,
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user.id, code='12345', code_type='sms')
response = client.post(url_for('main.check_and_resend_text_code'), response = client.post(url_for('main.check_and_resend_text_code'),
data={'mobile_number': '+441234123412'}) data={'mobile_number': '+441234123412'})
@@ -81,7 +81,7 @@ def test_should_update_email_address_resend_code(notifications_admin,
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='email') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='email')
response = client.post(url_for('main.check_and_resend_email_code'), response = client.post(url_for('main.check_and_resend_email_code'),
data={'email_address': 'new@address.gov.uk'}) data={'email_address': 'new@address.gov.uk'})
@@ -100,7 +100,7 @@ def test_should_update_mobile_number_resend_code(notifications_admin,
with client.session_transaction() as session: with client.session_transaction() as session:
_set_up_mocker(mocker) _set_up_mocker(mocker)
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
response = client.post(url_for('main.check_and_resend_text_code'), response = client.post(url_for('main.check_and_resend_text_code'),
data={'mobile_number': '+443456789012'}) data={'mobile_number': '+443456789012'})
@@ -117,7 +117,7 @@ def test_should_render_verification_code_not_received(notifications_admin,
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
response = client.get(url_for('main.verification_code_not_received')) response = client.get(url_for('main.verification_code_not_received'))
assert response.status_code == 200 assert response.status_code == 200
assert 'Resend verification code' in response.get_data(as_text=True) assert 'Resend verification code' in response.get_data(as_text=True)
@@ -133,7 +133,7 @@ def test_check_and_redirect_to_two_factor(notifications_admin,
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
_set_up_mocker(mocker) _set_up_mocker(mocker)
response = client.get(url_for('main.check_and_resend_verification_code')) response = client.get(url_for('main.check_and_resend_verification_code'))
assert response.status_code == 302 assert response.status_code == 302
@@ -148,7 +148,7 @@ def test_should_create_new_code_for_user(notifications_admin,
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
_set_up_mocker(mocker) _set_up_mocker(mocker)
response = client.get(url_for('main.check_and_resend_verification_code')) response = client.get(url_for('main.check_and_resend_verification_code'))

View File

@@ -7,8 +7,7 @@ def test_should_show_recent_jobs_on_dashboard(notifications_admin,
notify_db_session): notify_db_session):
with notifications_admin.test_request_context(): with notifications_admin.test_request_context():
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: user = create_test_user('active')
user = create_test_user('active')
client.login(user) client.login(user)
response = client.get(url_for('main.dashboard')) response = client.get(url_for('main.dashboard'))

View File

@@ -6,7 +6,13 @@ from tests.app.main import create_test_user
def test_should_render_two_factor_page(notifications_admin, notifications_admin_db, notify_db_session): def test_should_render_two_factor_page(notifications_admin, notifications_admin_db, notify_db_session):
with notifications_admin.test_request_context(): with notifications_admin.test_request_context():
response = notifications_admin.test_client().get(url_for('main.two_factor')) with notifications_admin.test_client() as client:
# TODO this lives here until we work out how to
# reassign the session after it is lost mid register process
with client.session_transaction() as session:
user = create_test_user('pending')
session['user_email'] = user.email_address
response = client.get(url_for('main.two_factor'))
assert response.status_code == 200 assert response.status_code == 200
assert '''We've sent you a text message with a verification code.''' in response.get_data(as_text=True) assert '''We've sent you a text message with a verification code.''' in response.get_data(as_text=True)
@@ -16,7 +22,7 @@ def test_should_login_user_and_redirect_to_dashboard(notifications_admin, notifi
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
response = client.post(url_for('main.two_factor'), response = client.post(url_for('main.two_factor'),
data={'sms_code': '12345'}) data={'sms_code': '12345'})
@@ -32,7 +38,7 @@ def test_should_return_200_with_sms_code_error_when_sms_code_is_wrong(notificati
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
response = client.post(url_for('main.two_factor'), response = client.post(url_for('main.two_factor'),
data={'sms_code': '23456'}) data={'sms_code': '23456'})
@@ -47,7 +53,7 @@ def test_should_login_user_when_multiple_valid_codes_exist(notifications_admin,
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('active') user = create_test_user('active')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='34567', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='34567', code_type='sms')

View File

@@ -10,7 +10,7 @@ def test_should_return_verify_template(notifications_admin, notifications_admin_
# reassign the session after it is lost mid register process # reassign the session after it is lost mid register process
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
response = client.get(url_for('main.verify')) response = client.get(url_for('main.verify'))
assert response.status_code == 200 assert response.status_code == 200
assert ( assert (
@@ -25,7 +25,7 @@ def test_should_redirect_to_add_service_when_code_are_correct(notifications_admi
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='email') verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='email')
response = client.post(url_for('main.verify'), response = client.post(url_for('main.verify'),
@@ -40,7 +40,7 @@ def test_should_activate_user_after_verify(notifications_admin, notifications_ad
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='email') verify_codes_dao.add_code(user_id=user.id, code='23456', code_type='email')
client.post(url_for('main.verify'), client.post(url_for('main.verify'),
@@ -56,12 +56,13 @@ def test_should_return_200_when_codes_are_wrong(notifications_admin, notificatio
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='23345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='23345', code_type='sms')
verify_codes_dao.add_code(user_id=user.id, code='98456', code_type='email') verify_codes_dao.add_code(user_id=user.id, code='98456', code_type='email')
response = client.post(url_for('main.verify'), response = client.post(url_for('main.verify'),
data={'sms_code': '12345', data={'sms_code': '12345',
'email_code': '23456'}) 'email_code': '23456'})
print(response.location)
assert response.status_code == 200 assert response.status_code == 200
resp_data = response.get_data(as_text=True) resp_data = response.get_data(as_text=True)
assert resp_data.count('Code does not match') == 2 assert resp_data.count('Code does not match') == 2
@@ -74,7 +75,7 @@ def test_should_mark_all_codes_as_used_when_many_codes_exist(notifications_admin
with notifications_admin.test_client() as client: with notifications_admin.test_client() as client:
with client.session_transaction() as session: with client.session_transaction() as session:
user = create_test_user('pending') user = create_test_user('pending')
session['user_id'] = user.id session['user_email'] = user.email_address
code1 = verify_codes_dao.add_code(user_id=user.id, code='23345', code_type='sms') code1 = verify_codes_dao.add_code(user_id=user.id, code='23345', code_type='sms')
code2 = verify_codes_dao.add_code(user_id=user.id, code='98456', code_type='email') code2 = verify_codes_dao.add_code(user_id=user.id, code='98456', code_type='email')
code3 = verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') code3 = verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')

View File

@@ -18,7 +18,7 @@ class TestClient(FlaskClient):
def login(self, user): def login(self, user):
# Skipping authentication here and just log them in # Skipping authentication here and just log them in
with self.session_transaction() as session: with self.session_transaction() as session:
session['user_id'] = user.id session['user_email'] = user.email_address
verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms') verify_codes_dao.add_code(user_id=user.id, code='12345', code_type='sms')
response = self.post( response = self.post(
url_for('main.two_factor'), data={'sms_code': '12345'}) url_for('main.two_factor'), data={'sms_code': '12345'})