New strategy for transaction management.

Introduce a contextmanger function to handle exceptions and nested
transactions. Using the nested_transaction will start a
nested transaction with `db.session.begin_nested`, once the nested
transaction is complete the commit will happen.
`@transactional` has been updated to commit unless in a nested
transaction.
This commit is contained in:
Rebecca Law
2021-04-13 15:02:46 +01:00
parent cf35135605
commit 93908bacda
9 changed files with 29 additions and 52 deletions

View File

@@ -1,6 +1,5 @@
from flask import Blueprint, jsonify, request from flask import Blueprint, jsonify, request
from app import db
from app.billing.billing_schemas import ( from app.billing.billing_schemas import (
create_or_update_free_sms_fragment_limit_schema, create_or_update_free_sms_fragment_limit_schema,
serialize_ft_billing_remove_emails, serialize_ft_billing_remove_emails,
@@ -84,7 +83,6 @@ def get_free_sms_fragment_limit(service_id):
annual_billing = dao_create_or_update_annual_billing_for_year(service_id, annual_billing = dao_create_or_update_annual_billing_for_year(service_id,
annual_billing.free_sms_fragment_limit, annual_billing.free_sms_fragment_limit,
financial_year_start) financial_year_start)
db.session.commit()
return jsonify(annual_billing.serialize_free_sms_items()), 200 return jsonify(annual_billing.serialize_free_sms_items()), 200
@@ -120,4 +118,3 @@ def update_free_sms_fragment_limit_data(service_id, free_sms_fragment_limit, fin
free_sms_fragment_limit, free_sms_fragment_limit,
financial_year_start financial_year_start
) )
db.session.commit()

View File

@@ -880,7 +880,6 @@ def populate_annual_billing_with_the_previous_years_allowance(year):
dao_create_or_update_annual_billing_for_year(service_id=row.id, dao_create_or_update_annual_billing_for_year(service_id=row.id,
free_sms_fragment_limit=free_allowance[0], free_sms_fragment_limit=free_allowance[0],
financial_year_start=int(year)) financial_year_start=int(year))
db.session.commit()
@notify_command(name='populate-annual-billing-with-defaults') @notify_command(name='populate-annual-billing-with-defaults')
@@ -915,4 +914,3 @@ def populate_annual_billing_with_defaults(year, missing_services_only):
for service in active_services: for service in active_services:
set_default_free_allowance_for_service(service, year) set_default_free_allowance_for_service(service, year)
db.session.commit()

View File

@@ -1,12 +1,12 @@
from flask import current_app from flask import current_app
from app import db from app import db
from app.dao.dao_utils import nested_transactional, transactional from app.dao.dao_utils import transactional
from app.dao.date_util import get_current_financial_year_start_year from app.dao.date_util import get_current_financial_year_start_year
from app.models import AnnualBilling from app.models import AnnualBilling
@nested_transactional @transactional
def dao_create_or_update_annual_billing_for_year(service_id, free_sms_fragment_limit, financial_year_start): 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) result = dao_get_free_sms_fragment_limit_for_year(service_id, financial_year_start)

View File

@@ -1,4 +1,5 @@
import itertools import itertools
from contextlib import contextmanager
from functools import wraps from functools import wraps
from app import db from app import db
@@ -10,7 +11,10 @@ def transactional(func):
def commit_or_rollback(*args, **kwargs): def commit_or_rollback(*args, **kwargs):
try: try:
res = func(*args, **kwargs) res = func(*args, **kwargs)
db.session.commit()
if not db.session.registry().transaction.nested:
db.session.commit()
return res return res
except Exception: except Exception:
db.session.rollback() db.session.rollback()
@@ -18,21 +22,18 @@ def transactional(func):
return commit_or_rollback return commit_or_rollback
def nested_transactional(func): @contextmanager
# This creates a save point for the nested transaction. def nested_transaction():
# You must manage the commit or rollback from outer most call of the nested of the transactions. try:
@wraps(func) db.session.begin_nested()
def commit_or_rollback(*args, **kwargs): yield
try: db.session.commit()
db.session.begin_nested()
res = func(*args, **kwargs)
db.session.commit()
return res
except Exception:
db.session.rollback()
raise
return commit_or_rollback if not db.session.registry().transaction.nested:
db.session.commit()
except Exception:
db.session.rollback()
raise
class VersionOptions(): class VersionOptions():

View File

@@ -1,12 +1,7 @@
from sqlalchemy.sql.expression import func from sqlalchemy.sql.expression import func
from app import db from app import db
from app.dao.dao_utils import ( from app.dao.dao_utils import VersionOptions, transactional, version_class
VersionOptions,
nested_transactional,
transactional,
version_class,
)
from app.models import Domain, Organisation, Service, User from app.models import Domain, Organisation, Service, User
@@ -110,7 +105,7 @@ def _update_organisation_services(organisation, attribute, only_where_none=True)
db.session.add(service) db.session.add(service)
@nested_transactional @transactional
@version_class(Service) @version_class(Service)
def dao_add_service_to_organisation(service, organisation_id): def dao_add_service_to_organisation(service, organisation_id):
organisation = Organisation.query.filter_by( organisation = Organisation.query.filter_by(

View File

@@ -7,12 +7,7 @@ from sqlalchemy.orm import joinedload
from sqlalchemy.sql.expression import and_, asc, case, func from sqlalchemy.sql.expression import and_, asc, case, func
from app import db from app import db
from app.dao.dao_utils import ( from app.dao.dao_utils import VersionOptions, transactional, version_class
VersionOptions,
nested_transactional,
transactional,
version_class,
)
from app.dao.date_util import get_current_financial_year from app.dao.date_util import get_current_financial_year
from app.dao.email_branding_dao import dao_get_email_branding_by_name from app.dao.email_branding_dao import dao_get_email_branding_by_name
from app.dao.letter_branding_dao import dao_get_letter_branding_by_name from app.dao.letter_branding_dao import dao_get_letter_branding_by_name
@@ -289,7 +284,7 @@ def dao_fetch_service_by_id_and_user(service_id, user_id):
).one() ).one()
@nested_transactional @transactional
@version_class(Service) @version_class(Service)
def dao_create_service( def dao_create_service(
service, service,

View File

@@ -1,10 +1,10 @@
from flask import Blueprint, abort, current_app, jsonify, request from flask import Blueprint, abort, current_app, jsonify, request
from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.exc import IntegrityError
from app import db
from app.config import QueueNames from app.config import QueueNames
from app.dao.annual_billing_dao import set_default_free_allowance_for_service from app.dao.annual_billing_dao import set_default_free_allowance_for_service
from app.dao.dao_utils import nested_transaction
from app.dao.fact_billing_dao import fetch_usage_year_for_organisation from app.dao.fact_billing_dao import fetch_usage_year_for_organisation
from app.dao.organisation_dao import ( from app.dao.organisation_dao import (
dao_add_service_to_organisation, dao_add_service_to_organisation,
@@ -120,13 +120,9 @@ def link_service_to_organisation(organisation_id):
service = dao_fetch_service_by_id(data['service_id']) service = dao_fetch_service_by_id(data['service_id'])
service.organisation = None service.organisation = None
try: with nested_transaction():
dao_add_service_to_organisation(service, organisation_id) dao_add_service_to_organisation(service, organisation_id)
set_default_free_allowance_for_service(service, year_start=None) set_default_free_allowance_for_service(service, year_start=None)
db.session.commit()
except SQLAlchemyError as e:
db.session.rollback()
raise e
return '', 204 return '', 204

View File

@@ -4,10 +4,9 @@ from datetime import datetime
from flask import Blueprint, current_app, jsonify, request from flask import Blueprint, current_app, jsonify, request
from notifications_utils.letter_timings import letter_can_be_cancelled from notifications_utils.letter_timings import letter_can_be_cancelled
from notifications_utils.timezones import convert_utc_to_bst from notifications_utils.timezones import convert_utc_to_bst
from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
from app import db
from app.aws import s3 from app.aws import s3
from app.config import QueueNames from app.config import QueueNames
from app.dao import fact_notification_status_dao, notifications_dao from app.dao import fact_notification_status_dao, notifications_dao
@@ -19,7 +18,7 @@ from app.dao.api_key_dao import (
save_model_api_key, save_model_api_key,
) )
from app.dao.broadcast_service_dao import set_broadcast_service_type 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.date_util import get_financial_year
from app.dao.fact_notification_status_dao import ( from app.dao.fact_notification_status_dao import (
fetch_monthly_template_usage_for_service, fetch_monthly_template_usage_for_service,
@@ -255,13 +254,9 @@ def create_service():
# unpack valid json into service object # unpack valid json into service object
valid_service = Service.from_json(data) valid_service = Service.from_json(data)
try: with nested_transaction():
dao_create_service(valid_service, user) dao_create_service(valid_service, user)
set_default_free_allowance_for_service(service=valid_service, year_start=None) set_default_free_allowance_for_service(service=valid_service, year_start=None)
db.session.commit()
except SQLAlchemyError as e:
db.session.rollback()
raise e
return jsonify(data=service_schema.dump(valid_service).data), 201 return jsonify(data=service_schema.dump(valid_service).data), 201

View File

@@ -233,7 +233,7 @@ def test_add_service_to_organisation(sample_service, sample_organisation):
sample_organisation.crown = False sample_organisation.crown = False
dao_add_service_to_organisation(sample_service, sample_organisation.id) dao_add_service_to_organisation(sample_service, sample_organisation.id)
db.session.commit()
assert len(sample_organisation.services) == 1 assert len(sample_organisation.services) == 1
assert sample_organisation.services[0].id == sample_service.id assert sample_organisation.services[0].id == sample_service.id