merge from main

This commit is contained in:
Kenneth Kehl
2024-10-28 13:42:43 -07:00
20 changed files with 727 additions and 502 deletions
+4 -4
View File
@@ -239,7 +239,7 @@
"filename": "tests/app/dao/test_services_dao.py", "filename": "tests/app/dao/test_services_dao.py",
"hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8", "hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8",
"is_verified": false, "is_verified": false,
"line_number": 265, "line_number": 289,
"is_secret": false "is_secret": false
} }
], ],
@@ -249,7 +249,7 @@
"filename": "tests/app/dao/test_users_dao.py", "filename": "tests/app/dao/test_users_dao.py",
"hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8", "hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8",
"is_verified": false, "is_verified": false,
"line_number": 52, "line_number": 69,
"is_secret": false "is_secret": false
}, },
{ {
@@ -257,7 +257,7 @@
"filename": "tests/app/dao/test_users_dao.py", "filename": "tests/app/dao/test_users_dao.py",
"hashed_secret": "f2c57870308dc87f432e5912d4de6f8e322721ba", "hashed_secret": "f2c57870308dc87f432e5912d4de6f8e322721ba",
"is_verified": false, "is_verified": false,
"line_number": 176, "line_number": 199,
"is_secret": false "is_secret": false
} }
], ],
@@ -384,5 +384,5 @@
} }
] ]
}, },
"generated_at": "2024-09-27T16:42:53Z" "generated_at": "2024-10-28T20:26:27Z"
} }
+64 -42
View File
@@ -1,7 +1,7 @@
from datetime import timedelta from datetime import timedelta
from flask import current_app from flask import current_app
from sqlalchemy import asc, desc, or_, select, text, union from sqlalchemy import asc, delete, desc, func, or_, select, text, union, update
from sqlalchemy.orm import joinedload from sqlalchemy.orm import joinedload
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
from sqlalchemy.sql import functions from sqlalchemy.sql import functions
@@ -109,11 +109,12 @@ def _update_notification_status(
def update_notification_status_by_id( def update_notification_status_by_id(
notification_id, status, sent_by=None, provider_response=None, carrier=None notification_id, status, sent_by=None, provider_response=None, carrier=None
): ):
notification = ( stmt = (
Notification.query.with_for_update() select(Notification)
.with_for_update()
.filter(Notification.id == notification_id) .filter(Notification.id == notification_id)
.first()
) )
notification = db.session.execute(stmt).scalars().first()
if not notification: if not notification:
current_app.logger.info( current_app.logger.info(
@@ -156,9 +157,8 @@ def update_notification_status_by_id(
@autocommit @autocommit
def update_notification_status_by_reference(reference, status): def update_notification_status_by_reference(reference, status):
# this is used to update emails # this is used to update emails
notification = Notification.query.filter( stmt = select(Notification).filter(Notification.reference == reference)
Notification.reference == reference notification = db.session.execute(stmt).scalars().first()
).first()
if not notification: if not notification:
current_app.logger.error( current_app.logger.error(
@@ -200,19 +200,20 @@ def get_notifications_for_job(
def dao_get_notification_count_for_job_id(*, job_id): def dao_get_notification_count_for_job_id(*, job_id):
return Notification.query.filter_by(job_id=job_id).count() stmt = select(func.count(Notification.id)).filter_by(job_id=job_id)
return db.session.execute(stmt).scalar()
def dao_get_notification_count_for_service(*, service_id): def dao_get_notification_count_for_service(*, service_id):
notification_count = Notification.query.filter_by(service_id=service_id).count() stmt = select(func.count(Notification.id)).filter_by(service_id=service_id)
return notification_count return db.session.execute(stmt).scalar()
def dao_get_failed_notification_count(): def dao_get_failed_notification_count():
failed_count = Notification.query.filter_by( stmt = select(func.count(Notification.id)).filter_by(
status=NotificationStatus.FAILED status=NotificationStatus.FAILED
).count() )
return failed_count return db.session.execute(stmt).scalar()
def get_notification_with_personalisation(service_id, notification_id, key_type): def get_notification_with_personalisation(service_id, notification_id, key_type):
@@ -220,11 +221,12 @@ def get_notification_with_personalisation(service_id, notification_id, key_type)
if key_type: if key_type:
filter_dict["key_type"] = key_type filter_dict["key_type"] = key_type
return ( stmt = (
Notification.query.filter_by(**filter_dict) select(Notification)
.filter_by(**filter_dict)
.options(joinedload(Notification.template)) .options(joinedload(Notification.template))
.one()
) )
return db.session.execute(stmt).scalars().one()
def get_notification_by_id(notification_id, service_id=None, _raise=False): def get_notification_by_id(notification_id, service_id=None, _raise=False):
@@ -233,9 +235,13 @@ def get_notification_by_id(notification_id, service_id=None, _raise=False):
if service_id: if service_id:
filters.append(Notification.service_id == service_id) filters.append(Notification.service_id == service_id)
query = Notification.query.filter(*filters) stmt = select(Notification).filter(*filters)
return query.one() if _raise else query.first() return (
db.session.execute(stmt).scalars().one()
if _raise
else db.session.execute(stmt).scalars().first()
)
def get_notifications_for_service( def get_notifications_for_service(
@@ -415,12 +421,13 @@ def move_notifications_to_notification_history(
deleted += delete_count_per_call deleted += delete_count_per_call
# Deleting test Notifications, test notifications are not persisted to NotificationHistory # Deleting test Notifications, test notifications are not persisted to NotificationHistory
Notification.query.filter( stmt = delete(Notification).filter(
Notification.notification_type == notification_type, Notification.notification_type == notification_type,
Notification.service_id == service_id, Notification.service_id == service_id,
Notification.created_at < timestamp_to_delete_backwards_from, Notification.created_at < timestamp_to_delete_backwards_from,
Notification.key_type == KeyType.TEST, Notification.key_type == KeyType.TEST,
).delete(synchronize_session=False) )
db.session.execute(stmt)
db.session.commit() db.session.commit()
return deleted return deleted
@@ -442,8 +449,9 @@ def dao_timeout_notifications(cutoff_time, limit=100000):
current_statuses = [NotificationStatus.SENDING, NotificationStatus.PENDING] current_statuses = [NotificationStatus.SENDING, NotificationStatus.PENDING]
new_status = NotificationStatus.TEMPORARY_FAILURE new_status = NotificationStatus.TEMPORARY_FAILURE
notifications = ( stmt = (
Notification.query.filter( select(Notification)
.filter(
Notification.created_at < cutoff_time, Notification.created_at < cutoff_time,
Notification.status.in_(current_statuses), Notification.status.in_(current_statuses),
Notification.notification_type.in_( Notification.notification_type.in_(
@@ -451,14 +459,15 @@ def dao_timeout_notifications(cutoff_time, limit=100000):
), ),
) )
.limit(limit) .limit(limit)
.all()
) )
notifications = db.session.execute(stmt).scalars().all()
Notification.query.filter( stmt = (
Notification.id.in_([n.id for n in notifications]), update(Notification)
).update( .filter(Notification.id.in_([n.id for n in notifications]))
{"status": new_status, "updated_at": updated_at}, synchronize_session=False .values({"status": new_status, "updated_at": updated_at})
) )
db.session.execute(stmt)
db.session.commit() db.session.commit()
return notifications return notifications
@@ -466,15 +475,23 @@ def dao_timeout_notifications(cutoff_time, limit=100000):
@autocommit @autocommit
def dao_update_notifications_by_reference(references, update_dict): def dao_update_notifications_by_reference(references, update_dict):
updated_count = Notification.query.filter( stmt = (
Notification.reference.in_(references) update(Notification)
).update(update_dict, synchronize_session=False) .filter(Notification.reference.in_(references))
.values(update_dict)
)
result = db.session.execute(stmt)
updated_count = result.rowcount
updated_history_count = 0 updated_history_count = 0
if updated_count != len(references): if updated_count != len(references):
updated_history_count = NotificationHistory.query.filter( stmt = (
NotificationHistory.reference.in_(references) update(NotificationHistory)
).update(update_dict, synchronize_session=False) .filter(NotificationHistory.reference.in_(references))
.values(update_dict)
)
result = db.session.execute(stmt)
updated_history_count = result.rowcount
return updated_count, updated_history_count return updated_count, updated_history_count
@@ -541,18 +558,21 @@ def dao_get_notifications_by_recipient_or_reference(
def dao_get_notification_by_reference(reference): def dao_get_notification_by_reference(reference):
return Notification.query.filter(Notification.reference == reference).one() stmt = select(Notification).filter(Notification.reference == reference)
return db.session.execute(stmt).scalars().one()
def dao_get_notification_history_by_reference(reference): def dao_get_notification_history_by_reference(reference):
try: try:
# This try except is necessary because in test keys and research mode does not create notification history. # This try except is necessary because in test keys and research mode does not create notification history.
# Otherwise we could just search for the NotificationHistory object # Otherwise we could just search for the NotificationHistory object
return Notification.query.filter(Notification.reference == reference).one() stmt = select(Notification).filter(Notification.reference == reference)
return db.session.execute(stmt).scalars().one()
except NoResultFound: except NoResultFound:
return NotificationHistory.query.filter( stmt = select(NotificationHistory).filter(
NotificationHistory.reference == reference NotificationHistory.reference == reference
).one() )
return db.session.execute(stmt).scalars().one()
def dao_get_notifications_processing_time_stats(start_date, end_date): def dao_get_notifications_processing_time_stats(start_date, end_date):
@@ -590,11 +610,12 @@ def dao_get_notifications_processing_time_stats(start_date, end_date):
def dao_get_last_notification_added_for_job_id(job_id): def dao_get_last_notification_added_for_job_id(job_id):
last_notification_added = ( stmt = (
Notification.query.filter(Notification.job_id == job_id) select(Notification)
.filter(Notification.job_id == job_id)
.order_by(Notification.job_row_number.desc()) .order_by(Notification.job_row_number.desc())
.first()
) )
last_notification_added = db.session.execute(stmt).scalars().first()
return last_notification_added return last_notification_added
@@ -602,11 +623,12 @@ def dao_get_last_notification_added_for_job_id(job_id):
def notifications_not_yet_sent(should_be_sending_after_seconds, notification_type): def notifications_not_yet_sent(should_be_sending_after_seconds, notification_type):
older_than_date = utc_now() - timedelta(seconds=should_be_sending_after_seconds) older_than_date = utc_now() - timedelta(seconds=should_be_sending_after_seconds)
notifications = Notification.query.filter( stmt = select(Notification).filter(
Notification.created_at <= older_than_date, Notification.created_at <= older_than_date,
Notification.notification_type == notification_type, Notification.notification_type == notification_type,
Notification.status == NotificationStatus.CREATED, Notification.status == NotificationStatus.CREATED,
).all() )
notifications = db.session.execute(stmt).scalars().all()
return notifications return notifications
+31 -22
View File
@@ -1,3 +1,4 @@
from sqlalchemy import delete, select, update
from sqlalchemy.sql.expression import func from sqlalchemy.sql.expression import func
from app import db from app import db
@@ -6,55 +7,57 @@ from app.models import Domain, Organization, Service, User
def dao_get_organizations(): def dao_get_organizations():
return Organization.query.order_by( stmt = select(Organization).order_by(
Organization.active.desc(), Organization.name.asc() Organization.active.desc(), Organization.name.asc()
).all() )
return db.session.execute(stmt).scalars().all()
def dao_count_organizations_with_live_services(): def dao_count_organizations_with_live_services():
return ( stmt = (
db.session.query(Organization.id) select(func.count(func.distinct(Organization.id)))
.join(Organization.services) .join(Organization.services)
.filter( .filter(
Service.active.is_(True), Service.active.is_(True),
Service.restricted.is_(False), Service.restricted.is_(False),
Service.count_as_live.is_(True), Service.count_as_live.is_(True),
) )
.distinct()
.count()
) )
return db.session.execute(stmt).scalar() or 0
def dao_get_organization_services(organization_id): def dao_get_organization_services(organization_id):
return Organization.query.filter_by(id=organization_id).one().services stmt = select(Organization).filter_by(id=organization_id)
return db.session.execute(stmt).scalars().one().services
def dao_get_organization_live_services(organization_id): def dao_get_organization_live_services(organization_id):
return Service.query.filter_by( stmt = select(Service).filter_by(organization_id=organization_id, restricted=False)
organization_id=organization_id, restricted=False return db.session.execute(stmt).scalars().all()
).all()
def dao_get_organization_by_id(organization_id): def dao_get_organization_by_id(organization_id):
return Organization.query.filter_by(id=organization_id).one() stmt = select(Organization).filter_by(id=organization_id)
return db.session.execute(stmt).scalars().one()
def dao_get_organization_by_email_address(email_address): def dao_get_organization_by_email_address(email_address):
email_address = email_address.lower().replace(".gsi.gov.uk", ".gov.uk") email_address = email_address.lower().replace(".gsi.gov.uk", ".gov.uk")
stmt = select(Domain).order_by(func.char_length(Domain.domain).desc())
for domain in Domain.query.order_by(func.char_length(Domain.domain).desc()).all(): domains = db.session.execute(stmt).scalars().all()
for domain in domains:
if email_address.endswith( if email_address.endswith(
"@{}".format(domain.domain) "@{}".format(domain.domain)
) or email_address.endswith(".{}".format(domain.domain)): ) or email_address.endswith(".{}".format(domain.domain)):
return Organization.query.filter_by(id=domain.organization_id).one() stmt = select(Organization).filter_by(id=domain.organization_id)
return db.session.execute(stmt).scalars().one()
return None return None
def dao_get_organization_by_service_id(service_id): def dao_get_organization_by_service_id(service_id):
return ( stmt = select(Organization).join(Organization.services).filter_by(id=service_id)
Organization.query.join(Organization.services).filter_by(id=service_id).first() return db.session.execute(stmt).scalars().first()
)
@autocommit @autocommit
@@ -65,10 +68,14 @@ def dao_create_organization(organization):
@autocommit @autocommit
def dao_update_organization(organization_id, **kwargs): def dao_update_organization(organization_id, **kwargs):
domains = kwargs.pop("domains", None) domains = kwargs.pop("domains", None)
num_updated = Organization.query.filter_by(id=organization_id).update(kwargs) stmt = (
update(Organization).where(Organization.id == organization_id).values(**kwargs)
)
num_updated = db.session.execute(stmt).rowcount
if isinstance(domains, list): if isinstance(domains, list):
Domain.query.filter_by(organization_id=organization_id).delete() stmt = delete(Domain).filter_by(organization_id=organization_id)
db.session.execute(stmt)
db.session.bulk_save_objects( db.session.bulk_save_objects(
[ [
Domain(domain=domain.lower(), organization_id=organization_id) Domain(domain=domain.lower(), organization_id=organization_id)
@@ -76,7 +83,7 @@ def dao_update_organization(organization_id, **kwargs):
] ]
) )
organization = Organization.query.get(organization_id) organization = db.session.get(Organization, organization_id)
if "organization_type" in kwargs: if "organization_type" in kwargs:
_update_organization_services( _update_organization_services(
organization, "organization_type", only_where_none=False organization, "organization_type", only_where_none=False
@@ -101,7 +108,8 @@ def _update_organization_services(organization, attribute, only_where_none=True)
@autocommit @autocommit
@version_class(Service) @version_class(Service)
def dao_add_service_to_organization(service, organization_id): def dao_add_service_to_organization(service, organization_id):
organization = Organization.query.filter_by(id=organization_id).one() stmt = select(Organization).filter_by(id=organization_id)
organization = db.session.execute(stmt).scalars().one()
service.organization_id = organization_id service.organization_id = organization_id
service.organization_type = organization.organization_type service.organization_type = organization.organization_type
@@ -122,7 +130,8 @@ def dao_get_users_for_organization(organization_id):
@autocommit @autocommit
def dao_add_user_to_organization(organization_id, user_id): def dao_add_user_to_organization(organization_id, user_id):
organization = dao_get_organization_by_id(organization_id) organization = dao_get_organization_by_id(organization_id)
user = User.query.filter_by(id=user_id).one() stmt = select(User).filter_by(id=user_id)
user = db.session.execute(stmt).scalars().one()
user.organizations.append(organization) user.organizations.append(organization)
db.session.add(organization) db.session.add(organization)
return user return user
+10 -6
View File
@@ -1,12 +1,14 @@
from sqlalchemy import delete, select
from app import db from app import db
from app.dao.dao_utils import autocommit from app.dao.dao_utils import autocommit
from app.models import ServicePermission from app.models import ServicePermission
def dao_fetch_service_permissions(service_id): def dao_fetch_service_permissions(service_id):
return ServicePermission.query.filter(
ServicePermission.service_id == service_id stmt = select(ServicePermission).filter(ServicePermission.service_id == service_id)
).all() return db.session.execute(stmt).scalars().all()
@autocommit @autocommit
@@ -16,9 +18,11 @@ def dao_add_service_permission(service_id, permission):
def dao_remove_service_permission(service_id, permission): def dao_remove_service_permission(service_id, permission):
deleted = ServicePermission.query.filter(
stmt = delete(ServicePermission).where(
ServicePermission.service_id == service_id, ServicePermission.service_id == service_id,
ServicePermission.permission == permission, ServicePermission.permission == permission,
).delete() )
result = db.session.execute(stmt)
db.session.commit() db.session.commit()
return deleted return result.rowcount
+9 -6
View File
@@ -1,4 +1,4 @@
from sqlalchemy import desc from sqlalchemy import desc, select
from app import db from app import db
from app.dao.dao_utils import autocommit from app.dao.dao_utils import autocommit
@@ -17,17 +17,20 @@ def insert_service_sms_sender(service, sms_sender):
def dao_get_service_sms_senders_by_id(service_id, service_sms_sender_id): def dao_get_service_sms_senders_by_id(service_id, service_sms_sender_id):
return ServiceSmsSender.query.filter_by( stmt = select(ServiceSmsSender).filter_by(
id=service_sms_sender_id, service_id=service_id, archived=False id=service_sms_sender_id, service_id=service_id, archived=False
).one() )
return db.session.execute(stmt).scalars().one()
def dao_get_sms_senders_by_service_id(service_id): def dao_get_sms_senders_by_service_id(service_id):
return (
ServiceSmsSender.query.filter_by(service_id=service_id, archived=False) stmt = (
select(ServiceSmsSender)
.filter_by(service_id=service_id, archived=False)
.order_by(desc(ServiceSmsSender.is_default)) .order_by(desc(ServiceSmsSender.is_default))
.all()
) )
return db.session.execute(stmt).scalars().all()
@autocommit @autocommit
+8 -10
View File
@@ -1,25 +1,23 @@
from sqlalchemy import select
from app import db from app import db
from app.dao.dao_utils import autocommit from app.dao.dao_utils import autocommit
from app.models import ServiceUser, User from app.models import ServiceUser, User
def dao_get_service_user(user_id, service_id): def dao_get_service_user(user_id, service_id):
# TODO: This has been changed to account for the test case failure stmt = select(ServiceUser).filter_by(user_id=user_id, service_id=service_id)
# that used this method but have any service user to return. Somehow, this return db.session.execute(stmt).scalars().one_or_none()
# started to throw an error with one() method in sqlalchemy 2.0 unlike 1.4
return ServiceUser.query.filter_by(
user_id=user_id, service_id=service_id
).one_or_none()
def dao_get_active_service_users(service_id): def dao_get_active_service_users(service_id):
query = (
db.session.query(ServiceUser) stmt = (
select(ServiceUser)
.join(User, User.id == ServiceUser.user_id) .join(User, User.id == ServiceUser.user_id)
.filter(User.state == "active", ServiceUser.service_id == service_id) .filter(User.state == "active", ServiceUser.service_id == service_id)
) )
return db.session.execute(stmt).scalars().all()
return query.all()
def dao_get_service_users_by_user_id(user_id): def dao_get_service_users_by_user_id(user_id):
+133 -105
View File
@@ -2,7 +2,7 @@ import uuid
from datetime import timedelta from datetime import timedelta
from flask import current_app from flask import current_app
from sqlalchemy import Float, cast, select from sqlalchemy import Float, cast, delete, select
from sqlalchemy.orm import joinedload from sqlalchemy.orm import joinedload
from sqlalchemy.sql.expression import and_, asc, case, func from sqlalchemy.sql.expression import and_, asc, case, func
@@ -51,34 +51,42 @@ from app.utils import (
def dao_fetch_all_services(only_active=False): def dao_fetch_all_services(only_active=False):
query = Service.query.order_by(asc(Service.created_at)).options(
joinedload(Service.users) stmt = select(Service)
)
if only_active: if only_active:
query = query.filter(Service.active) stmt = stmt.where(Service.active)
return query.all() stmt = stmt.order_by(asc(Service.created_at)).options(joinedload(Service.users))
result = db.session.execute(stmt)
return result.unique().scalars().all()
def get_services_by_partial_name(service_name): def get_services_by_partial_name(service_name):
service_name = escape_special_characters(service_name) service_name = escape_special_characters(service_name)
return Service.query.filter(Service.name.ilike("%{}%".format(service_name))).all() stmt = select(Service).where(Service.name.ilike("%{}%".format(service_name)))
result = db.session.execute(stmt)
return result.scalars().all()
def dao_count_live_services(): def dao_count_live_services():
return Service.query.filter_by( stmt = (
active=True, select(func.count())
restricted=False, .select_from(Service)
count_as_live=True, .where(
).count() Service.active, Service.count_as_live, Service.restricted == False # noqa
)
)
result = db.session.execute(stmt)
return result.scalar() # Retrieves the count
def dao_fetch_live_services_data(): def dao_fetch_live_services_data():
year_start_date, year_end_date = get_current_calendar_year() year_start_date, year_end_date = get_current_calendar_year()
most_recent_annual_billing = ( most_recent_annual_billing = (
db.session.query( select(
AnnualBilling.service_id, AnnualBilling.service_id,
func.max(AnnualBilling.financial_year_start).label("year"), func.max(AnnualBilling.financial_year_start).label("year"),
) )
@@ -86,13 +94,17 @@ def dao_fetch_live_services_data():
.subquery() .subquery()
) )
this_year_ft_billing = FactBilling.query.filter( this_year_ft_billing = (
FactBilling.local_date >= year_start_date, select(FactBilling)
FactBilling.local_date <= year_end_date, .filter(
).subquery() FactBilling.local_date >= year_start_date,
FactBilling.local_date <= year_end_date,
)
.subquery()
)
data = ( stmt = (
db.session.query( select(
Service.id.label("service_id"), Service.id.label("service_id"),
Service.name.label("service_name"), Service.name.label("service_name"),
Organization.name.label("organization_name"), Organization.name.label("organization_name"),
@@ -156,8 +168,9 @@ def dao_fetch_live_services_data():
AnnualBilling.free_sms_fragment_limit, AnnualBilling.free_sms_fragment_limit,
) )
.order_by(asc(Service.go_live_at)) .order_by(asc(Service.go_live_at))
.all()
) )
data = db.session.execute(stmt).all()
results = [] results = []
for row in data: for row in data:
existing_service = next( existing_service = next(
@@ -183,48 +196,55 @@ def dao_fetch_service_by_id(service_id, only_active=False):
stmt = stmt.where(Service.active) stmt = stmt.where(Service.active)
result = db.session.execute(stmt) result = db.session.execute(stmt)
return result.unique().scalars().one() return result.unique().scalars().unique().one()
def dao_fetch_service_by_inbound_number(number): def dao_fetch_service_by_inbound_number(number):
inbound_number = InboundNumber.query.filter( stmt = select(InboundNumber).where(
InboundNumber.number == number, InboundNumber.active InboundNumber.number == number, InboundNumber.active
).first() )
result = db.session.execute(stmt)
inbound_number = result.scalars().first()
if not inbound_number: if not inbound_number:
return None return None
return Service.query.filter(Service.id == inbound_number.service_id).first() stmt = select(Service).where(Service.id == inbound_number.service_id)
result = db.session.execute(stmt)
return result.scalars().first()
def dao_fetch_service_by_id_with_api_keys(service_id, only_active=False): def dao_fetch_service_by_id_with_api_keys(service_id, only_active=False):
query = Service.query.filter_by(id=service_id).options(joinedload(Service.api_keys)) stmt = (
select(Service).filter_by(id=service_id).options(joinedload(Service.api_keys))
)
if only_active: if only_active:
query = query.filter(Service.active) stmt = stmt.filter(Service.active)
return db.session.execute(stmt).scalars().unique().one()
return query.one()
def dao_fetch_all_services_by_user(user_id, only_active=False): def dao_fetch_all_services_by_user(user_id, only_active=False):
query = (
Service.query.filter(Service.users.any(id=user_id)) stmt = (
select(Service)
.filter(Service.users.any(id=user_id))
.order_by(asc(Service.created_at)) .order_by(asc(Service.created_at))
.options(joinedload(Service.users)) .options(joinedload(Service.users))
) )
if only_active: if only_active:
query = query.filter(Service.active) stmt = stmt.filter(Service.active)
return db.session.execute(stmt).scalars().unique().all()
return query.all()
def dao_fetch_all_services_created_by_user(user_id): def dao_fetch_all_services_created_by_user(user_id):
query = Service.query.filter_by(created_by_id=user_id).order_by(
asc(Service.created_at) stmt = (
select(Service)
.filter_by(created_by_id=user_id)
.order_by(asc(Service.created_at))
) )
return query.all() return db.session.execute(stmt).scalars().all()
@autocommit @autocommit
@@ -234,16 +254,15 @@ def dao_fetch_all_services_created_by_user(user_id):
VersionOptions(Template, history_class=TemplateHistory, must_write_history=False), VersionOptions(Template, history_class=TemplateHistory, must_write_history=False),
) )
def dao_archive_service(service_id): def dao_archive_service(service_id):
# have to eager load templates and api keys so that we don't flush when we loop through them stmt = (
# to ensure that db.session still contains the models when it comes to creating history objects select(Service)
service = ( .options(
Service.query.options(
joinedload(Service.templates).subqueryload(Template.template_redacted), joinedload(Service.templates).subqueryload(Template.template_redacted),
joinedload(Service.api_keys), joinedload(Service.api_keys),
) )
.filter(Service.id == service_id) .filter(Service.id == service_id)
.one()
) )
service = db.session.execute(stmt).scalars().unique().one()
service.active = False service.active = False
service.name = get_archived_db_column_value(service.name) service.name = get_archived_db_column_value(service.name)
@@ -259,11 +278,14 @@ def dao_archive_service(service_id):
def dao_fetch_service_by_id_and_user(service_id, user_id): def dao_fetch_service_by_id_and_user(service_id, user_id):
return (
Service.query.filter(Service.users.any(id=user_id), Service.id == service_id) stmt = (
select(Service)
.filter(Service.users.any(id=user_id), Service.id == service_id)
.options(joinedload(Service.users)) .options(joinedload(Service.users))
.one()
) )
result = db.session.execute(stmt).scalar_one()
return result
@autocommit @autocommit
@@ -366,39 +388,40 @@ def dao_remove_user_from_service(service, user):
def delete_service_and_all_associated_db_objects(service): def delete_service_and_all_associated_db_objects(service):
def _delete_commit(query): def _delete_commit(stmt):
query.delete(synchronize_session=False) db.session.execute(stmt)
db.session.commit() db.session.commit()
subq = db.session.query(Template.id).filter_by(service=service).subquery() subq = select(Template.id).filter_by(service=service).subquery()
_delete_commit(
TemplateRedacted.query.filter(TemplateRedacted.template_id.in_(subq))
)
_delete_commit(ServiceSmsSender.query.filter_by(service=service)) stmt = delete(TemplateRedacted).filter(TemplateRedacted.template_id.in_(subq))
_delete_commit(ServiceEmailReplyTo.query.filter_by(service=service)) _delete_commit(stmt)
_delete_commit(InvitedUser.query.filter_by(service=service))
_delete_commit(Permission.query.filter_by(service=service))
_delete_commit(NotificationHistory.query.filter_by(service=service))
_delete_commit(Notification.query.filter_by(service=service))
_delete_commit(Job.query.filter_by(service=service))
_delete_commit(Template.query.filter_by(service=service))
_delete_commit(TemplateHistory.query.filter_by(service_id=service.id))
_delete_commit(ServicePermission.query.filter_by(service_id=service.id))
_delete_commit(ApiKey.query.filter_by(service=service))
_delete_commit(ApiKey.get_history_model().query.filter_by(service_id=service.id))
_delete_commit(AnnualBilling.query.filter_by(service_id=service.id))
verify_codes = VerifyCode.query.join(User).filter( _delete_commit(delete(ServiceSmsSender).filter_by(service=service))
User.id.in_([x.id for x in service.users]) _delete_commit(delete(ServiceEmailReplyTo).filter_by(service=service))
_delete_commit(delete(InvitedUser).filter_by(service=service))
_delete_commit(delete(Permission).filter_by(service=service))
_delete_commit(delete(NotificationHistory).filter_by(service=service))
_delete_commit(delete(Notification).filter_by(service=service))
_delete_commit(delete(Job).filter_by(service=service))
_delete_commit(delete(Template).filter_by(service=service))
_delete_commit(delete(TemplateHistory).filter_by(service_id=service.id))
_delete_commit(delete(ServicePermission).filter_by(service_id=service.id))
_delete_commit(delete(ApiKey).filter_by(service=service))
_delete_commit(delete(ApiKey.get_history_model()).filter_by(service_id=service.id))
_delete_commit(delete(AnnualBilling).filter_by(service_id=service.id))
stmt = (
select(VerifyCode).join(User).filter(User.id.in_([x.id for x in service.users]))
) )
verify_codes = db.session.execute(stmt).scalars().all()
list(map(db.session.delete, verify_codes)) list(map(db.session.delete, verify_codes))
db.session.commit() db.session.commit()
users = [x for x in service.users] users = [x for x in service.users]
for user in users: for user in users:
user.organizations = [] user.organizations = []
service.users.remove(user) service.users.remove(user)
_delete_commit(Service.get_history_model().query.filter_by(id=service.id)) _delete_commit(delete(Service.get_history_model()).filter_by(id=service.id))
db.session.delete(service) db.session.delete(service)
db.session.commit() db.session.commit()
for user in users: for user in users:
@@ -409,8 +432,8 @@ def delete_service_and_all_associated_db_objects(service):
def dao_fetch_todays_stats_for_service(service_id): def dao_fetch_todays_stats_for_service(service_id):
today = utc_now().date() today = utc_now().date()
start_date = get_midnight_in_utc(today) start_date = get_midnight_in_utc(today)
return ( stmt = (
db.session.query( select(
Notification.notification_type, Notification.notification_type,
Notification.status, Notification.status,
func.count(Notification.id).label("count"), func.count(Notification.id).label("count"),
@@ -424,16 +447,16 @@ def dao_fetch_todays_stats_for_service(service_id):
Notification.notification_type, Notification.notification_type,
Notification.status, Notification.status,
) )
.all()
) )
return db.session.execute(stmt).all()
def dao_fetch_stats_for_service_from_days(service_id, start_date, end_date): def dao_fetch_stats_for_service_from_days(service_id, start_date, end_date):
start_date = get_midnight_in_utc(start_date) start_date = get_midnight_in_utc(start_date)
end_date = get_midnight_in_utc(end_date + timedelta(days=1)) end_date = get_midnight_in_utc(end_date + timedelta(days=1))
return ( stmt = (
db.session.query( select(
NotificationAllTimeView.notification_type, NotificationAllTimeView.notification_type,
NotificationAllTimeView.status, NotificationAllTimeView.status,
func.date_trunc("day", NotificationAllTimeView.created_at).label("day"), func.date_trunc("day", NotificationAllTimeView.created_at).label("day"),
@@ -450,8 +473,8 @@ def dao_fetch_stats_for_service_from_days(service_id, start_date, end_date):
NotificationAllTimeView.status, NotificationAllTimeView.status,
func.date_trunc("day", NotificationAllTimeView.created_at), func.date_trunc("day", NotificationAllTimeView.created_at),
) )
.all()
) )
return db.session.execute(stmt).scalars().all()
def dao_fetch_stats_for_service_from_days_for_user( def dao_fetch_stats_for_service_from_days_for_user(
@@ -460,13 +483,14 @@ def dao_fetch_stats_for_service_from_days_for_user(
start_date = get_midnight_in_utc(start_date) start_date = get_midnight_in_utc(start_date)
end_date = get_midnight_in_utc(end_date + timedelta(days=1)) end_date = get_midnight_in_utc(end_date + timedelta(days=1))
return ( stmt = (
db.session.query( select(
NotificationAllTimeView.notification_type, NotificationAllTimeView.notification_type,
NotificationAllTimeView.status, NotificationAllTimeView.status,
func.date_trunc("day", NotificationAllTimeView.created_at).label("day"), func.date_trunc("day", NotificationAllTimeView.created_at).label("day"),
func.count(NotificationAllTimeView.id).label("count"), func.count(NotificationAllTimeView.id).label("count"),
) )
.select_from(NotificationAllTimeView)
.filter( .filter(
NotificationAllTimeView.service_id == service_id, NotificationAllTimeView.service_id == service_id,
NotificationAllTimeView.key_type != KeyType.TEST, NotificationAllTimeView.key_type != KeyType.TEST,
@@ -479,8 +503,8 @@ def dao_fetch_stats_for_service_from_days_for_user(
NotificationAllTimeView.status, NotificationAllTimeView.status,
func.date_trunc("day", NotificationAllTimeView.created_at), func.date_trunc("day", NotificationAllTimeView.created_at),
) )
.all()
) )
return db.session.execute(stmt).scalars().all()
def dao_fetch_todays_stats_for_all_services( def dao_fetch_todays_stats_for_all_services(
@@ -491,7 +515,7 @@ def dao_fetch_todays_stats_for_all_services(
end_date = get_midnight_in_utc(today + timedelta(days=1)) end_date = get_midnight_in_utc(today + timedelta(days=1))
subquery = ( subquery = (
db.session.query( select(
Notification.notification_type, Notification.notification_type,
Notification.status, Notification.status,
Notification.service_id, Notification.service_id,
@@ -510,8 +534,8 @@ def dao_fetch_todays_stats_for_all_services(
subquery = subquery.subquery() subquery = subquery.subquery()
query = ( stmt = (
db.session.query( select(
Service.id.label("service_id"), Service.id.label("service_id"),
Service.name, Service.name,
Service.restricted, Service.restricted,
@@ -526,9 +550,9 @@ def dao_fetch_todays_stats_for_all_services(
) )
if only_active: if only_active:
query = query.filter(Service.active) stmt = stmt.filter(Service.active)
return query.all() return db.session.execute(stmt).all()
@autocommit @autocommit
@@ -537,15 +561,13 @@ def dao_fetch_todays_stats_for_all_services(
VersionOptions(Service), VersionOptions(Service),
) )
def dao_suspend_service(service_id): def dao_suspend_service(service_id):
# have to eager load api keys so that we don't flush when we loop through them
# to ensure that db.session still contains the models when it comes to creating history objects stmt = (
service = ( select(Service)
Service.query.options( .options(joinedload(Service.api_keys))
joinedload(Service.api_keys),
)
.filter(Service.id == service_id) .filter(Service.id == service_id)
.one()
) )
service = db.session.execute(stmt).scalars().unique().one()
for api_key in service.api_keys: for api_key in service.api_keys:
if not api_key.expiry_date: if not api_key.expiry_date:
@@ -557,19 +579,22 @@ def dao_suspend_service(service_id):
@autocommit @autocommit
@version_class(Service) @version_class(Service)
def dao_resume_service(service_id): def dao_resume_service(service_id):
service = Service.query.get(service_id) service = db.session.get(Service, service_id)
service.active = True service.active = True
def dao_fetch_active_users_for_service(service_id): def dao_fetch_active_users_for_service(service_id):
query = User.query.filter(User.services.any(id=service_id), User.state == "active")
return query.all() stmt = select(User).where(User.services.any(id=service_id), User.state == "active")
result = db.session.execute(stmt)
return result.scalars().all()
def dao_find_services_sending_to_tv_numbers(start_date, end_date, threshold=500): def dao_find_services_sending_to_tv_numbers(start_date, end_date, threshold=500):
return (
db.session.query( stmt = (
select(
Notification.service_id.label("service_id"), Notification.service_id.label("service_id"),
func.count(Notification.id).label("notification_count"), func.count(Notification.id).label("notification_count"),
) )
@@ -587,13 +612,13 @@ def dao_find_services_sending_to_tv_numbers(start_date, end_date, threshold=500)
Notification.service_id, Notification.service_id,
) )
.having(func.count(Notification.id) > threshold) .having(func.count(Notification.id) > threshold)
.all()
) )
return db.session.execute(stmt).all()
def dao_find_services_with_high_failure_rates(start_date, end_date, threshold=10000): def dao_find_services_with_high_failure_rates(start_date, end_date, threshold=10000):
subquery = ( subquery = (
db.session.query( select(
func.count(Notification.id).label("total_count"), func.count(Notification.id).label("total_count"),
Notification.service_id.label("service_id"), Notification.service_id.label("service_id"),
) )
@@ -614,8 +639,8 @@ def dao_find_services_with_high_failure_rates(start_date, end_date, threshold=10
subquery = subquery.subquery() subquery = subquery.subquery()
query = ( stmt = (
db.session.query( select(
Notification.service_id.label("service_id"), Notification.service_id.label("service_id"),
func.count(Notification.id).label("permanent_failure_count"), func.count(Notification.id).label("permanent_failure_count"),
subquery.c.total_count.label("total_count"), subquery.c.total_count.label("total_count"),
@@ -643,17 +668,19 @@ def dao_find_services_with_high_failure_rates(start_date, end_date, threshold=10
) )
) )
return query.all() return db.session.execute(stmt).all()
def get_live_services_with_organization(): def get_live_services_with_organization():
query = (
db.session.query( stmt = (
select(
Service.id.label("service_id"), Service.id.label("service_id"),
Service.name.label("service_name"), Service.name.label("service_name"),
Organization.id.label("organization_id"), Organization.id.label("organization_id"),
Organization.name.label("organization_name"), Organization.name.label("organization_name"),
) )
.select_from(Service)
.outerjoin(Service.organization) .outerjoin(Service.organization)
.filter( .filter(
Service.count_as_live.is_(True), Service.count_as_live.is_(True),
@@ -663,14 +690,15 @@ def get_live_services_with_organization():
.order_by(Organization.name, Service.name) .order_by(Organization.name, Service.name)
) )
return query.all() return db.session.execute(stmt).all()
def fetch_notification_stats_for_service_by_month_by_user( def fetch_notification_stats_for_service_by_month_by_user(
start_date, end_date, service_id, user_id start_date, end_date, service_id, user_id
): ):
return (
db.session.query( stmt = (
select(
func.date_trunc("month", NotificationAllTimeView.created_at).label("month"), func.date_trunc("month", NotificationAllTimeView.created_at).label("month"),
NotificationAllTimeView.notification_type, NotificationAllTimeView.notification_type,
(NotificationAllTimeView.status).label("notification_status"), (NotificationAllTimeView.status).label("notification_status"),
@@ -688,8 +716,8 @@ def fetch_notification_stats_for_service_by_month_by_user(
NotificationAllTimeView.notification_type, NotificationAllTimeView.notification_type,
NotificationAllTimeView.status, NotificationAllTimeView.status,
) )
.all()
) )
return db.session.execute(stmt).all()
def get_specific_days_stats(data, start_date, days=None, end_date=None): def get_specific_days_stats(data, start_date, days=None, end_date=None):
+7 -3
View File
@@ -1,16 +1,20 @@
from sqlalchemy import select
from app import db from app import db
from app.dao.dao_utils import autocommit from app.dao.dao_utils import autocommit
from app.models import TemplateFolder from app.models import TemplateFolder
def dao_get_template_folder_by_id_and_service_id(template_folder_id, service_id): def dao_get_template_folder_by_id_and_service_id(template_folder_id, service_id):
return TemplateFolder.query.filter( stmt = select(TemplateFolder).filter(
TemplateFolder.id == template_folder_id, TemplateFolder.service_id == service_id TemplateFolder.id == template_folder_id, TemplateFolder.service_id == service_id
).one() )
return db.session.execute(stmt).scalars().one()
def dao_get_valid_template_folders_by_id(folder_ids): def dao_get_valid_template_folders_by_id(folder_ids):
return TemplateFolder.query.filter(TemplateFolder.id.in_(folder_ids)).all() stmt = select(TemplateFolder).filter(TemplateFolder.id.in_(folder_ids))
return db.session.execute(stmt).scalars().all()
@autocommit @autocommit
+23 -17
View File
@@ -1,6 +1,6 @@
import uuid import uuid
from sqlalchemy import asc, desc from sqlalchemy import asc, desc, select
from app import db from app import db
from app.dao.dao_utils import VersionOptions, autocommit, version_class from app.dao.dao_utils import VersionOptions, autocommit, version_class
@@ -46,24 +46,29 @@ def dao_redact_template(template, user_id):
def dao_get_template_by_id_and_service_id(template_id, service_id, version=None): def dao_get_template_by_id_and_service_id(template_id, service_id, version=None):
if version is not None: if version is not None:
return TemplateHistory.query.filter_by( stmt = select(TemplateHistory).filter_by(
id=template_id, hidden=False, service_id=service_id, version=version id=template_id, hidden=False, service_id=service_id, version=version
).one() )
return Template.query.filter_by( return db.session.execute(stmt).scalars().one()
stmt = select(Template).filter_by(
id=template_id, hidden=False, service_id=service_id id=template_id, hidden=False, service_id=service_id
).one() )
return db.session.execute(stmt).scalars().one()
def dao_get_template_by_id(template_id, version=None): def dao_get_template_by_id(template_id, version=None):
if version is not None: if version is not None:
return TemplateHistory.query.filter_by(id=template_id, version=version).one() stmt = select(TemplateHistory).filter_by(id=template_id, version=version)
return Template.query.filter_by(id=template_id).one() return db.session.execute(stmt).scalars().one()
stmt = select(Template).filter_by(id=template_id)
return db.session.execute(stmt).scalars().one()
def dao_get_all_templates_for_service(service_id, template_type=None): def dao_get_all_templates_for_service(service_id, template_type=None):
if template_type is not None: if template_type is not None:
return ( stmt = (
Template.query.filter_by( select(Template)
.filter_by(
service_id=service_id, service_id=service_id,
template_type=template_type, template_type=template_type,
hidden=False, hidden=False,
@@ -73,26 +78,27 @@ def dao_get_all_templates_for_service(service_id, template_type=None):
asc(Template.name), asc(Template.name),
asc(Template.template_type), asc(Template.template_type),
) )
.all()
) )
return db.session.execute(stmt).scalars().all()
return ( stmt = (
Template.query.filter_by(service_id=service_id, hidden=False, archived=False) select(Template)
.filter_by(service_id=service_id, hidden=False, archived=False)
.order_by( .order_by(
asc(Template.name), asc(Template.name),
asc(Template.template_type), asc(Template.template_type),
) )
.all()
) )
return db.session.execute(stmt).scalars().all()
def dao_get_template_versions(service_id, template_id): def dao_get_template_versions(service_id, template_id):
return ( stmt = (
TemplateHistory.query.filter_by( select(TemplateHistory)
.filter_by(
service_id=service_id, service_id=service_id,
id=template_id, id=template_id,
hidden=False, hidden=False,
) )
.order_by(desc(TemplateHistory.version)) .order_by(desc(TemplateHistory.version))
.all()
) )
return db.session.execute(stmt).scalars().all()
+33 -22
View File
@@ -4,7 +4,7 @@ from secrets import randbelow
import sqlalchemy import sqlalchemy
from flask import current_app from flask import current_app
from sqlalchemy import func, text from sqlalchemy import delete, func, select, text
from sqlalchemy.orm import joinedload from sqlalchemy.orm import joinedload
from app import db from app import db
@@ -37,8 +37,8 @@ def get_login_gov_user(login_uuid, email_address):
login.gov uuids are. Eventually the code that checks by email address login.gov uuids are. Eventually the code that checks by email address
should be removed. should be removed.
""" """
stmt = select(User).filter_by(login_uuid=login_uuid)
user = User.query.filter_by(login_uuid=login_uuid).first() user = db.session.execute(stmt).scalars().first()
if user: if user:
if user.email_address != email_address: if user.email_address != email_address:
try: try:
@@ -54,7 +54,8 @@ def get_login_gov_user(login_uuid, email_address):
return user return user
# Remove this 1 July 2025, all users should have login.gov uuids by now # Remove this 1 July 2025, all users should have login.gov uuids by now
user = User.query.filter(User.email_address.ilike(email_address)).first() stmt = select(User).filter(User.email_address.ilike(email_address))
user = db.session.execute(stmt).scalars().first()
if user: if user:
save_user_attribute(user, {"login_uuid": login_uuid}) save_user_attribute(user, {"login_uuid": login_uuid})
@@ -102,24 +103,27 @@ def create_user_code(user, code, code_type):
def get_user_code(user, code, code_type): def get_user_code(user, code, code_type):
# Get the most recent codes to try and reduce the # Get the most recent codes to try and reduce the
# time searching for the correct code. # time searching for the correct code.
codes = VerifyCode.query.filter_by(user=user, code_type=code_type).order_by( stmt = (
VerifyCode.created_at.desc() select(VerifyCode)
.filter_by(user=user, code_type=code_type)
.order_by(VerifyCode.created_at.desc())
) )
codes = db.session.execute(stmt).scalars().all()
return next((x for x in codes if x.check_code(code)), None) return next((x for x in codes if x.check_code(code)), None)
def delete_codes_older_created_more_than_a_day_ago(): def delete_codes_older_created_more_than_a_day_ago():
deleted = ( stmt = delete(VerifyCode).filter(
db.session.query(VerifyCode) VerifyCode.created_at < utc_now() - timedelta(hours=24)
.filter(VerifyCode.created_at < utc_now() - timedelta(hours=24))
.delete()
) )
deleted = db.session.execute(stmt)
db.session.commit() db.session.commit()
return deleted return deleted
def use_user_code(id): def use_user_code(id):
verify_code = VerifyCode.query.get(id) verify_code = db.session.get(VerifyCode, id)
verify_code.code_used = True verify_code.code_used = True
db.session.add(verify_code) db.session.add(verify_code)
db.session.commit() db.session.commit()
@@ -131,36 +135,42 @@ def delete_model_user(user):
def delete_user_verify_codes(user): def delete_user_verify_codes(user):
VerifyCode.query.filter_by(user=user).delete() stmt = delete(VerifyCode).filter_by(user=user)
db.session.execute(stmt)
db.session.commit() db.session.commit()
def count_user_verify_codes(user): def count_user_verify_codes(user):
query = VerifyCode.query.filter( stmt = select(func.count(VerifyCode.id)).filter(
VerifyCode.user == user, VerifyCode.user == user,
VerifyCode.expiry_datetime > utc_now(), VerifyCode.expiry_datetime > utc_now(),
VerifyCode.code_used.is_(False), VerifyCode.code_used.is_(False),
) )
return query.count() result = db.session.execute(stmt).scalar()
return result or 0
def get_user_by_id(user_id=None): def get_user_by_id(user_id=None):
if user_id: if user_id:
return User.query.filter_by(id=user_id).one() stmt = select(User).filter_by(id=user_id)
return User.query.filter_by().all() return db.session.execute(stmt).scalars().one()
return get_users()
def get_users(): def get_users():
return User.query.all() stmt = select(User)
return db.session.execute(stmt).scalars().all()
def get_user_by_email(email): def get_user_by_email(email):
return User.query.filter(func.lower(User.email_address) == func.lower(email)).one() stmt = select(User).filter(func.lower(User.email_address) == func.lower(email))
return db.session.execute(stmt).scalars().one()
def get_users_by_partial_email(email): def get_users_by_partial_email(email):
email = escape_special_characters(email) email = escape_special_characters(email)
return User.query.filter(User.email_address.ilike("%{}%".format(email))).all() stmt = select(User).filter(User.email_address.ilike("%{}%".format(email)))
return db.session.execute(stmt).scalars().all()
def increment_failed_login_count(user): def increment_failed_login_count(user):
@@ -188,16 +198,17 @@ def get_user_and_accounts(user_id):
# TODO: With sqlalchemy 2.0 change as below because of the breaking change # TODO: With sqlalchemy 2.0 change as below because of the breaking change
# at User.organizations.services, we need to verify that the below subqueryload # at User.organizations.services, we need to verify that the below subqueryload
# that we have put is functionally doing the same thing as before # that we have put is functionally doing the same thing as before
return ( stmt = (
User.query.filter(User.id == user_id) select(User)
.filter(User.id == user_id)
.options( .options(
# eagerly load the user's services and organizations, and also the service's org and vice versa # eagerly load the user's services and organizations, and also the service's org and vice versa
# (so we can see if the user knows about it) # (so we can see if the user knows about it)
joinedload(User.services).joinedload(Service.organization), joinedload(User.services).joinedload(Service.organization),
joinedload(User.organizations).subqueryload(Organization.services), joinedload(User.organizations).subqueryload(Organization.services),
) )
.one()
) )
return db.session.execute(stmt).scalars().unique().one()
@autocommit @autocommit
@@ -31,10 +31,10 @@ def upgrade():
# #
# go_live = datetime.datetime.strptime('2016-05-18', '%Y-%m-%d') # go_live = datetime.datetime.strptime('2016-05-18', '%Y-%m-%d')
# notifications_history_start_date = datetime.datetime.strptime('2016-06-26 23:21:55', '%Y-%m-%d %H:%M:%S') # notifications_history_start_date = datetime.datetime.strptime('2016-06-26 23:21:55', '%Y-%m-%d %H:%M:%S')
# jobs = session.query(Job).join(Template).filter(Job.service_id == '95316ff0-e555-462d-a6e7-95d26fbfd091', # stmt = select(Job).join(Template).filter(Job.service_id == '95316ff0-e555-462d-a6e7-95d26fbfd091',
# Job.created_at >= go_live, # Job.created_at >= go_live,
# Job.created_at < notifications_history_start_date).all() # Job.created_at < notifications_history_start_date).all()
# # jobs = db.session.execute(stmt).scalars().all()
# for job in jobs: # for job in jobs:
# for i in range(0, job.notifications_delivered): # for i in range(0, job.notifications_delivered):
# notification = NotificationHistory(id=uuid.uuid4(), # notification = NotificationHistory(id=uuid.uuid4(),
@@ -76,12 +76,11 @@ def downgrade():
# #
# go_live = datetime.datetime.strptime('2016-05-18', '%Y-%m-%d') # go_live = datetime.datetime.strptime('2016-05-18', '%Y-%m-%d')
# notifications_history_start_date = datetime.datetime.strptime('2016-06-26 23:21:55', '%Y-%m-%d %H:%M:%S') # notifications_history_start_date = datetime.datetime.strptime('2016-06-26 23:21:55', '%Y-%m-%d %H:%M:%S')
# # stmt = delete(NotificationHistory).where(
# session.query(NotificationHistory).filter(
# NotificationHistory.created_at >= go_live, # NotificationHistory.created_at >= go_live,
# NotificationHistory.service_id == '95316ff0-e555-462d-a6e7-95d26fbfd091', # NotificationHistory.service_id == '95316ff0-e555-462d-a6e7-95d26fbfd091',
# NotificationHistory.created_at < notifications_history_start_date).delete() # NotificationHistory.created_at < notifications_history_start_date)
# # session.execute(stmt)
# session.commit() # session.commit()
# ### end Alembic commands ### # ### end Alembic commands ###
pass pass
Generated
+4 -4
View File
@@ -4519,13 +4519,13 @@ test = ["websockets"]
[[package]] [[package]]
name = "werkzeug" name = "werkzeug"
version = "3.0.3" version = "3.0.6"
description = "The comprehensive WSGI web application library." description = "The comprehensive WSGI web application library."
optional = false optional = false
python-versions = ">=3.8" python-versions = ">=3.8"
files = [ files = [
{file = "werkzeug-3.0.3-py3-none-any.whl", hash = "sha256:fc9645dc43e03e4d630d23143a04a7f947a9a3b5727cd535fdfe155a17cc48c8"}, {file = "werkzeug-3.0.6-py3-none-any.whl", hash = "sha256:1bc0c2310d2fbb07b1dd1105eba2f7af72f322e1e455f2f93c993bee8c8a5f17"},
{file = "werkzeug-3.0.3.tar.gz", hash = "sha256:097e5bfda9f0aba8da6b8545146def481d06aa7d3266e7448e2cccf67dd8bd18"}, {file = "werkzeug-3.0.6.tar.gz", hash = "sha256:a8dd59d4de28ca70471a34cba79bed5f7ef2e036a76b3ab0835474246eb41f8d"},
] ]
[package.dependencies] [package.dependencies]
@@ -4803,4 +4803,4 @@ multidict = ">=4.0"
[metadata] [metadata]
lock-version = "2.0" lock-version = "2.0"
python-versions = "^3.12.2" python-versions = "^3.12.2"
content-hash = "42172a923e16c5b0965ab06f717d41e8491ee35f7be674091b38014c48b7a89e" content-hash = "cf18ae74630e47eec18cc6c5fea9e554476809d20589d82c54a8d761bb2c3de0"
+1 -1
View File
@@ -47,7 +47,7 @@ psycopg2-binary = "==2.9.9"
pyjwt = "==2.8.0" pyjwt = "==2.8.0"
python-dotenv = "==1.0.1" python-dotenv = "==1.0.1"
sqlalchemy = "==2.0.31" sqlalchemy = "==2.0.31"
werkzeug = "^3.0.3" werkzeug = "^3.0.6"
faker = "^26.0.0" faker = "^26.0.0"
async-timeout = "^4.0.3" async-timeout = "^4.0.3"
bleach = "^6.1.0" bleach = "^6.1.0"
@@ -4,9 +4,11 @@ from functools import partial
import pytest import pytest
from freezegun import freeze_time from freezegun import freeze_time
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
from app import db
from app.dao.notifications_dao import ( from app.dao.notifications_dao import (
dao_create_notification, dao_create_notification,
dao_delete_notifications_by_id, dao_delete_notifications_by_id,
@@ -55,7 +57,10 @@ def test_should_by_able_to_update_status_by_reference(
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.get(notification.id).status == NotificationStatus.SENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.SENDING
)
notification.reference = "reference" notification.reference = "reference"
dao_update_notification(notification) dao_update_notification(notification)
@@ -64,7 +69,8 @@ def test_should_by_able_to_update_status_by_reference(
) )
assert updated.status == NotificationStatus.DELIVERED assert updated.status == NotificationStatus.DELIVERED
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
@@ -81,7 +87,10 @@ def test_should_by_able_to_update_status_by_id(
dao_create_notification(notification) dao_create_notification(notification)
assert notification.status == NotificationStatus.SENDING assert notification.status == NotificationStatus.SENDING
assert Notification.query.get(notification.id).status == NotificationStatus.SENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.SENDING
)
with freeze_time("2000-01-02 12:00:00"): with freeze_time("2000-01-02 12:00:00"):
updated = update_notification_status_by_id( updated = update_notification_status_by_id(
@@ -92,7 +101,8 @@ def test_should_by_able_to_update_status_by_id(
assert updated.status == NotificationStatus.DELIVERED assert updated.status == NotificationStatus.DELIVERED
assert updated.updated_at == datetime(2000, 1, 2, 12, 0, 0) assert updated.updated_at == datetime(2000, 1, 2, 12, 0, 0)
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
assert notification.updated_at == datetime(2000, 1, 2, 12, 0, 0) assert notification.updated_at == datetime(2000, 1, 2, 12, 0, 0)
assert notification.status == NotificationStatus.DELIVERED assert notification.status == NotificationStatus.DELIVERED
@@ -107,15 +117,17 @@ def test_should_not_update_status_by_id_if_not_sending_and_does_not_update_job(
job=sample_job, job=sample_job,
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
assert not update_notification_status_by_id( assert not update_notification_status_by_id(
notification.id, NotificationStatus.FAILED notification.id, NotificationStatus.FAILED
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
assert sample_job == Job.query.get(notification.job_id) assert sample_job == db.session.get(Job, notification.job_id)
def test_should_not_update_status_by_reference_if_not_sending_and_does_not_update_job( def test_should_not_update_status_by_reference_if_not_sending_and_does_not_update_job(
@@ -128,20 +140,22 @@ def test_should_not_update_status_by_reference_if_not_sending_and_does_not_updat
job=sample_job, job=sample_job,
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
assert not update_notification_status_by_reference( assert not update_notification_status_by_reference(
"reference", NotificationStatus.FAILED "reference", NotificationStatus.FAILED
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
assert sample_job == Job.query.get(notification.job_id) assert sample_job == db.session.get(Job, notification.job_id)
def test_should_update_status_by_id_if_created(sample_template, sample_notification): def test_should_update_status_by_id_if_created(sample_template, sample_notification):
assert ( assert (
Notification.query.get(sample_notification.id).status db.session.get(Notification, sample_notification.id).status
== NotificationStatus.CREATED == NotificationStatus.CREATED
) )
updated = update_notification_status_by_id( updated = update_notification_status_by_id(
@@ -149,7 +163,7 @@ def test_should_update_status_by_id_if_created(sample_template, sample_notificat
NotificationStatus.FAILED, NotificationStatus.FAILED,
) )
assert ( assert (
Notification.query.get(sample_notification.id).status db.session.get(Notification, sample_notification.id).status
== NotificationStatus.FAILED == NotificationStatus.FAILED
) )
assert updated.status == NotificationStatus.FAILED assert updated.status == NotificationStatus.FAILED
@@ -244,11 +258,17 @@ def test_should_not_update_status_by_reference_if_not_sending(sample_template):
status=NotificationStatus.CREATED, status=NotificationStatus.CREATED,
reference="reference", reference="reference",
) )
assert Notification.query.get(notification.id).status == NotificationStatus.CREATED assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.CREATED
)
updated = update_notification_status_by_reference( updated = update_notification_status_by_reference(
"reference", NotificationStatus.FAILED "reference", NotificationStatus.FAILED
) )
assert Notification.query.get(notification.id).status == NotificationStatus.CREATED assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.CREATED
)
assert not updated assert not updated
@@ -264,14 +284,18 @@ def test_should_by_able_to_update_status_by_id_from_pending_to_delivered(
assert update_notification_status_by_id( assert update_notification_status_by_id(
notification_id=notification.id, status=NotificationStatus.PENDING notification_id=notification.id, status=NotificationStatus.PENDING
) )
assert Notification.query.get(notification.id).status == NotificationStatus.PENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.PENDING
)
assert update_notification_status_by_id( assert update_notification_status_by_id(
notification.id, notification.id,
NotificationStatus.DELIVERED, NotificationStatus.DELIVERED,
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
@@ -289,7 +313,10 @@ def test_should_by_able_to_update_status_by_id_from_pending_to_temporary_failure
notification_id=notification.id, notification_id=notification.id,
status=NotificationStatus.PENDING, status=NotificationStatus.PENDING,
) )
assert Notification.query.get(notification.id).status == NotificationStatus.PENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.PENDING
)
assert update_notification_status_by_id( assert update_notification_status_by_id(
notification.id, notification.id,
@@ -297,7 +324,7 @@ def test_should_by_able_to_update_status_by_id_from_pending_to_temporary_failure
) )
assert ( assert (
Notification.query.get(notification.id).status db.session.get(Notification, notification.id).status
== NotificationStatus.TEMPORARY_FAILURE == NotificationStatus.TEMPORARY_FAILURE
) )
@@ -312,14 +339,17 @@ def test_should_by_able_to_update_status_by_id_from_sending_to_permanent_failure
) )
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.get(notification.id).status == NotificationStatus.SENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.SENDING
)
assert update_notification_status_by_id( assert update_notification_status_by_id(
notification.id, notification.id,
status=NotificationStatus.PERMANENT_FAILURE, status=NotificationStatus.PERMANENT_FAILURE,
) )
assert ( assert (
Notification.query.get(notification.id).status db.session.get(Notification, notification.id).status
== NotificationStatus.PERMANENT_FAILURE == NotificationStatus.PERMANENT_FAILURE
) )
@@ -331,7 +361,10 @@ def test_should_not_update_status_once_notification_status_is_delivered(
template=sample_email_template, template=sample_email_template,
status=NotificationStatus.SENDING, status=NotificationStatus.SENDING,
) )
assert Notification.query.get(notification.id).status == NotificationStatus.SENDING assert (
db.session.get(Notification, notification.id).status
== NotificationStatus.SENDING
)
notification.reference = "reference" notification.reference = "reference"
dao_update_notification(notification) dao_update_notification(notification)
@@ -340,7 +373,8 @@ def test_should_not_update_status_once_notification_status_is_delivered(
NotificationStatus.DELIVERED, NotificationStatus.DELIVERED,
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
update_notification_status_by_reference( update_notification_status_by_reference(
@@ -348,7 +382,8 @@ def test_should_not_update_status_once_notification_status_is_delivered(
NotificationStatus.FAILED, NotificationStatus.FAILED,
) )
assert ( assert (
Notification.query.get(notification.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification.id).status
== NotificationStatus.DELIVERED
) )
@@ -370,7 +405,7 @@ def test_create_notification_creates_notification_with_personalisation(
sample_template_with_placeholders, sample_template_with_placeholders,
sample_job, sample_job,
): ):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = create_notification( data = create_notification(
template=sample_template_with_placeholders, template=sample_template_with_placeholders,
@@ -379,8 +414,8 @@ def test_create_notification_creates_notification_with_personalisation(
status=NotificationStatus.CREATED, status=NotificationStatus.CREATED,
) )
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert data.to == notification_from_db.to assert data.to == notification_from_db.to
assert data.job_id == notification_from_db.job_id assert data.job_id == notification_from_db.job_id
@@ -393,15 +428,15 @@ def test_create_notification_creates_notification_with_personalisation(
def test_save_notification_creates_sms(sample_template, sample_job): def test_save_notification_creates_sms(sample_template, sample_job):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template, job_id=sample_job.id) data = _notification_json(sample_template, job_id=sample_job.id)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["job_id"] == notification_from_db.job_id assert data["job_id"] == notification_from_db.job_id
@@ -412,16 +447,36 @@ def test_save_notification_creates_sms(sample_template, sample_job):
assert notification_from_db.status == NotificationStatus.CREATED assert notification_from_db.status == NotificationStatus.CREATED
def _get_notification_query_all():
stmt = select(Notification)
return db.session.execute(stmt).scalars().all()
def _get_notification_query_one():
stmt = select(Notification)
return db.session.execute(stmt).scalars().one()
def _get_notification_query_count():
stmt = select(func.count(Notification.id))
return db.session.execute(stmt).scalar() or 0
def _get_notification_history_query_count():
stmt = select(func.count(NotificationHistory.id))
return db.session.execute(stmt).scalar() or 0
def test_save_notification_and_create_email(sample_email_template, sample_job): def test_save_notification_and_create_email(sample_email_template, sample_job):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_email_template, job_id=sample_job.id) data = _notification_json(sample_email_template, job_id=sample_job.id)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["job_id"] == notification_from_db.job_id assert data["job_id"] == notification_from_db.job_id
@@ -433,29 +488,29 @@ def test_save_notification_and_create_email(sample_email_template, sample_job):
def test_save_notification(sample_email_template, sample_job): def test_save_notification(sample_email_template, sample_job):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_email_template, job_id=sample_job.id) data = _notification_json(sample_email_template, job_id=sample_job.id)
notification_1 = Notification(**data) notification_1 = Notification(**data)
notification_2 = Notification(**data) notification_2 = Notification(**data)
dao_create_notification(notification_1) dao_create_notification(notification_1)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
dao_create_notification(notification_2) dao_create_notification(notification_2)
assert Notification.query.count() == 2 assert _get_notification_query_count() == 2
def test_save_notification_does_not_creates_history(sample_email_template, sample_job): def test_save_notification_does_not_creates_history(sample_email_template, sample_job):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_email_template, job_id=sample_job.id) data = _notification_json(sample_email_template, job_id=sample_job.id)
notification_1 = Notification(**data) notification_1 = Notification(**data)
dao_create_notification(notification_1) dao_create_notification(notification_1)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
assert NotificationHistory.query.count() == 0 assert _get_notification_history_query_count() == 0
def test_update_notification_with_research_mode_service_does_not_create_or_update_history( def test_update_notification_with_research_mode_service_does_not_create_or_update_history(
@@ -464,14 +519,14 @@ def test_update_notification_with_research_mode_service_does_not_create_or_updat
sample_template.service.research_mode = True sample_template.service.research_mode = True
notification = create_notification(template=sample_template) notification = create_notification(template=sample_template)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
assert NotificationHistory.query.count() == 0 assert _get_notification_history_query_count() == 0
notification.status = NotificationStatus.DELIVERED notification.status = NotificationStatus.DELIVERED
dao_update_notification(notification) dao_update_notification(notification)
assert Notification.query.one().status == NotificationStatus.DELIVERED assert _get_notification_query_one().status == NotificationStatus.DELIVERED
assert NotificationHistory.query.count() == 0 assert _get_notification_history_query_count() == 0
def test_not_save_notification_and_not_create_stats_on_commit_error( def test_not_save_notification_and_not_create_stats_on_commit_error(
@@ -479,26 +534,26 @@ def test_not_save_notification_and_not_create_stats_on_commit_error(
): ):
random_id = str(uuid.uuid4()) random_id = str(uuid.uuid4())
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template, job_id=random_id) data = _notification_json(sample_template, job_id=random_id)
notification = Notification(**data) notification = Notification(**data)
with pytest.raises(SQLAlchemyError): with pytest.raises(SQLAlchemyError):
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
assert Job.query.get(sample_job.id).notifications_sent == 0 assert db.session.get(Job, sample_job.id).notifications_sent == 0
def test_save_notification_and_increment_job(sample_template, sample_job, sns_provider): def test_save_notification_and_increment_job(sample_template, sample_job, sns_provider):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template, job_id=sample_job.id) data = _notification_json(sample_template, job_id=sample_job.id)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["job_id"] == notification_from_db.job_id assert data["job_id"] == notification_from_db.job_id
@@ -510,21 +565,21 @@ def test_save_notification_and_increment_job(sample_template, sample_job, sns_pr
notification_2 = Notification(**data) notification_2 = Notification(**data)
dao_create_notification(notification_2) dao_create_notification(notification_2)
assert Notification.query.count() == 2 assert _get_notification_query_count() == 2
def test_save_notification_and_increment_correct_job(sample_template, sns_provider): def test_save_notification_and_increment_correct_job(sample_template, sns_provider):
job_1 = create_job(sample_template) job_1 = create_job(sample_template)
job_2 = create_job(sample_template) job_2 = create_job(sample_template)
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template, job_id=job_1.id) data = _notification_json(sample_template, job_id=job_1.id)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["job_id"] == notification_from_db.job_id assert data["job_id"] == notification_from_db.job_id
@@ -537,14 +592,14 @@ def test_save_notification_and_increment_correct_job(sample_template, sns_provid
def test_save_notification_with_no_job(sample_template, sns_provider): def test_save_notification_with_no_job(sample_template, sns_provider):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template) data = _notification_json(sample_template)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["service"] == notification_from_db.service assert data["service"] == notification_from_db.service
@@ -592,7 +647,7 @@ def test_get_notification_by_id_when_notification_exists_for_different_service(
def test_get_notifications_by_reference(sample_template): def test_get_notifications_by_reference(sample_template):
client_reference = "some-client-ref" client_reference = "some-client-ref"
assert len(Notification.query.all()) == 0 assert len(_get_notification_query_all()) == 0
create_notification(sample_template, client_reference=client_reference) create_notification(sample_template, client_reference=client_reference)
create_notification(sample_template, client_reference=client_reference) create_notification(sample_template, client_reference=client_reference)
create_notification(sample_template, client_reference="other-ref") create_notification(sample_template, client_reference="other-ref")
@@ -603,14 +658,14 @@ def test_get_notifications_by_reference(sample_template):
def test_save_notification_no_job_id(sample_template): def test_save_notification_no_job_id(sample_template):
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
data = _notification_json(sample_template) data = _notification_json(sample_template)
notification = Notification(**data) notification = Notification(**data)
dao_create_notification(notification) dao_create_notification(notification)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
notification_from_db = Notification.query.all()[0] notification_from_db = _get_notification_query_all()[0]
assert notification_from_db.id assert notification_from_db.id
assert "1" == notification_from_db.to assert "1" == notification_from_db.to
assert data["service"] == notification_from_db.service assert data["service"] == notification_from_db.service
@@ -687,13 +742,13 @@ def test_update_notification_sets_status(sample_notification):
assert sample_notification.status == NotificationStatus.CREATED assert sample_notification.status == NotificationStatus.CREATED
sample_notification.status = NotificationStatus.FAILED sample_notification.status = NotificationStatus.FAILED
dao_update_notification(sample_notification) dao_update_notification(sample_notification)
notification_from_db = Notification.query.get(sample_notification.id) notification_from_db = db.session.get(Notification, sample_notification.id)
assert notification_from_db.status == NotificationStatus.FAILED assert notification_from_db.status == NotificationStatus.FAILED
@freeze_time("2016-01-10") @freeze_time("2016-01-10")
def test_should_limit_notifications_return_by_day_limit_plus_one(sample_template): def test_should_limit_notifications_return_by_day_limit_plus_one(sample_template):
assert len(Notification.query.all()) == 0 assert len(_get_notification_query_all()) == 0
# create one notification a day between 1st and 9th, # create one notification a day between 1st and 9th,
# with assumption that the local timezone is EST # with assumption that the local timezone is EST
@@ -706,7 +761,7 @@ def test_should_limit_notifications_return_by_day_limit_plus_one(sample_template
status=NotificationStatus.FAILED, status=NotificationStatus.FAILED,
) )
all_notifications = Notification.query.all() all_notifications = _get_notification_query_all()
assert len(all_notifications) == 10 assert len(all_notifications) == 10
all_notifications = get_notifications_for_service( all_notifications = get_notifications_for_service(
@@ -722,19 +777,19 @@ def test_should_limit_notifications_return_by_day_limit_plus_one(sample_template
def test_creating_notification_does_not_add_notification_history(sample_template): def test_creating_notification_does_not_add_notification_history(sample_template):
create_notification(template=sample_template) create_notification(template=sample_template)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
assert NotificationHistory.query.count() == 0 assert _get_notification_history_query_count() == 0
def test_should_delete_notification_for_id(sample_template): def test_should_delete_notification_for_id(sample_template):
notification = create_notification(template=sample_template) notification = create_notification(template=sample_template)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
assert NotificationHistory.query.count() == 0 assert _get_notification_history_query_count() == 0
dao_delete_notifications_by_id(notification.id) dao_delete_notifications_by_id(notification.id)
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
def test_should_delete_notification_and_ignore_history_for_research_mode( def test_should_delete_notification_and_ignore_history_for_research_mode(
@@ -744,31 +799,32 @@ def test_should_delete_notification_and_ignore_history_for_research_mode(
notification = create_notification(template=sample_template) notification = create_notification(template=sample_template)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
dao_delete_notifications_by_id(notification.id) dao_delete_notifications_by_id(notification.id)
assert Notification.query.count() == 0 assert _get_notification_query_count() == 0
def test_should_delete_only_notification_with_id(sample_template): def test_should_delete_only_notification_with_id(sample_template):
notification_1 = create_notification(template=sample_template) notification_1 = create_notification(template=sample_template)
notification_2 = create_notification(template=sample_template) notification_2 = create_notification(template=sample_template)
assert Notification.query.count() == 2 assert _get_notification_query_count() == 2
dao_delete_notifications_by_id(notification_1.id) dao_delete_notifications_by_id(notification_1.id)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
assert Notification.query.first().id == notification_2.id stmt = select(Notification)
assert db.session.execute(stmt).scalars().first().id == notification_2.id
def test_should_delete_no_notifications_if_no_matching_ids(sample_template): def test_should_delete_no_notifications_if_no_matching_ids(sample_template):
create_notification(template=sample_template) create_notification(template=sample_template)
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
dao_delete_notifications_by_id(uuid.uuid4()) dao_delete_notifications_by_id(uuid.uuid4())
assert Notification.query.count() == 1 assert _get_notification_query_count() == 1
def _notification_json(sample_template, job_id=None, id=None, status=None): def _notification_json(sample_template, job_id=None, id=None, status=None):
@@ -814,16 +870,19 @@ def test_dao_timeout_notifications(sample_template):
temporary_failure_notifications = dao_timeout_notifications(utc_now()) temporary_failure_notifications = dao_timeout_notifications(utc_now())
assert len(temporary_failure_notifications) == 2 assert len(temporary_failure_notifications) == 2
assert Notification.query.get(created.id).status == NotificationStatus.CREATED assert db.session.get(Notification, created.id).status == NotificationStatus.CREATED
assert ( assert (
Notification.query.get(sending.id).status db.session.get(Notification, sending.id).status
== NotificationStatus.TEMPORARY_FAILURE == NotificationStatus.TEMPORARY_FAILURE
) )
assert ( assert (
Notification.query.get(pending.id).status db.session.get(Notification, pending.id).status
== NotificationStatus.TEMPORARY_FAILURE == NotificationStatus.TEMPORARY_FAILURE
) )
assert Notification.query.get(delivered.id).status == NotificationStatus.DELIVERED assert (
db.session.get(Notification, delivered.id).status
== NotificationStatus.DELIVERED
)
def test_dao_timeout_notifications_only_updates_for_older_notifications( def test_dao_timeout_notifications_only_updates_for_older_notifications(
@@ -842,8 +901,8 @@ def test_dao_timeout_notifications_only_updates_for_older_notifications(
temporary_failure_notifications = dao_timeout_notifications(utc_now()) temporary_failure_notifications = dao_timeout_notifications(utc_now())
assert len(temporary_failure_notifications) == 0 assert len(temporary_failure_notifications) == 0
assert Notification.query.get(sending.id).status == NotificationStatus.SENDING assert db.session.get(Notification, sending.id).status == NotificationStatus.SENDING
assert Notification.query.get(pending.id).status == NotificationStatus.PENDING assert db.session.get(Notification, pending.id).status == NotificationStatus.PENDING
def test_should_return_notifications_excluding_jobs_by_default( def test_should_return_notifications_excluding_jobs_by_default(
@@ -935,7 +994,7 @@ def test_get_notifications_created_by_api_or_csv_are_returned_correctly_excludin
key_type=sample_test_api_key.key_type, key_type=sample_test_api_key.key_type,
) )
all_notifications = Notification.query.all() all_notifications = _get_notification_query_all()
assert len(all_notifications) == 4 assert len(all_notifications) == 4
# returns all real API derived notifications # returns all real API derived notifications
@@ -982,7 +1041,7 @@ def test_get_notifications_with_a_live_api_key_type(
key_type=sample_test_api_key.key_type, key_type=sample_test_api_key.key_type,
) )
all_notifications = Notification.query.all() all_notifications = _get_notification_query_all()
assert len(all_notifications) == 4 assert len(all_notifications) == 4
# only those created with normal API key, no jobs # only those created with normal API key, no jobs
@@ -1114,7 +1173,7 @@ def test_should_exclude_test_key_notifications_by_default(
key_type=sample_test_api_key.key_type, key_type=sample_test_api_key.key_type,
) )
all_notifications = Notification.query.all() all_notifications = _get_notification_query_all()
assert len(all_notifications) == 4 assert len(all_notifications) == 4
all_notifications = get_notifications_for_service( all_notifications = get_notifications_for_service(
@@ -1757,10 +1816,10 @@ def test_dao_update_notifications_by_reference_updated_notifications(sample_temp
update_dict={"status": NotificationStatus.DELIVERED, "billable_units": 2}, update_dict={"status": NotificationStatus.DELIVERED, "billable_units": 2},
) )
assert updated_count == 2 assert updated_count == 2
updated_1 = Notification.query.get(notification_1.id) updated_1 = db.session.get(Notification, notification_1.id)
assert updated_1.billable_units == 2 assert updated_1.billable_units == 2
assert updated_1.status == NotificationStatus.DELIVERED assert updated_1.status == NotificationStatus.DELIVERED
updated_2 = Notification.query.get(notification_2.id) updated_2 = db.session.get(Notification, notification_2.id)
assert updated_2.billable_units == 2 assert updated_2.billable_units == 2
assert updated_2.status == NotificationStatus.DELIVERED assert updated_2.status == NotificationStatus.DELIVERED
@@ -1823,10 +1882,11 @@ def test_dao_update_notifications_by_reference_updates_history_when_one_of_two_n
assert updated_count == 1 assert updated_count == 1
assert updated_history_count == 1 assert updated_history_count == 1
assert ( assert (
Notification.query.get(notification2.id).status == NotificationStatus.DELIVERED db.session.get(Notification, notification2.id).status
== NotificationStatus.DELIVERED
) )
assert ( assert (
NotificationHistory.query.get(notification1.id).status db.session.get(NotificationHistory, notification1.id).status
== NotificationStatus.DELIVERED == NotificationStatus.DELIVERED
) )
+17 -12
View File
@@ -1,6 +1,7 @@
import uuid import uuid
import pytest import pytest
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError, SQLAlchemyError from sqlalchemy.exc import IntegrityError, SQLAlchemyError
from app import db from app import db
@@ -57,7 +58,8 @@ def test_get_organization_by_id_gets_correct_organization(notify_db_session):
def test_update_organization(notify_db_session): def test_update_organization(notify_db_session):
create_organization() create_organization()
organization = Organization.query.one() stmt = select(Organization)
organization = db.session.execute(stmt).scalars().one()
user = create_user() user = create_user()
email_branding = create_email_branding() email_branding = create_email_branding()
@@ -78,7 +80,8 @@ def test_update_organization(notify_db_session):
dao_update_organization(organization.id, **data) dao_update_organization(organization.id, **data)
organization = Organization.query.one() stmt = select(Organization)
organization = db.session.execute(stmt).scalars().one()
for attribute, value in data.items(): for attribute, value in data.items():
assert getattr(organization, attribute) == value assert getattr(organization, attribute) == value
@@ -102,7 +105,8 @@ def test_update_organization_domains_lowercases(
): ):
create_organization() create_organization()
organization = Organization.query.one() stmt = select(Organization)
organization = db.session.execute(stmt).scalars().one()
# Seed some domains # Seed some domains
dao_update_organization(organization.id, domains=["123", "456"]) dao_update_organization(organization.id, domains=["123", "456"])
@@ -121,7 +125,8 @@ def test_update_organization_domains_lowercases_integrity_error(
): ):
create_organization() create_organization()
organization = Organization.query.one() stmt = select(Organization)
organization = db.session.execute(stmt).scalars().one()
# Seed some domains # Seed some domains
dao_update_organization(organization.id, domains=["123", "456"]) dao_update_organization(organization.id, domains=["123", "456"])
@@ -175,11 +180,11 @@ def test_update_organization_updates_the_service_org_type_if_org_type_is_provide
assert sample_organization.organization_type == OrganizationType.FEDERAL assert sample_organization.organization_type == OrganizationType.FEDERAL
assert sample_service.organization_type == OrganizationType.FEDERAL assert sample_service.organization_type == OrganizationType.FEDERAL
stmt = select(Service.get_history_model()).filter_by(
id=sample_service.id, version=2
)
assert ( assert (
Service.get_history_model() db.session.execute(stmt).scalars().one().organization_type
.query.filter_by(id=sample_service.id, version=2)
.one()
.organization_type
== OrganizationType.FEDERAL == OrganizationType.FEDERAL
) )
@@ -229,11 +234,11 @@ def test_add_service_to_organization(sample_service, sample_organization):
assert sample_organization.services[0].id == sample_service.id assert sample_organization.services[0].id == sample_service.id
assert sample_service.organization_type == sample_organization.organization_type assert sample_service.organization_type == sample_organization.organization_type
stmt = select(Service.get_history_model()).filter_by(
id=sample_service.id, version=2
)
assert ( assert (
Service.get_history_model() db.session.execute(stmt).scalars().one().organization_type
.query.filter_by(id=sample_service.id, version=2)
.one()
.organization_type
== sample_organization.organization_type == sample_organization.organization_type
) )
assert sample_service.organization_id == sample_organization.id assert sample_service.organization_id == sample_organization.id
+16 -13
View File
@@ -1,8 +1,10 @@
import uuid import uuid
import pytest import pytest
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.exc import SQLAlchemyError
from app import db
from app.dao.service_sms_sender_dao import ( from app.dao.service_sms_sender_dao import (
archive_sms_sender, archive_sms_sender,
dao_add_sms_sender_for_service, dao_add_sms_sender_for_service,
@@ -97,10 +99,8 @@ def test_dao_add_sms_sender_for_service(notify_db_session):
is_default=False, is_default=False,
inbound_number_id=None, inbound_number_id=None,
) )
stmt = select(ServiceSmsSender).order_by(ServiceSmsSender.created_at)
service_sms_senders = ServiceSmsSender.query.order_by( service_sms_senders = db.session.execute(stmt).scalars().all()
ServiceSmsSender.created_at
).all()
assert len(service_sms_senders) == 2 assert len(service_sms_senders) == 2
assert service_sms_senders[0].sms_sender == "testing" assert service_sms_senders[0].sms_sender == "testing"
assert service_sms_senders[0].is_default assert service_sms_senders[0].is_default
@@ -116,10 +116,8 @@ def test_dao_add_sms_sender_for_service_switches_default(notify_db_session):
is_default=True, is_default=True,
inbound_number_id=None, inbound_number_id=None,
) )
stmt = select(ServiceSmsSender).order_by(ServiceSmsSender.created_at)
service_sms_senders = ServiceSmsSender.query.order_by( service_sms_senders = db.session.execute(stmt).scalars().all()
ServiceSmsSender.created_at
).all()
assert len(service_sms_senders) == 2 assert len(service_sms_senders) == 2
assert service_sms_senders[0].sms_sender == "testing" assert service_sms_senders[0].sms_sender == "testing"
assert not service_sms_senders[0].is_default assert not service_sms_senders[0].is_default
@@ -128,7 +126,8 @@ def test_dao_add_sms_sender_for_service_switches_default(notify_db_session):
def test_dao_update_service_sms_sender(notify_db_session): def test_dao_update_service_sms_sender(notify_db_session):
service = create_service() service = create_service()
service_sms_senders = ServiceSmsSender.query.filter_by(service_id=service.id).all() stmt = select(ServiceSmsSender).filter_by(service_id=service.id)
service_sms_senders = db.session.execute(stmt).scalars().all()
assert len(service_sms_senders) == 1 assert len(service_sms_senders) == 1
sms_sender_to_update = service_sms_senders[0] sms_sender_to_update = service_sms_senders[0]
@@ -138,7 +137,8 @@ def test_dao_update_service_sms_sender(notify_db_session):
is_default=True, is_default=True,
sms_sender="updated", sms_sender="updated",
) )
sms_senders = ServiceSmsSender.query.filter_by(service_id=service.id).all() stmt = select(ServiceSmsSender).filter_by(service_id=service.id)
sms_senders = db.session.execute(stmt).scalars().all()
assert len(sms_senders) == 1 assert len(sms_senders) == 1
assert sms_senders[0].is_default assert sms_senders[0].is_default
assert sms_senders[0].sms_sender == "updated" assert sms_senders[0].sms_sender == "updated"
@@ -159,7 +159,8 @@ def test_dao_update_service_sms_sender_switches_default(notify_db_session):
is_default=True, is_default=True,
sms_sender="updated", sms_sender="updated",
) )
sms_senders = ServiceSmsSender.query.filter_by(service_id=service.id).all() stmt = select(ServiceSmsSender).filter_by(service_id=service.id)
sms_senders = db.session.execute(stmt).scalars().all()
expected = {("testing", False), ("updated", True)} expected = {("testing", False), ("updated", True)}
results = {(sender.sms_sender, sender.is_default) for sender in sms_senders} results = {(sender.sms_sender, sender.is_default) for sender in sms_senders}
@@ -190,7 +191,8 @@ def test_update_existing_sms_sender_with_inbound_number(notify_db_session):
service = create_service() service = create_service()
inbound_number = create_inbound_number(number="12345", service_id=service.id) inbound_number = create_inbound_number(number="12345", service_id=service.id)
existing_sms_sender = ServiceSmsSender.query.filter_by(service_id=service.id).one() stmt = select(ServiceSmsSender).filter_by(service_id=service.id)
existing_sms_sender = db.session.execute(stmt).scalars().one()
sms_sender = update_existing_sms_sender_with_inbound_number( sms_sender = update_existing_sms_sender_with_inbound_number(
service_sms_sender=existing_sms_sender, service_sms_sender=existing_sms_sender,
sms_sender=inbound_number.number, sms_sender=inbound_number.number,
@@ -206,7 +208,8 @@ def test_update_existing_sms_sender_with_inbound_number_raises_exception_if_inbo
notify_db_session, notify_db_session,
): ):
service = create_service() service = create_service()
existing_sms_sender = ServiceSmsSender.query.filter_by(service_id=service.id).one() stmt = select(ServiceSmsSender).filter_by(service_id=service.id)
existing_sms_sender = db.session.execute(stmt).scalars().one()
with pytest.raises(expected_exception=SQLAlchemyError): with pytest.raises(expected_exception=SQLAlchemyError):
update_existing_sms_sender_with_inbound_number( update_existing_sms_sender_with_inbound_number(
service_sms_sender=existing_sms_sender, service_sms_sender=existing_sms_sender,
+135 -107
View File
@@ -6,6 +6,7 @@ from unittest.mock import Mock
import pytest import pytest
import sqlalchemy import sqlalchemy
from freezegun import freeze_time from freezegun import freeze_time
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
@@ -89,9 +90,32 @@ from tests.app.db import (
) )
def _get_service_query_count():
stmt = select(func.count(Service.id))
return db.session.execute(stmt).scalar() or 0
def _get_service_history_query_count():
stmt = select(func.count(Service.get_history_model().id))
return db.session.execute(stmt).scalar() or 0
def _get_first_service():
stmt = select(Service).limit(1)
service = db.session.execute(stmt).scalars().first()
return service
def _get_service_by_id(service_id):
stmt = select(Service).filter(Service.id == service_id)
service = db.session.execute(stmt).scalars().one()
return service
def test_create_service(notify_db_session): def test_create_service(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -101,8 +125,8 @@ def test_create_service(notify_db_session):
created_by=user, created_by=user,
) )
dao_create_service(service, user) dao_create_service(service, user)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
service_db = Service.query.one() service_db = _get_first_service()
assert service_db.name == "service_name" assert service_db.name == "service_name"
assert service_db.id == service.id assert service_db.id == service.id
assert service_db.email_from == "email_from" assert service_db.email_from == "email_from"
@@ -120,7 +144,7 @@ def test_create_service_with_organization(notify_db_session):
organization_type=OrganizationType.STATE, organization_type=OrganizationType.STATE,
domains=["local-authority.gov.uk"], domains=["local-authority.gov.uk"],
) )
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -130,9 +154,9 @@ def test_create_service_with_organization(notify_db_session):
created_by=user, created_by=user,
) )
dao_create_service(service, user) dao_create_service(service, user)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
service_db = Service.query.one() service_db = _get_first_service()
organization = Organization.query.get(organization.id) organization = db.session.get(Organization, organization.id)
assert service_db.name == "service_name" assert service_db.name == "service_name"
assert service_db.id == service.id assert service_db.id == service.id
assert service_db.email_from == "email_from" assert service_db.email_from == "email_from"
@@ -151,7 +175,7 @@ def test_fetch_service_by_id_with_api_keys(notify_db_session):
organization_type=OrganizationType.STATE, organization_type=OrganizationType.STATE,
domains=["local-authority.gov.uk"], domains=["local-authority.gov.uk"],
) )
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -161,9 +185,9 @@ def test_fetch_service_by_id_with_api_keys(notify_db_session):
created_by=user, created_by=user,
) )
dao_create_service(service, user) dao_create_service(service, user)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
service_db = Service.query.one() service_db = _get_first_service()
organization = Organization.query.get(organization.id) organization = db.session.get(Organization, organization.id)
assert service_db.name == "service_name" assert service_db.name == "service_name"
assert service_db.id == service.id assert service_db.id == service.id
assert service_db.email_from == "email_from" assert service_db.email_from == "email_from"
@@ -183,7 +207,7 @@ def test_fetch_service_by_id_with_api_keys(notify_db_session):
def test_cannot_create_two_services_with_same_name(notify_db_session): def test_cannot_create_two_services_with_same_name(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service1 = Service( service1 = Service(
name="service_name", name="service_name",
email_from="email_from1", email_from="email_from1",
@@ -209,7 +233,7 @@ def test_cannot_create_two_services_with_same_name(notify_db_session):
def test_cannot_create_two_services_with_same_email_from(notify_db_session): def test_cannot_create_two_services_with_same_email_from(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service1 = Service( service1 = Service(
name="service_name1", name="service_name1",
email_from="email_from", email_from="email_from",
@@ -235,7 +259,7 @@ def test_cannot_create_two_services_with_same_email_from(notify_db_session):
def test_cannot_create_service_with_no_user(notify_db_session): def test_cannot_create_service_with_no_user(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -258,7 +282,7 @@ def test_should_add_user_to_service(notify_db_session):
created_by=user, created_by=user,
) )
dao_create_service(service, user) dao_create_service(service, user)
assert user in Service.query.first().users assert user in _get_first_service().users
new_user = User( new_user = User(
name="Test User", name="Test User",
email_address="new_user@digital.fake.gov", email_address="new_user@digital.fake.gov",
@@ -267,7 +291,7 @@ def test_should_add_user_to_service(notify_db_session):
) )
save_model_user(new_user, validated_email_access=True) save_model_user(new_user, validated_email_access=True)
dao_add_user_to_service(service, new_user) dao_add_user_to_service(service, new_user)
assert new_user in Service.query.first().users assert new_user in _get_first_service().users
def test_dao_add_user_to_service_sets_folder_permissions(sample_user, sample_service): def test_dao_add_user_to_service_sets_folder_permissions(sample_user, sample_service):
@@ -314,7 +338,8 @@ def test_dao_add_user_to_service_raises_error_if_adding_folder_permissions_for_a
other_service_folder = create_template_folder(other_service) other_service_folder = create_template_folder(other_service)
folder_permissions = [str(other_service_folder.id)] folder_permissions = [str(other_service_folder.id)]
assert ServiceUser.query.count() == 2 stmt = select(func.count(ServiceUser.service_id))
assert db.session.execute(stmt).scalar() == 2
with pytest.raises(IntegrityError) as e: with pytest.raises(IntegrityError) as e:
dao_add_user_to_service( dao_add_user_to_service(
@@ -326,7 +351,8 @@ def test_dao_add_user_to_service_raises_error_if_adding_folder_permissions_for_a
'insert or update on table "user_folder_permissions" violates foreign key constraint' 'insert or update on table "user_folder_permissions" violates foreign key constraint'
in str(e.value) in str(e.value)
) )
assert ServiceUser.query.count() == 2 stmt = select(func.count(ServiceUser.service_id))
assert db.session.execute(stmt).scalar() == 2
def test_should_remove_user_from_service(notify_db_session): def test_should_remove_user_from_service(notify_db_session):
@@ -347,9 +373,9 @@ def test_should_remove_user_from_service(notify_db_session):
) )
save_model_user(new_user, validated_email_access=True) save_model_user(new_user, validated_email_access=True)
dao_add_user_to_service(service, new_user) dao_add_user_to_service(service, new_user)
assert new_user in Service.query.first().users assert new_user in _get_first_service().users
dao_remove_user_from_service(service, new_user) dao_remove_user_from_service(service, new_user)
assert new_user not in Service.query.first().users assert new_user not in _get_first_service().users
def test_should_remove_user_from_service_exception(notify_db_session): def test_should_remove_user_from_service_exception(notify_db_session):
@@ -382,11 +408,12 @@ def test_should_remove_user_from_service_exception(notify_db_session):
def test_removing_a_user_from_a_service_deletes_their_permissions( def test_removing_a_user_from_a_service_deletes_their_permissions(
sample_user, sample_service sample_user, sample_service
): ):
assert len(Permission.query.all()) == 7 stmt = select(Permission)
assert len(db.session.execute(stmt).all()) == 7
dao_remove_user_from_service(sample_service, sample_user) dao_remove_user_from_service(sample_service, sample_user)
assert Permission.query.all() == [] assert db.session.execute(stmt).all() == []
def test_removing_a_user_from_a_service_deletes_their_folder_permissions_for_that_service( def test_removing_a_user_from_a_service_deletes_their_folder_permissions_for_that_service(
@@ -668,8 +695,8 @@ def test_removing_all_permission_returns_service_with_no_permissions(notify_db_s
def test_create_service_creates_a_history_record_with_current_data(notify_db_session): def test_create_service_creates_a_history_record_with_current_data(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
assert Service.get_history_model().query.count() == 0 assert _get_service_history_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -678,11 +705,12 @@ def test_create_service_creates_a_history_record_with_current_data(notify_db_ses
created_by=user, created_by=user,
) )
dao_create_service(service, user) dao_create_service(service, user)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
assert Service.get_history_model().query.count() == 1 assert _get_service_history_query_count() == 1
service_from_db = Service.query.first() service_from_db = _get_first_service()
service_history = Service.get_history_model().query.first() stmt = select(Service.get_history_model())
service_history = db.session.execute(stmt).scalars().first()
assert service_from_db.id == service_history.id assert service_from_db.id == service_history.id
assert service_from_db.name == service_history.name assert service_from_db.name == service_history.name
@@ -694,8 +722,8 @@ def test_create_service_creates_a_history_record_with_current_data(notify_db_ses
def test_update_service_creates_a_history_record_with_current_data(notify_db_session): def test_update_service_creates_a_history_record_with_current_data(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
assert Service.get_history_model().query.count() == 0 assert _get_service_history_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -705,39 +733,31 @@ def test_update_service_creates_a_history_record_with_current_data(notify_db_ses
) )
dao_create_service(service, user) dao_create_service(service, user)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
assert Service.query.first().version == 1 assert _get_first_service().version == 1
assert Service.get_history_model().query.count() == 1 assert _get_service_history_query_count() == 1
service.name = "updated_service_name" service.name = "updated_service_name"
dao_update_service(service) dao_update_service(service)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
assert Service.get_history_model().query.count() == 2 assert _get_service_history_query_count() == 2
service_from_db = Service.query.first() service_from_db = _get_first_service()
assert service_from_db.version == 2 assert service_from_db.version == 2
stmt = select(Service.get_history_model()).filter_by(name="service_name")
assert ( assert db.session.execute(stmt).scalars().one().version == 1
Service.get_history_model().query.filter_by(name="service_name").one().version stmt = select(Service.get_history_model()).filter_by(name="updated_service_name")
== 1 assert db.session.execute(stmt).scalars().one().version == 2
)
assert (
Service.get_history_model()
.query.filter_by(name="updated_service_name")
.one()
.version
== 2
)
def test_update_service_permission_creates_a_history_record_with_current_data( def test_update_service_permission_creates_a_history_record_with_current_data(
notify_db_session, notify_db_session,
): ):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
assert Service.get_history_model().query.count() == 0 assert _get_service_history_query_count() == 0
service = Service( service = Service(
name="service_name", name="service_name",
email_from="email_from", email_from="email_from",
@@ -755,17 +775,17 @@ def test_update_service_permission_creates_a_history_record_with_current_data(
], ],
) )
assert Service.query.count() == 1 assert _get_service_query_count() == 1
service.permissions.append( service.permissions.append(
ServicePermission(service_id=service.id, permission=ServicePermissionType.EMAIL) ServicePermission(service_id=service.id, permission=ServicePermissionType.EMAIL)
) )
dao_update_service(service) dao_update_service(service)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
assert Service.get_history_model().query.count() == 2 assert _get_service_history_query_count() == 2
service_from_db = Service.query.first() service_from_db = _get_first_service()
assert service_from_db.version == 2 assert service_from_db.version == 2
@@ -784,10 +804,10 @@ def test_update_service_permission_creates_a_history_record_with_current_data(
service.permissions.remove(permission) service.permissions.remove(permission)
dao_update_service(service) dao_update_service(service)
assert Service.query.count() == 1 assert _get_service_query_count() == 1
assert Service.get_history_model().query.count() == 3 assert _get_service_history_query_count() == 3
service_from_db = Service.query.first() service_from_db = _get_first_service()
assert service_from_db.version == 3 assert service_from_db.version == 3
_assert_service_permissions( _assert_service_permissions(
service.permissions, service.permissions,
@@ -797,21 +817,20 @@ def test_update_service_permission_creates_a_history_record_with_current_data(
), ),
) )
history = ( stmt = (
Service.get_history_model() select(Service.get_history_model())
.query.filter_by(name="service_name") .filter_by(name="service_name")
.order_by("version") .order_by("version")
.all()
) )
history = db.session.execute(stmt).scalars().all()
assert len(history) == 3 assert len(history) == 3
assert history[2].version == 3 assert history[2].version == 3
def test_create_service_and_history_is_transactional(notify_db_session): def test_create_service_and_history_is_transactional(notify_db_session):
user = create_user() user = create_user()
assert Service.query.count() == 0 assert _get_service_query_count() == 0
assert Service.get_history_model().query.count() == 0 assert _get_service_history_query_count() == 0
service = Service( service = Service(
name=None, name=None,
email_from="email_from", email_from="email_from",
@@ -828,8 +847,8 @@ def test_create_service_and_history_is_transactional(notify_db_session):
in str(seeei) in str(seeei)
) )
assert Service.query.count() == 0 assert _get_service_query_count() == 0
assert Service.get_history_model().query.count() == 0 assert _get_service_history_query_count() == 0
def test_delete_service_and_associated_objects(notify_db_session): def test_delete_service_and_associated_objects(notify_db_session):
@@ -845,8 +864,8 @@ def test_delete_service_and_associated_objects(notify_db_session):
create_notification(template=template, api_key=api_key) create_notification(template=template, api_key=api_key)
create_invited_user(service=service) create_invited_user(service=service)
user.organizations = [organization] user.organizations = [organization]
stmt = select(func.count(ServicePermission.service_id))
assert ServicePermission.query.count() == len( assert db.session.execute(stmt).scalar() == len(
( (
ServicePermissionType.SMS, ServicePermissionType.SMS,
ServicePermissionType.EMAIL, ServicePermissionType.EMAIL,
@@ -855,21 +874,35 @@ def test_delete_service_and_associated_objects(notify_db_session):
) )
delete_service_and_all_associated_db_objects(service) delete_service_and_all_associated_db_objects(service)
assert VerifyCode.query.count() == 0 stmt = select(VerifyCode)
assert ApiKey.query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert ApiKey.get_history_model().query.count() == 0 stmt = select(ApiKey)
assert Template.query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert TemplateHistory.query.count() == 0 stmt = select(ApiKey.get_history_model())
assert Job.query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert Notification.query.count() == 0 stmt = select(Template)
assert Permission.query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert User.query.count() == 0 stmt = select(TemplateHistory)
assert InvitedUser.query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert Service.query.count() == 0 stmt = select(Job)
assert Service.get_history_model().query.count() == 0 assert db.session.execute(stmt).scalar() is None
assert ServicePermission.query.count() == 0 stmt = select(Notification)
assert db.session.execute(stmt).scalar() is None
stmt = select(Permission)
assert db.session.execute(stmt).scalar() is None
stmt = select(User)
assert db.session.execute(stmt).scalar() is None
stmt = select(InvitedUser)
assert db.session.execute(stmt).scalar() is None
assert _get_service_query_count() == 0
assert _get_service_history_query_count() == 0
stmt = select(ServicePermission)
assert db.session.execute(stmt).scalar() is None
# the organization hasn't been deleted # the organization hasn't been deleted
assert Organization.query.count() == 1 stmt = select(func.count(Organization.id))
assert db.session.execute(stmt).scalar() == 1
def test_add_existing_user_to_another_service_doesnot_change_old_permissions( def test_add_existing_user_to_another_service_doesnot_change_old_permissions(
@@ -887,9 +920,8 @@ def test_add_existing_user_to_another_service_doesnot_change_old_permissions(
dao_create_service(service_one, user) dao_create_service(service_one, user)
assert user.id == service_one.users[0].id assert user.id == service_one.users[0].id
test_user_permissions = Permission.query.filter_by( stmt = select(Permission).filter_by(service=service_one, user=user)
service=service_one, user=user test_user_permissions = db.session.execute(stmt).all()
).all()
assert len(test_user_permissions) == 7 assert len(test_user_permissions) == 7
other_user = User( other_user = User(
@@ -909,14 +941,12 @@ def test_add_existing_user_to_another_service_doesnot_change_old_permissions(
dao_create_service(service_two, other_user) dao_create_service(service_two, other_user)
assert other_user.id == service_two.users[0].id assert other_user.id == service_two.users[0].id
other_user_permissions = Permission.query.filter_by( stmt = select(Permission).filter_by(service=service_two, user=other_user)
service=service_two, user=other_user other_user_permissions = db.session.execute(stmt).all()
).all()
assert len(other_user_permissions) == 7 assert len(other_user_permissions) == 7
stmt = select(Permission).filter_by(service=service_one, user=other_user)
other_user_service_one_permissions = db.session.execute(stmt).all()
other_user_service_one_permissions = Permission.query.filter_by(
service=service_one, user=other_user
).all()
assert len(other_user_service_one_permissions) == 0 assert len(other_user_service_one_permissions) == 0
# adding the other_user to service_one should leave all other_user permissions on service_two intact # adding the other_user to service_one should leave all other_user permissions on service_two intact
@@ -925,15 +955,12 @@ def test_add_existing_user_to_another_service_doesnot_change_old_permissions(
permissions.append(Permission(permission=p)) permissions.append(Permission(permission=p))
dao_add_user_to_service(service_one, other_user, permissions=permissions) dao_add_user_to_service(service_one, other_user, permissions=permissions)
stmt = select(Permission).filter_by(service=service_one, user=other_user)
other_user_service_one_permissions = Permission.query.filter_by( other_user_service_one_permissions = db.session.execute(stmt).all()
service=service_one, user=other_user
).all()
assert len(other_user_service_one_permissions) == 2 assert len(other_user_service_one_permissions) == 2
other_user_service_two_permissions = Permission.query.filter_by( stmt = select(Permission).filter_by(service=service_two, user=other_user)
service=service_two, user=other_user other_user_service_two_permissions = db.session.execute(stmt).all()
).all()
assert len(other_user_service_two_permissions) == 7 assert len(other_user_service_two_permissions) == 7
@@ -956,9 +983,10 @@ def test_fetch_stats_filters_on_service(notify_db_session):
def test_fetch_stats_ignores_historical_notification_data(sample_template): def test_fetch_stats_ignores_historical_notification_data(sample_template):
create_notification_history(template=sample_template) create_notification_history(template=sample_template)
stmt = select(func.count(Notification.id))
assert Notification.query.count() == 0 assert db.session.execute(stmt).scalar() == 0
assert NotificationHistory.query.count() == 1 stmt = select(func.count(NotificationHistory.id))
assert db.session.execute(stmt).scalar() == 1
stats = dao_fetch_todays_stats_for_service(sample_template.service_id) stats = dao_fetch_todays_stats_for_service(sample_template.service_id)
assert len(stats) == 0 assert len(stats) == 0
@@ -1316,7 +1344,7 @@ def test_dao_fetch_todays_stats_for_all_services_can_exclude_from_test_key(
def test_dao_suspend_service_with_no_api_keys(notify_db_session): def test_dao_suspend_service_with_no_api_keys(notify_db_session):
service = create_service() service = create_service()
dao_suspend_service(service.id) dao_suspend_service(service.id)
service = Service.query.get(service.id) service = _get_service_by_id(service.id)
assert not service.active assert not service.active
assert service.name == service.name assert service.name == service.name
assert service.api_keys == [] assert service.api_keys == []
@@ -1329,11 +1357,11 @@ def test_dao_suspend_service_marks_service_as_inactive_and_expires_api_keys(
service = create_service() service = create_service()
api_key = create_api_key(service=service) api_key = create_api_key(service=service)
dao_suspend_service(service.id) dao_suspend_service(service.id)
service = Service.query.get(service.id) service = _get_service_by_id(service.id)
assert not service.active assert not service.active
assert service.name == service.name assert service.name == service.name
api_key = ApiKey.query.get(api_key.id) api_key = db.session.get(ApiKey, api_key.id)
assert api_key.expiry_date == datetime(2001, 1, 1, 23, 59, 00) assert api_key.expiry_date == datetime(2001, 1, 1, 23, 59, 00)
@@ -1344,13 +1372,13 @@ def test_dao_resume_service_marks_service_as_active_and_api_keys_are_still_revok
service = create_service() service = create_service()
api_key = create_api_key(service=service) api_key = create_api_key(service=service)
dao_suspend_service(service.id) dao_suspend_service(service.id)
service = Service.query.get(service.id) service = _get_service_by_id(service.id)
assert not service.active assert not service.active
dao_resume_service(service.id) dao_resume_service(service.id)
assert Service.query.get(service.id).active assert _get_service_by_id(service.id).active
api_key = ApiKey.query.get(api_key.id) api_key = db.session.get(ApiKey, api_key.id)
assert api_key.expiry_date == datetime(2001, 1, 1, 23, 59, 00) assert api_key.expiry_date == datetime(2001, 1, 1, 23, 59, 00)
+4 -2
View File
@@ -1,3 +1,5 @@
from sqlalchemy import select
from app import db from app import db
from app.dao.service_user_dao import dao_get_service_user from app.dao.service_user_dao import dao_get_service_user
from app.dao.template_folder_dao import ( from app.dao.template_folder_dao import (
@@ -17,5 +19,5 @@ def test_dao_delete_template_folder_deletes_user_folder_permissions(
dao_update_template_folder(folder) dao_update_template_folder(folder)
dao_delete_template_folder(folder) dao_delete_template_folder(folder)
stmt = select(user_folder_permissions)
assert db.session.query(user_folder_permissions).all() == [] assert db.session.execute(stmt).scalars().all() == []
+44 -23
View File
@@ -2,8 +2,10 @@ from datetime import datetime
import pytest import pytest
from freezegun import freeze_time from freezegun import freeze_time
from sqlalchemy import func, select
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
from app import db
from app.dao.templates_dao import ( from app.dao.templates_dao import (
dao_create_template, dao_create_template,
dao_get_all_templates_for_service, dao_get_all_templates_for_service,
@@ -17,6 +19,16 @@ from app.models import Template, TemplateHistory, TemplateRedacted
from tests.app.db import create_template from tests.app.db import create_template
def template_query_count():
stmt = select(func.count()).select_from(Template)
return db.session.execute(stmt).scalar() or 0
def template_history_query_count():
stmt = select(func.count()).select_from(TemplateHistory)
return db.session.execute(stmt).scalar() or 0
@pytest.mark.parametrize( @pytest.mark.parametrize(
"template_type, subject", "template_type, subject",
[ [
@@ -37,7 +49,7 @@ def test_create_template(sample_service, sample_user, template_type, subject):
template = Template(**data) template = Template(**data)
dao_create_template(template) dao_create_template(template)
assert Template.query.count() == 1 assert template_query_count() == 1
assert len(dao_get_all_templates_for_service(sample_service.id)) == 1 assert len(dao_get_all_templates_for_service(sample_service.id)) == 1
assert ( assert (
dao_get_all_templates_for_service(sample_service.id)[0].name dao_get_all_templates_for_service(sample_service.id)[0].name
@@ -50,11 +62,13 @@ def test_create_template(sample_service, sample_user, template_type, subject):
def test_create_template_creates_redact_entry(sample_service): def test_create_template_creates_redact_entry(sample_service):
assert TemplateRedacted.query.count() == 0 stmt = select(func.count()).select_from(TemplateRedacted)
assert db.session.execute(stmt).scalar() == 0
template = create_template(sample_service) template = create_template(sample_service)
redacted = TemplateRedacted.query.one() stmt = select(TemplateRedacted)
redacted = db.session.execute(stmt).scalars().one()
assert redacted.template_id == template.id assert redacted.template_id == template.id
assert redacted.redact_personalisation is False assert redacted.redact_personalisation is False
assert redacted.updated_by_id == sample_service.created_by_id assert redacted.updated_by_id == sample_service.created_by_id
@@ -79,7 +93,8 @@ def test_update_template(sample_service, sample_user):
def test_redact_template(sample_template): def test_redact_template(sample_template):
redacted = TemplateRedacted.query.one() stmt = select(TemplateRedacted)
redacted = db.session.execute(stmt).scalars().one()
assert redacted.template_id == sample_template.id assert redacted.template_id == sample_template.id
assert redacted.redact_personalisation is False assert redacted.redact_personalisation is False
@@ -96,7 +111,7 @@ def test_get_all_templates_for_service(service_factory):
service_1 = service_factory.get("service 1", email_from="service.1") service_1 = service_factory.get("service 1", email_from="service.1")
service_2 = service_factory.get("service 2", email_from="service.2") service_2 = service_factory.get("service 2", email_from="service.2")
assert Template.query.count() == 2 assert template_query_count() == 2
assert len(dao_get_all_templates_for_service(service_1.id)) == 1 assert len(dao_get_all_templates_for_service(service_1.id)) == 1
assert len(dao_get_all_templates_for_service(service_2.id)) == 1 assert len(dao_get_all_templates_for_service(service_2.id)) == 1
@@ -119,7 +134,7 @@ def test_get_all_templates_for_service(service_factory):
content="Template content", content="Template content",
) )
assert Template.query.count() == 5 assert template_query_count() == 5
assert len(dao_get_all_templates_for_service(service_1.id)) == 3 assert len(dao_get_all_templates_for_service(service_1.id)) == 3
assert len(dao_get_all_templates_for_service(service_2.id)) == 2 assert len(dao_get_all_templates_for_service(service_2.id)) == 2
@@ -144,7 +159,7 @@ def test_get_all_templates_for_service_is_alphabetised(sample_service):
service=sample_service, service=sample_service,
) )
assert Template.query.count() == 3 assert template_query_count() == 3
assert ( assert (
dao_get_all_templates_for_service(sample_service.id)[0].name dao_get_all_templates_for_service(sample_service.id)[0].name
== "Sample Template 1" == "Sample Template 1"
@@ -171,7 +186,7 @@ def test_get_all_templates_for_service_is_alphabetised(sample_service):
def test_get_all_returns_empty_list_if_no_templates(sample_service): def test_get_all_returns_empty_list_if_no_templates(sample_service):
assert Template.query.count() == 0 assert template_query_count() == 0
assert len(dao_get_all_templates_for_service(sample_service.id)) == 0 assert len(dao_get_all_templates_for_service(sample_service.id)) == 0
@@ -257,8 +272,8 @@ def test_get_template_by_id_and_service_returns_none_if_no_template(
def test_create_template_creates_a_history_record_with_current_data( def test_create_template_creates_a_history_record_with_current_data(
sample_service, sample_user sample_service, sample_user
): ):
assert Template.query.count() == 0 assert template_query_count() == 0
assert TemplateHistory.query.count() == 0 assert template_history_query_count() == 0
data = { data = {
"name": "Sample Template", "name": "Sample Template",
"template_type": TemplateType.EMAIL, "template_type": TemplateType.EMAIL,
@@ -270,10 +285,12 @@ def test_create_template_creates_a_history_record_with_current_data(
template = Template(**data) template = Template(**data)
dao_create_template(template) dao_create_template(template)
assert Template.query.count() == 1 assert template_query_count() == 1
template_from_db = Template.query.first() stmt = select(Template)
template_history = TemplateHistory.query.first() template_from_db = db.session.execute(stmt).scalars().first()
stmt = select(TemplateHistory)
template_history = db.session.execute(stmt).scalars().first()
assert template_from_db.id == template_history.id assert template_from_db.id == template_history.id
assert template_from_db.name == template_history.name assert template_from_db.name == template_history.name
@@ -286,8 +303,8 @@ def test_create_template_creates_a_history_record_with_current_data(
def test_update_template_creates_a_history_record_with_current_data( def test_update_template_creates_a_history_record_with_current_data(
sample_service, sample_user sample_service, sample_user
): ):
assert Template.query.count() == 0 assert template_query_count() == 0
assert TemplateHistory.query.count() == 0 assert template_history_query_count() == 0
data = { data = {
"name": "Sample Template", "name": "Sample Template",
"template_type": TemplateType.EMAIL, "template_type": TemplateType.EMAIL,
@@ -301,22 +318,26 @@ def test_update_template_creates_a_history_record_with_current_data(
created = dao_get_all_templates_for_service(sample_service.id)[0] created = dao_get_all_templates_for_service(sample_service.id)[0]
assert created.name == "Sample Template" assert created.name == "Sample Template"
assert Template.query.count() == 1 assert template_query_count() == 1
assert Template.query.first().version == 1 stmt = select(Template)
assert TemplateHistory.query.count() == 1 assert db.session.execute(stmt).scalars().first().version == 1
assert template_history_query_count() == 1
created.name = "new name" created.name = "new name"
dao_update_template(created) dao_update_template(created)
assert Template.query.count() == 1 assert template_query_count() == 1
assert TemplateHistory.query.count() == 2 assert template_history_query_count() == 2
template_from_db = Template.query.first() stmt = select(Template)
template_from_db = db.session.execute(stmt).scalars().first()
assert template_from_db.version == 2 assert template_from_db.version == 2
assert TemplateHistory.query.filter_by(name="Sample Template").one().version == 1 stmt = select(TemplateHistory).filter_by(name="Sample Template")
assert TemplateHistory.query.filter_by(name="new name").one().version == 2 assert db.session.execute(stmt).scalars().one().version == 1
stmt = select(TemplateHistory).filter_by(name="new name")
assert db.session.execute(stmt).scalars().one().version == 2
def test_get_template_history_version(sample_user, sample_service, sample_template): def test_get_template_history_version(sample_user, sample_service, sample_template):
+32 -10
View File
@@ -3,6 +3,7 @@ from datetime import timedelta
import pytest import pytest
from freezegun import freeze_time from freezegun import freeze_time
from sqlalchemy import func, select
from sqlalchemy.exc import DataError from sqlalchemy.exc import DataError
from sqlalchemy.orm.exc import NoResultFound from sqlalchemy.orm.exc import NoResultFound
@@ -37,6 +38,21 @@ from tests.app.db import (
) )
def _get_user_query_count():
stmt = select(func.count(User.id))
return db.session.execute(stmt).scalar() or 0
def _get_user_query_first():
stmt = select(User)
return db.session.execute(stmt).scalars().first()
def _get_verify_code_query_count():
stmt = select(func.count(VerifyCode.id))
return db.session.execute(stmt).scalar() or 0
@freeze_time("2020-01-28T12:00:00") @freeze_time("2020-01-28T12:00:00")
@pytest.mark.parametrize( @pytest.mark.parametrize(
"phone_number, expected_phone_number", "phone_number, expected_phone_number",
@@ -55,8 +71,10 @@ def test_create_user(notify_db_session, phone_number, expected_phone_number):
} }
user = User(**data) user = User(**data)
save_model_user(user, password="password", validated_email_access=True) save_model_user(user, password="password", validated_email_access=True)
assert User.query.count() == 1 stmt = select(func.count(User.id))
user_query = User.query.first() assert db.session.execute(stmt).scalar() == 1
stmt = select(User)
user_query = db.session.execute(stmt).scalars().first()
assert user_query.email_address == email assert user_query.email_address == email
assert user_query.id == user.id assert user_query.id == user.id
assert user_query.mobile_number == expected_phone_number assert user_query.mobile_number == expected_phone_number
@@ -68,7 +86,8 @@ def test_get_all_users(notify_db_session):
create_user(email="1@test.com") create_user(email="1@test.com")
create_user(email="2@test.com") create_user(email="2@test.com")
assert User.query.count() == 2 stmt = select(func.count(User.id))
assert db.session.execute(stmt).scalar() == 2
assert len(get_user_by_id()) == 2 assert len(get_user_by_id()) == 2
@@ -89,9 +108,10 @@ def test_get_user_invalid_id(notify_db_session):
def test_delete_users(sample_user): def test_delete_users(sample_user):
assert User.query.count() == 1 stmt = select(func.count(User.id))
assert db.session.execute(stmt).scalar() == 1
delete_model_user(sample_user) delete_model_user(sample_user)
assert User.query.count() == 0 assert db.session.execute(stmt).scalar() == 0
def test_increment_failed_login_should_increment_failed_logins(sample_user): def test_increment_failed_login_should_increment_failed_logins(sample_user):
@@ -127,9 +147,10 @@ def test_get_user_by_email_is_case_insensitive(sample_user):
def test_should_delete_all_verification_codes_more_than_one_day_old(sample_user): def test_should_delete_all_verification_codes_more_than_one_day_old(sample_user):
make_verify_code(sample_user, age=timedelta(hours=24), code="54321") make_verify_code(sample_user, age=timedelta(hours=24), code="54321")
make_verify_code(sample_user, age=timedelta(hours=24), code="54321") make_verify_code(sample_user, age=timedelta(hours=24), code="54321")
assert VerifyCode.query.count() == 2 stmt = select(func.count(VerifyCode.id))
assert db.session.execute(stmt).scalar() == 2
delete_codes_older_created_more_than_a_day_ago() delete_codes_older_created_more_than_a_day_ago()
assert VerifyCode.query.count() == 0 assert db.session.execute(stmt).scalar() == 0
def test_should_not_delete_verification_codes_less_than_one_day_old(sample_user): def test_should_not_delete_verification_codes_less_than_one_day_old(sample_user):
@@ -137,10 +158,11 @@ def test_should_not_delete_verification_codes_less_than_one_day_old(sample_user)
sample_user, age=timedelta(hours=23, minutes=59, seconds=59), code="12345" sample_user, age=timedelta(hours=23, minutes=59, seconds=59), code="12345"
) )
make_verify_code(sample_user, age=timedelta(hours=24), code="54321") make_verify_code(sample_user, age=timedelta(hours=24), code="54321")
stmt = select(func.count(VerifyCode.id))
assert VerifyCode.query.count() == 2 assert db.session.execute(stmt).scalar() == 2
delete_codes_older_created_more_than_a_day_ago() delete_codes_older_created_more_than_a_day_ago()
assert VerifyCode.query.one()._code == "12345" stmt = select(VerifyCode)
assert db.session.execute(stmt).scalars().one()._code == "12345"
def make_verify_code(user, age=None, expiry_age=None, code="12335", code_used=False): def make_verify_code(user, age=None, expiry_age=None, code="12335", code_used=False):