Compare commits

...

4 Commits

Author SHA1 Message Date
Ben Thorner
db81c9c355 Try a more magical way of doing nested transactions 2021-04-13 09:10:59 +01:00
Rebecca Law
6704be4021 Adding @nested_transactional for transactions that require more than one
db update/insert.

Using a savepoint for the multiple transactions allows us to rollback if
there is an error when executing the second db transaction.
However, this does add a bit of complexity. Developers need to manage
the db session when calling multiple nested tranactions.

Unit tests have been added to test this functionality and some end to
end tests have been done to make sure all transactions are rollback if
there is an exception while executing the transaction.
2021-04-12 13:59:57 +01:00
Rebecca Law
7806bb5246 When a service is created add the default annual billing for the service.
This will need to be merged before https://github.com/alphagov/notifications-admin/pull/3855, it will be that until the admin PR is merged the annual billing will be set twice, but that's not an issue.
2021-04-07 09:58:12 +01:00
Rebecca Law
964f7b4b52 When a service is associated with a organisation set the free allowance to
the default free allowance for the organisation type.

The update/insert for the default free allowance is done in a separate
transaction. Updates to services need to happen in a transaction to
trigger the insert into the ServicesHistory table. For that reason the
call to set_default_free_allowance_for_service is done after the service
is updated.
I've added a try/except around the set_default_free_allowance_for_service call to ensure we still get the update to the service but get an exception log if the update to annual_billing fails. I believe it's important to preserve the update to the service in the unlikely event that the annual_billing upsert fails.
2021-04-06 13:42:18 +01:00
7 changed files with 170 additions and 14 deletions

View File

@@ -6,7 +6,6 @@ from app.dao.date_util import get_current_financial_year_start_year
from app.models import AnnualBilling
@transactional
def dao_create_or_update_annual_billing_for_year(service_id, free_sms_fragment_limit, financial_year_start):
result = dao_get_free_sms_fragment_limit_for_year(service_id, financial_year_start)
@@ -90,6 +89,7 @@ def set_default_free_allowance_for_service(service, year_start=None):
}
if not year_start:
year_start = get_current_financial_year_start_year()
# handle cases where the year is less than 2020 or greater than 2021
if year_start < 2020:
year_start = 2020
if year_start > 2021:

View File

@@ -1,16 +1,34 @@
import itertools
from contextlib import contextmanager
from functools import wraps
from app import db
from app.history_meta import create_history
@contextmanager
def nested_transaction():
try:
db.session.begin_nested()
yield
db.session.commit()
if not db.session.registry().transaction.nested:
db.session.commit()
except Exception:
db.session.rollback()
raise
def transactional(func):
@wraps(func)
def commit_or_rollback(*args, **kwargs):
try:
res = func(*args, **kwargs)
db.session.commit()
if not db.session.registry().transaction.nested:
db.session.commit()
return res
except Exception:
db.session.rollback()

View File

@@ -1,8 +1,11 @@
from flask import Blueprint, abort, current_app, jsonify, request
from sqlalchemy.exc import IntegrityError
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from app.dao.dao_utils import nested_transaction
from app.config import QueueNames
from app.dao.annual_billing_dao import set_default_free_allowance_for_service
from app.dao.fact_billing_dao import fetch_usage_year_for_organisation
from app.dao.organisation_dao import (
dao_add_service_to_organisation,
@@ -118,7 +121,9 @@ def link_service_to_organisation(organisation_id):
service = dao_fetch_service_by_id(data['service_id'])
service.organisation = None
dao_add_service_to_organisation(service, organisation_id)
with nested_transaction():
dao_add_service_to_organisation(service, organisation_id)
set_default_free_allowance_for_service(service, year_start=None)
return '', 204

View File

@@ -4,12 +4,13 @@ from datetime import datetime
from flask import Blueprint, current_app, jsonify, request
from notifications_utils.letter_timings import letter_can_be_cancelled
from notifications_utils.timezones import convert_utc_to_bst
from sqlalchemy.exc import IntegrityError
from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from sqlalchemy.orm.exc import NoResultFound
from app.aws import s3
from app.config import QueueNames
from app.dao import fact_notification_status_dao, notifications_dao
from app.dao.annual_billing_dao import set_default_free_allowance_for_service
from app.dao.api_key_dao import (
expire_api_key,
get_model_api_keys,
@@ -17,7 +18,7 @@ from app.dao.api_key_dao import (
save_model_api_key,
)
from app.dao.broadcast_service_dao import set_broadcast_service_type
from app.dao.dao_utils import dao_rollback
from app.dao.dao_utils import dao_rollback, nested_transaction
from app.dao.date_util import get_financial_year
from app.dao.fact_notification_status_dao import (
fetch_monthly_template_usage_for_service,
@@ -253,7 +254,9 @@ def create_service():
# unpack valid json into service object
valid_service = Service.from_json(data)
dao_create_service(valid_service, user)
with nested_transaction():
dao_create_service(valid_service, user)
set_default_free_allowance_for_service(valid_service, year_start=None)
return jsonify(data=service_schema.dump(valid_service).data), 201

View File

@@ -91,3 +91,21 @@ def test_set_default_free_allowance_for_service_using_correct_year(sample_servic
25000,
2020
)
@freeze_time('2021-04-01 14:02:00')
def test_set_default_free_allowance_for_service_updates_existing_year(sample_service):
set_default_free_allowance_for_service(service=sample_service, year_start=None)
annual_billing = AnnualBilling.query.all()
assert not sample_service.organisation_type
assert len(annual_billing) == 1
assert annual_billing[0].service_id == sample_service.id
assert annual_billing[0].free_sms_fragment_limit == 10000
sample_service.organisation_type = 'central'
set_default_free_allowance_for_service(service=sample_service, year_start=None)
annual_billing = AnnualBilling.query.all()
assert len(annual_billing) == 1
assert annual_billing[0].service_id == sample_service.id
assert annual_billing[0].free_sms_fragment_limit == 150000

View File

@@ -3,13 +3,14 @@ from datetime import datetime
import pytest
from freezegun import freeze_time
from sqlalchemy.exc import SQLAlchemyError
from app.dao.organisation_dao import (
dao_add_service_to_organisation,
dao_add_user_to_organisation,
)
from app.dao.services_dao import dao_archive_service
from app.models import Organisation
from app.models import AnnualBilling, Organisation
from tests.app.db import (
create_annual_billing,
create_domain,
@@ -491,19 +492,65 @@ def test_post_update_organisation_set_mou_emails_signed_by(
}
def test_post_link_service_to_organisation(admin_request, sample_service, sample_organisation):
def test_post_link_service_to_organisation(admin_request, sample_service):
data = {
'service_id': str(sample_service.id)
}
organisation = create_organisation(organisation_type='central')
assert len(organisation.services) == 0
assert len(AnnualBilling.query.all()) == 0
admin_request.post(
'organisation.link_service_to_organisation',
_data=data,
organisation_id=sample_organisation.id,
organisation_id=organisation.id,
_expected_status=204
)
assert len(organisation.services) == 1
assert sample_service.organisation_type == 'central'
def test_post_link_service_to_organisation_inserts_annual_billing(admin_request, sample_service):
data = {
'service_id': str(sample_service.id)
}
organisation = create_organisation(organisation_type='central')
assert len(organisation.services) == 0
assert len(AnnualBilling.query.all()) == 0
admin_request.post(
'organisation.link_service_to_organisation',
_data=data,
organisation_id=organisation.id,
_expected_status=204
)
assert len(sample_organisation.services) == 1
annual_billing = AnnualBilling.query.all()
assert len(annual_billing) == 1
assert annual_billing[0].free_sms_fragment_limit == 150000
def test_post_link_service_to_organisation_rollback_service_if_annual_billing_update_fails(
admin_request, sample_service, mocker
):
mocker.patch('app.dao.annual_billing_dao.dao_create_or_update_annual_billing_for_year',
side_effect=SQLAlchemyError)
data = {
'service_id': str(sample_service.id)
}
assert not sample_service.organisation_type
organisation = create_organisation(organisation_type='central')
assert len(organisation.services) == 0
assert len(AnnualBilling.query.all()) == 0
with pytest.raises(expected_exception=SQLAlchemyError):
admin_request.post(
'organisation.link_service_to_organisation',
_data=data,
organisation_id=organisation.id,
_expected_status=404
)
assert not sample_service.organisation_type
assert len(organisation.services) == 0
assert len(AnnualBilling.query.all()) == 0
def test_post_link_service_to_another_org(
@@ -511,7 +558,8 @@ def test_post_link_service_to_another_org(
data = {
'service_id': str(sample_service.id)
}
assert len(sample_organisation.services) == 0
assert not sample_service.organisation_type
admin_request.post(
'organisation.link_service_to_organisation',
_data=data,
@@ -520,8 +568,9 @@ def test_post_link_service_to_another_org(
)
assert len(sample_organisation.services) == 1
assert not sample_service.organisation_type
new_org = create_organisation()
new_org = create_organisation(organisation_type='central')
admin_request.post(
'organisation.link_service_to_organisation',
_data=data,
@@ -530,6 +579,10 @@ def test_post_link_service_to_another_org(
)
assert not sample_organisation.services
assert len(new_org.services) == 1
assert sample_service.organisation_type == 'central'
annual_billing = AnnualBilling.query.all()
assert len(annual_billing) == 1
assert annual_billing[0].free_sms_fragment_limit == 150000
def test_post_link_service_to_organisation_nonexistent_organisation(
@@ -569,6 +622,23 @@ def test_post_link_service_to_organisation_missing_payload(
)
def test_link_service_to_organisation_does_not_updates_service_if_annual_billing_update_fails(
mocker, admin_request, sample_service, sample_organisation
):
mocker.patch('app.organisation.rest.set_default_free_allowance_for_service', side_effect=SQLAlchemyError)
data = {
'service_id': str(sample_service.id)
}
with pytest.raises(expected_exception=SQLAlchemyError):
admin_request.post(
'organisation.link_service_to_organisation',
organisation_id=str(sample_organisation.id),
_data=data,
)
assert not sample_service.organisation_id
assert len(AnnualBilling.query.all()) == 0
def test_rest_get_organisation_services(
admin_request, sample_organisation, sample_service):
dao_add_service_to_organisation(sample_service, sample_organisation.id)

View File

@@ -6,6 +6,7 @@ from unittest.mock import ANY
import pytest
from flask import current_app, url_for
from freezegun import freeze_time
from sqlalchemy.exc import SQLAlchemyError
from app.dao.organisation_dao import dao_add_service_to_organisation
from app.dao.service_sms_sender_dao import dao_get_sms_senders_by_service_id
@@ -31,6 +32,7 @@ from app.models import (
SERVICE_PERMISSION_TYPES,
SMS_TYPE,
UPLOAD_LETTERS,
AnnualBilling,
EmailBranding,
InboundNumber,
Notification,
@@ -482,6 +484,46 @@ def test_create_service_with_domain_sets_organisation(
assert json_resp['data']['organisation'] is None
def test_create_service_should_create_annual_billing_for_service(
admin_request, sample_user
):
data = {
'name': 'created service',
'user_id': str(sample_user.id),
'message_limit': 1000,
'restricted': False,
'active': False,
'email_from': 'created.service',
'created_by': str(sample_user.id)
}
assert len(AnnualBilling.query.all()) == 0
admin_request.post('service.create_service', _data=data, _expected_status=201)
annual_billing = AnnualBilling.query.all()
assert len(annual_billing) == 1
def test_create_service_should_create_service_if_annual_billing_query_fails(
admin_request, sample_user, mocker
):
mocker.patch('app.service.rest.set_default_free_allowance_for_service', side_effect=SQLAlchemyError)
data = {
'name': 'created service',
'user_id': str(sample_user.id),
'message_limit': 1000,
'restricted': False,
'active': False,
'email_from': 'created.service',
'created_by': str(sample_user.id)
}
assert len(AnnualBilling.query.all()) == 0
admin_request.post('service.create_service', _data=data, _expected_status=201)
annual_billing = AnnualBilling.query.all()
assert len(annual_billing) == 0
assert len(Service.query.filter(Service.name == 'created service').all()) == 1
def test_create_service_inherits_branding_from_organisation(
admin_request,
sample_user,