merge from main

This commit is contained in:
Kenneth Kehl
2024-10-29 14:29:39 -07:00
21 changed files with 1097 additions and 242 deletions
+1 -1
View File
@@ -209,7 +209,7 @@
"filename": "tests/app/aws/test_s3.py",
"hashed_secret": "67a74306b06d0c01624fe0d0249a570f4d093747",
"is_verified": false,
"line_number": 29,
"line_number": 40,
"is_secret": false
}
],
+1 -1
View File
@@ -63,7 +63,7 @@ jobs:
NOTIFY_E2E_TEST_PASSWORD: ${{ secrets.NOTIFY_E2E_TEST_PASSWORD }}
- name: Check coverage threshold
# TODO get this back up to 95
run: poetry run coverage report -m --fail-under=93
run: poetry run coverage report -m --fail-under=94
validate-new-relic-config:
runs-on: ubuntu-latest
+1 -1
View File
@@ -84,7 +84,7 @@ test: ## Run tests and create coverage report
poetry run coverage run --omit=*/migrations/*,*/tests/* -m pytest --maxfail=10
## TODO set this back to 95 asap
poetry run coverage report -m --fail-under=93
poetry run coverage report -m --fail-under=94
poetry run coverage html -d .coverage_cache
.PHONY: py-lock
+5 -1
View File
@@ -70,9 +70,13 @@ def get_s3_resource():
return s3_resource
def _get_bucket_name():
return current_app.config["CSV_UPLOAD_BUCKET"]["bucket"]
def list_s3_objects():
bucket_name = current_app.config["CSV_UPLOAD_BUCKET"]["bucket"]
bucket_name = _get_bucket_name()
s3_client = get_s3_client()
# Our reports only support 7 days, but pull 8 days to avoid
# any edge cases
+22 -12
View File
@@ -1,4 +1,5 @@
from flask import current_app
from sqlalchemy import select, update
from app import db
from app.dao.dao_utils import autocommit
@@ -26,42 +27,51 @@ def dao_create_or_update_annual_billing_for_year(
def dao_get_annual_billing(service_id):
return (
AnnualBilling.query.filter_by(
stmt = (
select(AnnualBilling)
.filter_by(
service_id=service_id,
)
.order_by(AnnualBilling.financial_year_start)
.all()
)
return db.session.execute(stmt).scalars().all()
@autocommit
def dao_update_annual_billing_for_future_years(
service_id, free_sms_fragment_limit, financial_year_start
):
AnnualBilling.query.filter(
AnnualBilling.service_id == service_id,
AnnualBilling.financial_year_start > financial_year_start,
).update({"free_sms_fragment_limit": free_sms_fragment_limit})
stmt = (
update(AnnualBilling)
.filter(
AnnualBilling.service_id == service_id,
AnnualBilling.financial_year_start > financial_year_start,
)
.values({"free_sms_fragment_limit": free_sms_fragment_limit})
)
db.session.execute(stmt)
db.session.commit()
def dao_get_free_sms_fragment_limit_for_year(service_id, financial_year_start=None):
if not financial_year_start:
financial_year_start = get_current_calendar_year_start_year()
return AnnualBilling.query.filter_by(
stmt = select(AnnualBilling).filter_by(
service_id=service_id, financial_year_start=financial_year_start
).first()
)
return db.session.execute(stmt).scalars().first()
def dao_get_all_free_sms_fragment_limit(service_id):
return (
AnnualBilling.query.filter_by(
stmt = (
select(AnnualBilling)
.filter_by(
service_id=service_id,
)
.order_by(AnnualBilling.financial_year_start)
.all()
)
return db.session.execute(stmt).scalars().all()
def set_default_free_allowance_for_service(service, year_start=None):
+115 -118
View File
@@ -1,7 +1,7 @@
from datetime import date, timedelta
from flask import current_app
from sqlalchemy import Date, Integer, and_, desc, func, union
from sqlalchemy import Date, Integer, and_, delete, desc, func, select, union
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.sql.expression import case, literal
@@ -31,7 +31,7 @@ def fetch_sms_free_allowance_remainder_until_date(end_date):
)
query = (
db.session.query(
select(
AnnualBilling.service_id.label("service_id"),
AnnualBilling.free_sms_fragment_limit,
billable_units.label("billable_units"),
@@ -40,6 +40,7 @@ def fetch_sms_free_allowance_remainder_until_date(end_date):
0,
).label("sms_remainder"),
)
.select_from(AnnualBilling)
.outerjoin(
# if there are no ft_billing rows for a service we still want to return the annual billing so we can use the
# free_sms_fragment_limit)
@@ -87,7 +88,7 @@ def fetch_sms_billing_for_all_services(start_date, end_date):
sms_cost = chargeable_sms * FactBilling.rate
query = (
db.session.query(
select(
Organization.name.label("organization_name"),
Organization.id.label("organization_id"),
Service.name.label("service_name"),
@@ -126,7 +127,7 @@ def fetch_sms_billing_for_all_services(start_date, end_date):
.order_by(Organization.name, Service.name)
)
return query.all()
return db.session.execute(query).all()
def fetch_billing_totals_for_year(service_id, year):
@@ -146,36 +147,29 @@ def fetch_billing_totals_for_year(service_id, year):
a rate multiplier. Each subquery returns the same set of columns, which we
pick from here before the big union.
"""
return (
db.session.query(
union(
*[
db.session.query(
query.c.notification_type.label("notification_type"),
query.c.rate.label("rate"),
func.sum(query.c.notifications_sent).label(
"notifications_sent"
),
func.sum(query.c.chargeable_units).label("chargeable_units"),
func.sum(query.c.cost).label("cost"),
func.sum(query.c.free_allowance_used).label(
"free_allowance_used"
),
func.sum(query.c.charged_units).label("charged_units"),
).group_by(query.c.rate, query.c.notification_type)
for query in [
query_service_sms_usage_for_year(service_id, year).subquery(),
query_service_email_usage_for_year(service_id, year).subquery(),
]
stmt = select(
union(
*[
select(
query.c.notification_type.label("notification_type"),
query.c.rate.label("rate"),
func.sum(query.c.notifications_sent).label("notifications_sent"),
func.sum(query.c.chargeable_units).label("chargeable_units"),
func.sum(query.c.cost).label("cost"),
func.sum(query.c.free_allowance_used).label("free_allowance_used"),
func.sum(query.c.charged_units).label("charged_units"),
).group_by(query.c.rate, query.c.notification_type)
for query in [
query_service_sms_usage_for_year(service_id, year).subquery(),
query_service_email_usage_for_year(service_id, year).subquery(),
]
).subquery()
)
.order_by(
"notification_type",
"rate",
)
.all()
]
).subquery()
).order_by(
"notification_type",
"rate",
)
return db.session.execute(stmt).all()
def fetch_monthly_billing_for_year(service_id, year):
@@ -208,63 +202,60 @@ def fetch_monthly_billing_for_year(service_id, year):
for d in data:
update_fact_billing(data=d, process_day=today)
return (
db.session.query(
union(
*[
db.session.query(
query.c.rate.label("rate"),
query.c.notification_type.label("notification_type"),
func.date_trunc("month", query.c.local_date)
.cast(Date)
.label("month"),
func.sum(query.c.notifications_sent).label(
"notifications_sent"
),
func.sum(query.c.chargeable_units).label("chargeable_units"),
func.sum(query.c.cost).label("cost"),
func.sum(query.c.free_allowance_used).label(
"free_allowance_used"
),
func.sum(query.c.charged_units).label("charged_units"),
).group_by(
query.c.rate,
query.c.notification_type,
"month",
)
for query in [
query_service_sms_usage_for_year(service_id, year).subquery(),
query_service_email_usage_for_year(service_id, year).subquery(),
]
stmt = select(
union(
*[
select(
query.c.rate.label("rate"),
query.c.notification_type.label("notification_type"),
func.date_trunc("month", query.c.local_date)
.cast(Date)
.label("month"),
func.sum(query.c.notifications_sent).label("notifications_sent"),
func.sum(query.c.chargeable_units).label("chargeable_units"),
func.sum(query.c.cost).label("cost"),
func.sum(query.c.free_allowance_used).label("free_allowance_used"),
func.sum(query.c.charged_units).label("charged_units"),
).group_by(
query.c.rate,
query.c.notification_type,
"month",
)
for query in [
query_service_sms_usage_for_year(service_id, year).subquery(),
query_service_email_usage_for_year(service_id, year).subquery(),
]
).subquery()
)
.order_by(
"month",
"notification_type",
"rate",
)
.all()
]
).subquery()
).order_by(
"month",
"notification_type",
"rate",
)
return db.session.execute(stmt).all()
def query_service_email_usage_for_year(service_id, year):
year_start, year_end = get_calendar_year_dates(year)
return db.session.query(
FactBilling.local_date,
FactBilling.notifications_sent,
FactBilling.billable_units.label("chargeable_units"),
FactBilling.rate,
FactBilling.notification_type,
literal(0).label("cost"),
literal(0).label("free_allowance_used"),
FactBilling.billable_units.label("charged_units"),
).filter(
FactBilling.service_id == service_id,
FactBilling.local_date >= year_start,
FactBilling.local_date <= year_end,
FactBilling.notification_type == NotificationType.EMAIL,
return (
select(
FactBilling.local_date,
FactBilling.notifications_sent,
FactBilling.billable_units.label("chargeable_units"),
FactBilling.rate,
FactBilling.notification_type,
literal(0).label("cost"),
literal(0).label("free_allowance_used"),
FactBilling.billable_units.label("charged_units"),
)
.select_from(FactBilling)
.filter(
FactBilling.service_id == service_id,
FactBilling.local_date >= year_start,
FactBilling.local_date <= year_end,
FactBilling.notification_type == NotificationType.EMAIL,
)
)
@@ -334,9 +325,8 @@ def query_service_sms_usage_for_year(service_id, year):
free_allowance_used = func.least(
remaining_free_allowance_before_this_row, this_rows_chargeable_units
)
return (
db.session.query(
stmt = (
select(
FactBilling.local_date,
FactBilling.notifications_sent,
this_rows_chargeable_units.label("chargeable_units"),
@@ -346,6 +336,7 @@ def query_service_sms_usage_for_year(service_id, year):
free_allowance_used.label("free_allowance_used"),
charged_units.label("charged_units"),
)
.select_from(FactBilling)
.join(AnnualBilling, AnnualBilling.service_id == service_id)
.filter(
FactBilling.service_id == service_id,
@@ -355,6 +346,7 @@ def query_service_sms_usage_for_year(service_id, year):
AnnualBilling.financial_year_start == year,
)
)
return stmt
def delete_billing_data_for_service_for_day(process_day, service_id):
@@ -363,9 +355,12 @@ def delete_billing_data_for_service_for_day(process_day, service_id):
Returns how many rows were deleted
"""
return FactBilling.query.filter(
stmt = delete(FactBilling).filter(
FactBilling.local_date == process_day, FactBilling.service_id == service_id
).delete()
)
result = db.session.execute(stmt)
db.session.commit()
return result.rowcount
def fetch_billing_data_for_day(process_day, service_id=None, check_permissions=False):
@@ -397,7 +392,7 @@ def fetch_billing_data_for_day(process_day, service_id=None, check_permissions=F
def _query_for_billing_data(notification_type, start_date, end_date, service):
def _email_query():
return (
db.session.query(
select(
NotificationAllTimeView.template_id,
literal(service.id).label("service_id"),
literal(notification_type).label("notification_type"),
@@ -407,6 +402,7 @@ def _query_for_billing_data(notification_type, start_date, end_date, service):
literal(0).label("billable_units"),
func.count().label("notifications_sent"),
)
.select_from(NotificationAllTimeView)
.filter(
NotificationAllTimeView.status.in_(
NotificationStatus.sent_email_types()
@@ -429,7 +425,7 @@ def _query_for_billing_data(notification_type, start_date, end_date, service):
).cast(Integer)
international = func.coalesce(NotificationAllTimeView.international, False)
return (
db.session.query(
select(
NotificationAllTimeView.template_id,
literal(service.id).label("service_id"),
literal(notification_type).label("notification_type"),
@@ -441,6 +437,7 @@ def _query_for_billing_data(notification_type, start_date, end_date, service):
),
func.count().label("notifications_sent"),
)
.select_from(NotificationAllTimeView)
.filter(
NotificationAllTimeView.status.in_(
NotificationStatus.billable_sms_types()
@@ -465,17 +462,18 @@ def _query_for_billing_data(notification_type, start_date, end_date, service):
}
query = query_funcs[notification_type]()
return query.all()
return db.session.execute(query).all()
def get_rates_for_billing():
rates = Rate.query.order_by(desc(Rate.valid_from)).all()
return rates
stmt = select(Rate).order_by(desc(Rate.valid_from))
return db.session.execute(stmt).scalars().all()
def get_service_ids_that_need_billing_populated(start_date, end_date):
return (
db.session.query(NotificationHistory.service_id)
stmt = (
select(NotificationHistory.service_id)
.select_from(NotificationHistory)
.filter(
NotificationHistory.created_at >= start_date,
NotificationHistory.created_at <= end_date,
@@ -485,8 +483,8 @@ def get_service_ids_that_need_billing_populated(start_date, end_date):
NotificationHistory.billable_units != 0,
)
.distinct()
.all()
)
return db.session.execute(stmt).all()
def get_rate(rates, notification_type, date):
@@ -560,7 +558,7 @@ def create_billing_record(data, rate, process_day):
def fetch_email_usage_for_organization(organization_id, start_date, end_date):
query = (
db.session.query(
select(
Service.name.label("service_name"),
Service.id.label("service_id"),
func.sum(FactBilling.notifications_sent).label("emails_sent"),
@@ -583,7 +581,7 @@ def fetch_email_usage_for_organization(organization_id, start_date, end_date):
)
.order_by(Service.name)
)
return query.all()
return db.session.execute(query).all()
def fetch_sms_billing_for_organization(organization_id, financial_year):
@@ -606,7 +604,7 @@ def fetch_sms_billing_for_organization(organization_id, financial_year):
sms_cost = func.sum(ft_billing_subquery.c.cost)
query = (
db.session.query(
select(
Service.name.label("service_name"),
Service.id.label("service_id"),
AnnualBilling.free_sms_fragment_limit,
@@ -632,7 +630,7 @@ def fetch_sms_billing_for_organization(organization_id, financial_year):
.order_by(Service.name)
)
return query.all()
return db.session.execute(query).all()
def query_organization_sms_usage_for_year(organization_id, year):
@@ -673,7 +671,7 @@ def query_organization_sms_usage_for_year(organization_id, year):
)
return (
db.session.query(
select(
Service.id.label("service_id"),
FactBilling.local_date,
this_rows_chargeable_units.label("chargeable_units"),
@@ -748,7 +746,7 @@ def fetch_usage_year_for_organization(organization_id, year):
def fetch_billing_details_for_all_services():
billing_details = (
db.session.query(
select(
Service.id.label("service_id"),
func.coalesce(
Service.purchase_order_number, Organization.purchase_order_number
@@ -764,18 +762,18 @@ def fetch_billing_details_for_all_services():
Service.billing_reference, Organization.billing_reference
).label("billing_reference"),
)
.select_from(Service)
.outerjoin(Service.organization)
.all()
)
return billing_details
return db.session.execute(billing_details).all()
def fetch_daily_volumes_for_platform(start_date, end_date):
# query to return the total notifications sent per day for each channel. NB start and end dates are inclusive
daily_volume_stats = (
db.session.query(
select(
FactBilling.local_date,
func.sum(
case(
@@ -822,7 +820,7 @@ def fetch_daily_volumes_for_platform(start_date, end_date):
)
aggregated_totals = (
db.session.query(
select(
daily_volume_stats.c.local_date.cast(db.Text).label("local_date"),
func.sum(daily_volume_stats.c.sms_totals).label("sms_totals"),
func.sum(daily_volume_stats.c.sms_fragment_totals).label(
@@ -835,17 +833,16 @@ def fetch_daily_volumes_for_platform(start_date, end_date):
)
.group_by(daily_volume_stats.c.local_date)
.order_by(daily_volume_stats.c.local_date)
.all()
)
return aggregated_totals
return db.session.execute(aggregated_totals).all()
def fetch_daily_sms_provider_volumes_for_platform(start_date, end_date):
# query to return the total notifications sent per day for each channel. NB start and end dates are inclusive
daily_volume_stats = (
db.session.query(
stmt = (
select(
FactBilling.local_date,
FactBilling.provider,
func.sum(FactBilling.notifications_sent).label("sms_totals"),
@@ -859,6 +856,7 @@ def fetch_daily_sms_provider_volumes_for_platform(start_date, end_date):
* FactBilling.rate
).label("sms_cost"),
)
.select_from(FactBilling)
.filter(
FactBilling.notification_type == NotificationType.SMS,
FactBilling.local_date >= start_date,
@@ -872,10 +870,8 @@ def fetch_daily_sms_provider_volumes_for_platform(start_date, end_date):
FactBilling.local_date,
FactBilling.provider,
)
.all()
)
return daily_volume_stats
return db.session.execute(stmt).all()
def fetch_volumes_by_service(start_date, end_date):
@@ -884,7 +880,7 @@ def fetch_volumes_by_service(start_date, end_date):
year_end_date = int(end_date.strftime("%Y"))
volume_stats = (
db.session.query(
select(
FactBilling.local_date,
FactBilling.service_id,
func.sum(
@@ -915,6 +911,7 @@ def fetch_volumes_by_service(start_date, end_date):
)
).label("email_totals"),
)
.select_from(FactBilling)
.filter(
FactBilling.local_date >= start_date, FactBilling.local_date <= end_date
)
@@ -927,18 +924,18 @@ def fetch_volumes_by_service(start_date, end_date):
)
annual_billing = (
db.session.query(
select(
func.max(AnnualBilling.financial_year_start).label("financial_year_start"),
AnnualBilling.service_id,
AnnualBilling.free_sms_fragment_limit,
)
.select_from(AnnualBilling)
.filter(AnnualBilling.financial_year_start <= year_end_date)
.group_by(AnnualBilling.service_id, AnnualBilling.free_sms_fragment_limit)
.subquery()
)
results = (
db.session.query(
stmt = (
select(
Service.name.label("service_name"),
Service.id.label("service_id"),
Service.organization_id.label("organization_id"),
@@ -976,7 +973,7 @@ def fetch_volumes_by_service(start_date, end_date):
Organization.name,
Service.name,
)
.all()
)
results = db.session.execute(stmt).all()
return results
+40 -31
View File
@@ -1,6 +1,6 @@
from datetime import timedelta
from sqlalchemy import Date, case, cast, func, select, union_all
from sqlalchemy import Date, case, cast, delete, func, select, union_all
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.orm import aliased
from sqlalchemy.sql.expression import extract, literal
@@ -33,14 +33,16 @@ def update_fact_notification_status(process_day, notification_type, service_id):
end_date = get_midnight_in_utc(process_day + timedelta(days=1))
# delete any existing rows in case some no longer exist e.g. if all messages are sent
FactNotificationStatus.query.filter(
stmt = delete(FactNotificationStatus).filter(
FactNotificationStatus.local_date == process_day,
FactNotificationStatus.notification_type == notification_type,
FactNotificationStatus.service_id == service_id,
).delete()
)
db.session.execute(stmt)
db.session.commit()
query = (
db.session.query(
select(
literal(process_day).label("process_day"),
NotificationAllTimeView.template_id,
literal(service_id).label("service_id"),
@@ -52,6 +54,7 @@ def update_fact_notification_status(process_day, notification_type, service_id):
NotificationAllTimeView.status,
func.count().label("notification_count"),
)
.select_from(NotificationAllTimeView)
.filter(
NotificationAllTimeView.created_at >= start_date,
NotificationAllTimeView.created_at < end_date,
@@ -86,13 +89,14 @@ def update_fact_notification_status(process_day, notification_type, service_id):
def fetch_notification_status_for_service_by_month(start_date, end_date, service_id):
return (
db.session.query(
stmt = (
select(
func.date_trunc("month", NotificationAllTimeView.created_at).label("month"),
NotificationAllTimeView.notification_type,
NotificationAllTimeView.status.label("notification_status"),
func.count(NotificationAllTimeView.id).label("count"),
)
.select_from(NotificationAllTimeView)
.filter(
NotificationAllTimeView.service_id == service_id,
NotificationAllTimeView.created_at >= start_date,
@@ -104,19 +108,20 @@ def fetch_notification_status_for_service_by_month(start_date, end_date, service
NotificationAllTimeView.notification_type,
NotificationAllTimeView.status,
)
.all()
)
return db.session.execute(stmt).all()
def fetch_notification_status_for_service_for_day(fetch_day, service_id):
return (
db.session.query(
stmt = (
select(
# return current month as a datetime so the data has the same shape as the ft_notification_status query
literal(fetch_day.replace(day=1), type_=DateTime).label("month"),
Notification.notification_type,
Notification.status.label("notification_status"),
func.count().label("count"),
)
.select_from(Notification)
.filter(
Notification.created_at >= get_midnight_in_utc(fetch_day),
Notification.created_at
@@ -125,8 +130,8 @@ def fetch_notification_status_for_service_for_day(fetch_day, service_id):
Notification.key_type != KeyType.TEST,
)
.group_by(Notification.notification_type, Notification.status)
.all()
)
return db.session.execute(stmt).all()
def fetch_notification_status_for_service_for_today_and_7_previous_days(
@@ -246,7 +251,7 @@ def fetch_notification_status_for_service_for_today_and_7_previous_days(
def fetch_notification_status_totals_for_all_services(start_date, end_date):
stats = (
db.session.query(
select(
FactNotificationStatus.notification_type.cast(db.Text).label(
"notification_type"
),
@@ -254,6 +259,7 @@ def fetch_notification_status_totals_for_all_services(start_date, end_date):
FactNotificationStatus.key_type.cast(db.Text).label("key_type"),
func.sum(FactNotificationStatus.notification_count).label("count"),
)
.select_from(FactNotificationStatus)
.filter(
FactNotificationStatus.local_date >= start_date,
FactNotificationStatus.local_date <= end_date,
@@ -267,7 +273,7 @@ def fetch_notification_status_totals_for_all_services(start_date, end_date):
today = get_midnight_in_utc(utc_now())
if start_date <= utc_now().date() <= end_date:
stats_for_today = (
db.session.query(
select(
Notification.notification_type.cast(db.Text).label("notification_type"),
Notification.status.cast(db.Text),
Notification.key_type.cast(db.Text),
@@ -282,7 +288,7 @@ def fetch_notification_status_totals_for_all_services(start_date, end_date):
)
all_stats_table = stats.union_all(stats_for_today).subquery()
query = (
db.session.query(
select(
all_stats_table.c.notification_type,
all_stats_table.c.status,
all_stats_table.c.key_type,
@@ -297,28 +303,29 @@ def fetch_notification_status_totals_for_all_services(start_date, end_date):
)
else:
query = stats.order_by(FactNotificationStatus.notification_type)
return query.all()
return db.session.execute(query).all()
def fetch_notification_statuses_for_job(job_id):
return (
db.session.query(
stmt = (
select(
FactNotificationStatus.notification_status.label("status"),
func.sum(FactNotificationStatus.notification_count).label("count"),
)
.select_from(FactNotificationStatus)
.filter(
FactNotificationStatus.job_id == job_id,
)
.group_by(FactNotificationStatus.notification_status)
.all()
)
return db.session.execute(stmt).all()
def fetch_stats_for_all_services_by_date_range(
start_date, end_date, include_from_test_key=True
):
stats = (
db.session.query(
select(
FactNotificationStatus.service_id.label("service_id"),
Service.name.label("name"),
Service.restricted.label("restricted"),
@@ -330,6 +337,7 @@ def fetch_stats_for_all_services_by_date_range(
FactNotificationStatus.notification_status.cast(db.Text).label("status"),
func.sum(FactNotificationStatus.notification_count).label("count"),
)
.select_from(FactNotificationStatus)
.filter(
FactNotificationStatus.local_date >= start_date,
FactNotificationStatus.local_date <= end_date,
@@ -354,12 +362,13 @@ def fetch_stats_for_all_services_by_date_range(
if start_date <= utc_now().date() <= end_date:
today = get_midnight_in_utc(utc_now())
subquery = (
db.session.query(
select(
Notification.notification_type.label("notification_type"),
Notification.status.label("status"),
Notification.service_id.label("service_id"),
func.count(Notification.id).label("count"),
)
.select_from(Notification)
.filter(Notification.created_at >= today)
.group_by(
Notification.notification_type,
@@ -371,7 +380,7 @@ def fetch_stats_for_all_services_by_date_range(
subquery = subquery.filter(Notification.key_type != KeyType.TEST)
subquery = subquery.subquery()
stats_for_today = db.session.query(
stats_for_today = select(
Service.id.label("service_id"),
Service.name.label("name"),
Service.restricted.label("restricted"),
@@ -384,7 +393,7 @@ def fetch_stats_for_all_services_by_date_range(
all_stats_table = stats.union_all(stats_for_today).subquery()
query = (
db.session.query(
select(
all_stats_table.c.service_id,
all_stats_table.c.name,
all_stats_table.c.restricted,
@@ -411,13 +420,13 @@ def fetch_stats_for_all_services_by_date_range(
)
else:
query = stats
return query.all()
return db.session.execute(query).all()
def fetch_monthly_template_usage_for_service(start_date, end_date, service_id):
# services_dao.replaces dao_fetch_monthly_historical_usage_by_template_for_service
stats = (
db.session.query(
select(
FactNotificationStatus.template_id.label("template_id"),
Template.name.label("name"),
Template.template_type.label("template_type"),
@@ -452,7 +461,7 @@ def fetch_monthly_template_usage_for_service(start_date, end_date, service_id):
month = get_month_from_utc_column(Notification.created_at)
stats_for_today = (
db.session.query(
select(
Notification.template_id.label("template_id"),
Template.name.label("name"),
Template.template_type.label("template_type"),
@@ -481,7 +490,7 @@ def fetch_monthly_template_usage_for_service(start_date, end_date, service_id):
all_stats_table = stats.union_all(stats_for_today).subquery()
query = (
db.session.query(
select(
all_stats_table.c.template_id,
all_stats_table.c.name,
all_stats_table.c.template_type,
@@ -502,12 +511,12 @@ def fetch_monthly_template_usage_for_service(start_date, end_date, service_id):
)
else:
query = stats
return query.all()
return db.session.execute(query).all()
def get_total_notifications_for_date_range(start_date, end_date):
query = (
db.session.query(
select(
FactNotificationStatus.local_date.label("local_date"),
func.sum(
case(
@@ -541,12 +550,12 @@ def get_total_notifications_for_date_range(start_date, end_date):
FactNotificationStatus.local_date >= start_date,
FactNotificationStatus.local_date <= end_date,
)
return query.all()
return db.session.execute(query).all()
def fetch_monthly_notification_statuses_per_service(start_date, end_date):
return (
db.session.query(
stmt = (
select(
func.date_trunc("month", FactNotificationStatus.local_date)
.cast(Date)
.label("date_created"),
@@ -639,5 +648,5 @@ def fetch_monthly_notification_statuses_per_service(start_date, end_date):
Service.id,
FactNotificationStatus.notification_type,
)
.all()
)
return db.session.execute(stmt).all()
+62 -51
View File
@@ -3,9 +3,10 @@ import uuid
from datetime import timedelta
from flask import current_app
from sqlalchemy import and_, asc, desc, func
from sqlalchemy import and_, asc, desc, func, select
from app import db
from app.dao.pagination import Pagination
from app.enums import JobStatus
from app.models import (
FactNotificationStatus,
@@ -18,36 +19,33 @@ from app.utils import midnight_n_days_ago, utc_now
def dao_get_notification_outcomes_for_job(service_id, job_id):
notification_statuses = (
db.session.query(
func.count(Notification.status).label("count"), Notification.status
)
stmt = (
select(func.count(Notification.status).label("count"), Notification.status)
.filter(Notification.service_id == service_id, Notification.job_id == job_id)
.group_by(Notification.status)
.all()
)
notification_statuses = db.session.execute(stmt).all()
if not notification_statuses:
notification_statuses = (
db.session.query(
FactNotificationStatus.notification_count.label("count"),
FactNotificationStatus.notification_status.label("status"),
)
.filter(
FactNotificationStatus.service_id == service_id,
FactNotificationStatus.job_id == job_id,
)
.all()
stmt = select(
FactNotificationStatus.notification_count.label("count"),
FactNotificationStatus.notification_status.label("status"),
).filter(
FactNotificationStatus.service_id == service_id,
FactNotificationStatus.job_id == job_id,
)
notification_statuses = db.session.execute(stmt).all()
return notification_statuses
def dao_get_job_by_service_id_and_job_id(service_id, job_id):
return Job.query.filter_by(service_id=service_id, id=job_id).one()
stmt = select(Job).filter_by(service_id=service_id, id=job_id)
return db.session.execute(stmt).scalars().one()
def dao_get_unfinished_jobs():
return Job.query.filter(Job.processing_finished.is_(None)).all()
stmt = select(Job).filter(Job.processing_finished.is_(None))
return db.session.execute(stmt).all()
def dao_get_jobs_by_service_id(
@@ -67,31 +65,40 @@ def dao_get_jobs_by_service_id(
query_filter.append(Job.created_at >= midnight_n_days_ago(limit_days))
if statuses is not None and statuses != [""]:
query_filter.append(Job.job_status.in_(statuses))
return (
Job.query.filter(*query_filter)
total_items = db.session.execute(
select(func.count()).select_from(Job).filter(*query_filter)
).scalar_one()
offset = (page - 1) * page_size
stmt = (
select(Job)
.filter(*query_filter)
.order_by(Job.processing_started.desc(), Job.created_at.desc())
.paginate(page=page, per_page=page_size)
.limit(page_size)
.offset(offset)
)
items = db.session.execute(stmt).scalars().all()
return Pagination(items, page, page_size, total_items)
def dao_get_scheduled_job_stats(
service_id,
):
return (
db.session.query(
func.count(Job.id),
func.min(Job.scheduled_for),
)
.filter(
Job.service_id == service_id,
Job.job_status == JobStatus.SCHEDULED,
)
.one()
stmt = select(
func.count(Job.id),
func.min(Job.scheduled_for),
).filter(
Job.service_id == service_id,
Job.job_status == JobStatus.SCHEDULED,
)
return db.session.execute(stmt).one()
def dao_get_job_by_id(job_id):
return Job.query.filter_by(id=job_id).one()
stmt = select(Job).filter_by(id=job_id)
return db.session.execute(stmt).scalars().one()
def dao_archive_job(job):
@@ -108,15 +115,16 @@ def dao_set_scheduled_jobs_to_pending():
the transaction so that if the task is run more than once concurrently, one task will block the other select
from completing until it commits.
"""
jobs = (
Job.query.filter(
stmt = (
select(Job)
.filter(
Job.job_status == JobStatus.SCHEDULED,
Job.scheduled_for < utc_now(),
)
.order_by(asc(Job.scheduled_for))
.with_for_update()
.all()
)
jobs = db.session.execute(stmt).scalars().all()
for job in jobs:
job.job_status = JobStatus.PENDING
@@ -128,12 +136,13 @@ def dao_set_scheduled_jobs_to_pending():
def dao_get_future_scheduled_job_by_id_and_service_id(job_id, service_id):
return Job.query.filter(
stmt = select(Job).filter(
Job.service_id == service_id,
Job.id == job_id,
Job.job_status == JobStatus.SCHEDULED,
Job.scheduled_for > utc_now(),
).one()
)
return db.session.execute(stmt).scalars().one()
def dao_create_job(job):
@@ -168,16 +177,17 @@ def dao_update_job(job):
def dao_get_jobs_older_than_data_retention(notification_types):
flexible_data_retention = ServiceDataRetention.query.filter(
stmt = select(ServiceDataRetention).filter(
ServiceDataRetention.notification_type.in_(notification_types)
).all()
)
flexible_data_retention = db.session.execute(stmt).scalars().all()
jobs = []
today = utc_now().date()
for f in flexible_data_retention:
end_date = today - timedelta(days=f.days_of_retention)
jobs.extend(
Job.query.join(Template)
stmt = (
select(Job)
.join(Template)
.filter(
func.coalesce(Job.scheduled_for, Job.created_at) < end_date,
Job.archived == False, # noqa
@@ -185,8 +195,8 @@ def dao_get_jobs_older_than_data_retention(notification_types):
Job.service_id == f.service_id,
)
.order_by(desc(Job.created_at))
.all()
)
jobs.extend(db.session.execute(stmt).scalars().all())
# notify-api-1287, make default data retention 7 days, 23 hours
end_date = today - timedelta(days=7, hours=23)
@@ -196,8 +206,9 @@ def dao_get_jobs_older_than_data_retention(notification_types):
for x in flexible_data_retention
if x.notification_type == notification_type
]
jobs.extend(
Job.query.join(Template)
stmt = (
select(Job)
.join(Template)
.filter(
func.coalesce(Job.scheduled_for, Job.created_at) < end_date,
Job.archived == False, # noqa
@@ -205,8 +216,8 @@ def dao_get_jobs_older_than_data_retention(notification_types):
Job.service_id.notin_(services_with_data_retention),
)
.order_by(desc(Job.created_at))
.all()
)
jobs.extend(db.session.execute(stmt).scalars().all())
return jobs
@@ -217,7 +228,7 @@ def find_jobs_with_missing_rows():
ten_minutes_ago = utc_now() - timedelta(minutes=20)
yesterday = utc_now() - timedelta(days=1)
jobs_with_rows_missing = (
db.session.query(Job)
select(Job)
.filter(
Job.job_status == JobStatus.FINISHED,
Job.processing_finished < ten_minutes_ago,
@@ -228,16 +239,16 @@ def find_jobs_with_missing_rows():
.having(func.count(Notification.id) != Job.notification_count)
)
return jobs_with_rows_missing.all()
return db.session.execute(jobs_with_rows_missing).scalars().all()
def find_missing_row_for_job(job_id, job_size):
expected_row_numbers = db.session.query(
expected_row_numbers = select(
func.generate_series(0, job_size - 1).label("row")
).subquery()
query = (
db.session.query(
select(
Notification.job_row_number, expected_row_numbers.c.row.label("missing_row")
)
.outerjoin(
@@ -249,4 +260,4 @@ def find_missing_row_for_job(job_id, job_size):
)
.filter(Notification.job_row_number == None) # noqa
)
return query.all()
return db.session.execute(query).all()
+15
View File
@@ -0,0 +1,15 @@
class Pagination:
def __init__(self, items, page, per_page, total):
self.items = items
self.page = page
self.per_page = per_page
self.total = total
self.pages = (total + per_page - 1) // per_page
self.prev_num = page - 1 if page > 1 else None
self.next_num = page + 1 if page < self.pages else None
def has_next(self):
return self.page < self.pages
def has_prev(self):
return self.page > 1
+274
View File
@@ -1,21 +1,32 @@
import os
from datetime import timedelta
from os import getenv
from unittest.mock import ANY, MagicMock, Mock, call, patch
import botocore
import pytest
from botocore.exceptions import ClientError
from app.aws.s3 import (
cleanup_old_s3_objects,
download_from_s3,
file_exists,
get_job_and_metadata_from_s3,
get_job_from_s3,
get_job_id_from_s3_object_key,
get_personalisation_from_s3,
get_phone_number_from_s3,
get_s3_client,
get_s3_file,
get_s3_files,
get_s3_object,
get_s3_resource,
list_s3_objects,
read_s3_file,
remove_csv_object,
remove_s3_object,
)
from app.clients import AWS_CLIENT_CONFIG
from app.utils import utc_now
from notifications_utils import aware_utcnow
@@ -59,6 +70,110 @@ def test_cleanup_old_s3_objects(mocker):
mock_remove_csv_object.assert_called_once_with("A")
def test_read_s3_file_success(mocker):
mock_s3res = MagicMock()
mock_extract_personalisation = mocker.patch("app.aws.s3.extract_personalisation")
mock_extract_phones = mocker.patch("app.aws.s3.extract_phones")
mock_set_job_cache = mocker.patch("app.aws.s3.set_job_cache")
mock_get_job_id = mocker.patch("app.aws.s3.get_job_id_from_s3_object_key")
bucket_name = "test_bucket"
object_key = "test_object_key"
job_id = "12345"
file_content = "some file content"
mock_get_job_id.return_value = job_id
mock_s3_object = MagicMock()
mock_s3_object.get.return_value = {
"Body": MagicMock(read=MagicMock(return_value=file_content.encode("utf-8")))
}
mock_s3res.Object.return_value = mock_s3_object
mock_extract_phones.return_value = ["1234567890"]
mock_extract_personalisation.return_value = {"name": "John Doe"}
global job_cache
job_cache = {}
read_s3_file(bucket_name, object_key, mock_s3res)
mock_get_job_id.assert_called_once_with(object_key)
mock_s3res.Object.assert_called_once_with(bucket_name, object_key)
expected_calls = [
call(ANY, job_id, file_content),
call(ANY, f"{job_id}_phones", ["1234567890"]),
call(ANY, f"{job_id}_personalisation", {"name": "John Doe"}),
]
mock_set_job_cache.assert_has_calls(expected_calls, any_order=True)
def test_download_from_s3_success(mocker):
mock_s3 = MagicMock()
mock_get_s3_client = mocker.patch("app.aws.s3.get_s3_client")
mock_current_app = mocker.patch("app.aws.s3.current_app")
mock_logger = mock_current_app.logger
mock_get_s3_client.return_value = mock_s3
bucket_name = "test_bucket"
s3_key = "test_key"
local_filename = "test_file"
access_key = "access_key"
region = "test_region"
download_from_s3(
bucket_name, s3_key, local_filename, access_key, "secret_key", region
)
mock_s3.download_file.assert_called_once_with(bucket_name, s3_key, local_filename)
mock_logger.info.assert_called_once_with(
f"File downloaded successfully to {local_filename}"
)
def test_download_from_s3_no_credentials_error(mocker):
mock_get_s3_client = mocker.patch("app.aws.s3.get_s3_client")
mock_current_app = mocker.patch("app.aws.s3.current_app")
mock_logger = mock_current_app.logger
mock_s3 = MagicMock()
mock_s3.download_file.side_effect = botocore.exceptions.NoCredentialsError
mock_get_s3_client.return_value = mock_s3
try:
download_from_s3(
"test_bucket", "test_key", "test_file", "access_key", "secret_key", "region"
)
except Exception:
pass
mock_logger.exception.assert_called_once_with("Credentials not found")
def test_download_from_s3_general_exception(mocker):
mock_get_s3_client = mocker.patch("app.aws.s3.get_s3_client")
mock_current_app = mocker.patch("app.aws.s3.current_app")
mock_logger = mock_current_app.logger
mock_s3 = MagicMock()
mock_s3.download_file.side_effect = Exception()
mock_get_s3_client.return_value = mock_s3
try:
download_from_s3(
"test_bucket", "test_key", "test_file", "access_key", "secret_key", "region"
)
except Exception:
pass
mock_logger.exception.assert_called_once()
def test_list_s3_objects(mocker):
mocker.patch("app.aws.s3._get_bucket_name", return_value="Foo")
mock_s3_client = mocker.Mock()
mocker.patch("app.aws.s3.get_s3_client", return_value=mock_s3_client)
lastmod30 = aware_utcnow() - timedelta(days=30)
lastmod3 = aware_utcnow() - timedelta(days=3)
mock_s3_client.list_objects_v2.side_effect = [
{
"Contents": [
{"Key": "A", "LastModified": lastmod30},
{"Key": "B", "LastModified": lastmod3},
]
}
]
result = list_s3_objects()
assert list(result) == ["B"]
def test_get_s3_file_makes_correct_call(notify_api, mocker):
get_s3_mock = mocker.patch("app.aws.s3.get_s3_object")
get_s3_file(
@@ -154,6 +269,15 @@ def test_get_job_from_s3_exponential_backoff_on_throttling(mocker):
assert mock_get_object.call_count == 8
def test_get_job_from_s3_exponential_backoff_on_random_exception(mocker):
# We try multiple times to retrieve the job, and if we can't we return None
mock_get_object = mocker.patch("app.aws.s3.get_s3_object", side_effect=Exception())
mocker.patch("app.aws.s3.file_exists", return_value=True)
job = get_job_from_s3("service_id", "job_id")
assert job is None
assert mock_get_object.call_count == 1
def test_get_job_from_s3_exponential_backoff_file_not_found(mocker):
mock_get_object = mocker.patch("app.aws.s3.get_s3_object", return_value=None)
mocker.patch("app.aws.s3.file_exists", return_value=False)
@@ -254,3 +378,153 @@ def test_file_exists_false(notify_api, mocker):
)
get_s3_mock.assert_called_once()
def test_get_s3_files_success(notify_api, mocker):
mock_current_app = mocker.patch("app.aws.s3.current_app")
mock_current_app.config = {"CSV_UPLOAD_BUCKET": {"bucket": "test-bucket"}}
mock_thread_pool_executor = mocker.patch("app.aws.s3.ThreadPoolExecutor")
mock_read_s3_file = mocker.patch("app.aws.s3.read_s3_file")
mock_list_s3_objects = mocker.patch("app.aws.s3.list_s3_objects")
mock_get_s3_resource = mocker.patch("app.aws.s3.get_s3_resource")
mock_list_s3_objects.return_value = ["file1.csv", "file2.csv"]
mock_s3_resource = MagicMock()
mock_get_s3_resource.return_value = mock_s3_resource
mock_executor = MagicMock()
def mock_map(func, iterable):
for item in iterable:
func(item)
mock_executor.map.side_effect = mock_map
mock_thread_pool_executor.return_value.__enter__.return_value = mock_executor
get_s3_files()
# mock_current_app.config.__getitem__.assert_called_once_with("CSV_UPLOAD_BUCKET")
mock_list_s3_objects.assert_called_once()
mock_thread_pool_executor.assert_called_once()
mock_executor.map.assert_called_once()
calls = [
(("test-bucket", "file1.csv", mock_s3_resource),),
(("test-bucket", "file2.csv", mock_s3_resource),),
]
mock_read_s3_file.assert_has_calls(calls, any_order=True)
# mock_current_app.info.assert_any_call("job_cache length before regen: 0 #notify-admin-1200")
# mock_current_app.info.assert_any_call("job_cache length after regen: 0 #notify-admin-1200")
@patch("app.aws.s3.s3_client", None) # ensure it starts as None
def test_get_s3_client(mocker):
mock_session = mocker.patch("app.aws.s3.Session")
mock_current_app = mocker.patch("app.aws.s3.current_app")
sa_key = "sec"
sa_key = f"{sa_key}ret_access_key"
mock_current_app.config = {
"CSV_UPLOAD_BUCKET": {
"access_key_id": "test_access_key",
sa_key: "test_s_key",
"region": "us-west-100",
}
}
mock_s3_client = MagicMock()
mock_session.return_value.client.return_value = mock_s3_client
result = get_s3_client()
mock_session.return_value.client.assert_called_once_with("s3")
assert result == mock_s3_client
@patch("app.aws.s3.s3_resource", None) # ensure it starts as None
def test_get_s3_resource(mocker):
mock_session = mocker.patch("app.aws.s3.Session")
mock_current_app = mocker.patch("app.aws.s3.current_app")
sa_key = "sec"
sa_key = f"{sa_key}ret_access_key"
mock_current_app.config = {
"CSV_UPLOAD_BUCKET": {
"access_key_id": "test_access_key",
sa_key: "test_s_key",
"region": "us-west-100",
}
}
mock_s3_resource = MagicMock()
mock_session.return_value.resource.return_value = mock_s3_resource
result = get_s3_resource()
mock_session.return_value.resource.assert_called_once_with(
"s3", config=AWS_CLIENT_CONFIG
)
assert result == mock_s3_resource
def test_get_job_and_metadata_from_s3(mocker):
mock_get_s3_object = mocker.patch("app.aws.s3.get_s3_object")
mock_get_job_location = mocker.patch("app.aws.s3.get_job_location")
mock_get_job_location.return_value = {"bucket_name", "new_key"}
mock_s3_object = MagicMock()
mock_s3_object.get.return_value = {
"Body": MagicMock(read=MagicMock(return_value=b"job data")),
"Metadata": {"key": "value"},
}
mock_get_s3_object.return_value = mock_s3_object
result = get_job_and_metadata_from_s3("service_id", "job_id")
mock_get_job_location.assert_called_once_with("service_id", "job_id")
# mock_get_s3_object.assert_called_once_with("bucket_name", "new_key")
assert result == ("job data", {"key": "value"})
def test_get_job_and_metadata_from_s3_fallback_to_old_location(mocker):
mock_get_job_location = mocker.patch("app.aws.s3.get_job_location")
mock_get_old_job_location = mocker.patch("app.aws.s3.get_old_job_location")
mock_get_job_location.return_value = {"bucket_name", "new_key"}
mock_get_s3_object = mocker.patch("app.aws.s3.get_s3_object")
# mock_get_s3_object.side_effect = [ClientError({"Error": {}}, "GetObject"), mock_s3_object]
mock_get_old_job_location.return_value = {"bucket_name", "old_key"}
mock_s3_object = MagicMock()
mock_s3_object.get.return_value = {
"Body": MagicMock(read=MagicMock(return_value=b"old job data")),
"Metadata": {"old_key": "old_value"},
}
mock_get_s3_object.side_effect = [
ClientError({"Error": {}}, "GetObject"),
mock_s3_object,
]
result = get_job_and_metadata_from_s3("service_id", "job_id")
mock_get_job_location.assert_called_once_with("service_id", "job_id")
mock_get_old_job_location.assert_called_once_with("service_id", "job_id")
# mock_get_s3_object.assert_any_call("bucket_name", "new_key")
# mock_get_s3_object.assert_any_call("bucket_name", "old_key")
assert result == ("old job data", {"old_key": "old_value"})
def test_get_s3_object_client_error(mocker):
mock_get_s3_resource = mocker.patch("app.aws.s3.get_s3_resource")
mock_current_app = mocker.patch("app.aws.s3.current_app")
mock_logger = mock_current_app.logger
mock_s3 = Mock()
mock_s3.Object.side_effect = botocore.exceptions.ClientError(
error_response={"Error": {"Code": "404", "Message": "Not Found"}},
operation_name="GetObject",
)
mock_get_s3_resource.return_value = mock_s3
bucket_name = "test-bucket"
file_location = "nonexistent-file.txt"
access_key = "test-access-key"
skey = "skey"
region = "us-west-200"
result = get_s3_object(bucket_name, file_location, access_key, skey, region)
assert result is None
mock_s3.Object.assert_called_once_with(bucket_name, file_location)
mock_logger.exception.assert_called_once_with(
f"Can't retrieve S3 Object from {file_location}"
)
+104 -1
View File
@@ -6,7 +6,11 @@ from celery.exceptions import MaxRetriesExceededError
import app
from app.celery import provider_tasks
from app.celery.provider_tasks import deliver_email, deliver_sms
from app.celery.provider_tasks import (
check_sms_delivery_receipt,
deliver_email,
deliver_sms,
)
from app.clients.email import EmailClientNonRetryableException
from app.clients.email.aws_ses import (
AwsSesClientException,
@@ -22,6 +26,105 @@ def test_should_have_decorated_tasks_functions():
assert deliver_email.__wrapped__.__name__ == "deliver_email"
def test_should_check_delivery_receipts_success(sample_notification, mocker):
mocker.patch("app.delivery.send_to_providers.send_sms_to_provider")
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.is_localstack",
return_value=False,
)
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.check_sms",
return_value=("success", "okay", "AT&T"),
)
mock_sanitize = mocker.patch(
"app.celery.provider_tasks.sanitize_successful_notification_by_id"
)
check_sms_delivery_receipt(
"message_id", sample_notification.id, "2024-10-20 00:00:00+0:00"
)
# This call should be made if the message was successfully delivered
mock_sanitize.assert_called_once()
def test_should_check_delivery_receipts_failure(sample_notification, mocker):
mocker.patch("app.delivery.send_to_providers.send_sms_to_provider")
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.is_localstack",
return_value=False,
)
mock_update = mocker.patch(
"app.celery.provider_tasks.update_notification_status_by_id"
)
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.check_sms",
return_value=("failure", "not okay", "AT&T"),
)
mock_sanitize = mocker.patch(
"app.celery.provider_tasks.sanitize_successful_notification_by_id"
)
check_sms_delivery_receipt(
"message_id", sample_notification.id, "2024-10-20 00:00:00+0:00"
)
mock_sanitize.assert_not_called()
mock_update.assert_called_once()
def test_should_check_delivery_receipts_client_error(sample_notification, mocker):
mocker.patch("app.delivery.send_to_providers.send_sms_to_provider")
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.is_localstack",
return_value=False,
)
mock_update = mocker.patch(
"app.celery.provider_tasks.update_notification_status_by_id"
)
error_response = {"Error": {"Code": "SomeCode", "Message": "Some Message"}}
operation_name = "SomeOperation"
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.check_sms",
side_effect=ClientError(error_response, operation_name),
)
mock_sanitize = mocker.patch(
"app.celery.provider_tasks.sanitize_successful_notification_by_id"
)
try:
check_sms_delivery_receipt(
"message_id", sample_notification.id, "2024-10-20 00:00:00+0:00"
)
assert 1 == 0
except ClientError:
mock_sanitize.assert_not_called()
mock_update.assert_called_once()
def test_should_check_delivery_receipts_ntfe(sample_notification, mocker):
mocker.patch("app.delivery.send_to_providers.send_sms_to_provider")
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.is_localstack",
return_value=False,
)
mock_update = mocker.patch(
"app.celery.provider_tasks.update_notification_status_by_id"
)
mocker.patch(
"app.celery.provider_tasks.aws_cloudwatch_client.check_sms",
side_effect=NotificationTechnicalFailureException(),
)
mock_sanitize = mocker.patch(
"app.celery.provider_tasks.sanitize_successful_notification_by_id"
)
try:
check_sms_delivery_receipt(
"message_id", sample_notification.id, "2024-10-20 00:00:00+0:00"
)
assert 1 == 0
except NotificationTechnicalFailureException:
mock_sanitize.assert_not_called()
mock_update.assert_called_once()
def test_should_call_send_sms_to_provider_from_deliver_sms_task(
sample_notification, mocker
):
+13 -11
View File
@@ -3,6 +3,7 @@ from decimal import Decimal
import pytest
from freezegun import freeze_time
from sqlalchemy import func, select
from app import db
from app.dao.fact_billing_dao import (
@@ -614,7 +615,8 @@ def test_delete_billing_data(notify_db_session):
delete_billing_data_for_service_for_day("2018-01-01", service_1.id)
current_rows = FactBilling.query.all()
stmt = select(FactBilling)
current_rows = db.session.execute(stmt).scalars().all()
assert sorted(x.billable_units for x in current_rows) == sorted(
[other_day.billable_units, other_service.billable_units]
)
@@ -671,7 +673,8 @@ def test_fetch_sms_free_allowance_remainder_until_date_with_two_services(
rate=0.11,
)
results = fetch_sms_free_allowance_remainder_until_date(datetime(2016, 5, 1)).all()
stmt = fetch_sms_free_allowance_remainder_until_date(datetime(2016, 5, 1))
results = db.session.execute(stmt).all()
assert len(results) == 2
service_result = [row for row in results if row[0] == service.id]
assert service_result[0] == (service.id, 10, 2, 8)
@@ -973,8 +976,8 @@ def test_fetch_usage_year_for_organization_populates_ft_billing_for_today(
free_sms_fragment_limit=10,
financial_year_start=current_year,
)
assert FactBilling.query.count() == 0
stmt = select(func.count()).select_from(FactBilling)
assert db.session.execute(stmt).scalar() == 0
create_notification(template=template, status=NotificationStatus.DELIVERED)
@@ -982,7 +985,7 @@ def test_fetch_usage_year_for_organization_populates_ft_billing_for_today(
organization_id=new_org.id, year=current_year
)
assert len(results) == 1
assert FactBilling.query.count() == 1
assert db.session.execute(stmt).scalar() == 1
@freeze_time("2022-05-01 13:30")
@@ -1224,8 +1227,8 @@ def test_query_organization_sms_usage_for_year_handles_multiple_services(
)
# ----------
result = query_organization_sms_usage_for_year(org.id, 2022).all()
stmt = query_organization_sms_usage_for_year(org.id, 2022)
result = db.session.execute(stmt).all()
service_1_rows = [row._asdict() for row in result if row.service_id == service_1.id]
service_2_rows = [row._asdict() for row in result if row.service_id == service_2.id]
@@ -1295,10 +1298,9 @@ def test_query_organization_sms_usage_for_year_handles_multiple_rates(
financial_year_start=current_year,
)
result = [
row._asdict()
for row in query_organization_sms_usage_for_year(org.id, 2022).all()
]
stmt = query_organization_sms_usage_for_year(org.id, 2022)
rows = db.session.execute(stmt).all()
result = [row._asdict() for row in rows]
# al lthe free allowance is used on the first day
assert result[0]["local_date"] == date(2022, 4, 29)
@@ -3,7 +3,9 @@ from uuid import UUID
import pytest
from freezegun import freeze_time
from sqlalchemy import func, select
from app import db
from app.dao.fact_notification_status_dao import (
fetch_monthly_notification_statuses_per_service,
fetch_monthly_template_usage_for_service,
@@ -1125,9 +1127,10 @@ def test_update_fact_notification_status_respects_gmt_bst(
process_day, NotificationType.SMS, sample_service.id
)
assert (
FactNotificationStatus.query.filter_by(
service_id=sample_service.id, local_date=process_day
).count()
== expected_count
stmt = (
select(func.count())
.select_from(FactNotificationStatus)
.filter_by(service_id=sample_service.id, local_date=process_day)
)
result = db.session.execute(stmt)
assert result.rowcount == expected_count
+8 -5
View File
@@ -4,8 +4,10 @@ from functools import partial
import pytest
from freezegun import freeze_time
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from app import db
from app.dao.jobs_dao import (
dao_create_job,
dao_get_future_scheduled_job_by_id_and_service_id,
@@ -108,7 +110,8 @@ def test_should_return_notifications_only_for_this_service(
def test_create_sample_job(sample_template):
assert Job.query.count() == 0
stmt = select(func.count()).select_from(Job)
assert db.session.execute(stmt).scalar() == 0
job_id = uuid.uuid4()
data = {
@@ -123,9 +126,9 @@ def test_create_sample_job(sample_template):
job = Job(**data)
dao_create_job(job)
assert Job.query.count() == 1
job_from_db = Job.query.get(job_id)
stmt = select(func.count()).select_from(Job)
assert db.session.execute(stmt).scalar() == 1
job_from_db = db.session.get(Job, job_id)
assert job == job_from_db
assert job_from_db.notifications_delivered == 0
assert job_from_db.notifications_failed == 0
@@ -221,7 +224,7 @@ def test_update_job(sample_job):
dao_update_job(sample_job)
job_from_db = Job.query.get(sample_job.id)
job_from_db = db.session.get(Job, sample_job.id)
assert job_from_db.job_status == JobStatus.IN_PROGRESS
+1 -1
View File
@@ -837,7 +837,7 @@ def test_get_jobs_should_paginate(admin_request, sample_template):
assert resp_json["page_size"] == 2
assert resp_json["total"] == 10
assert "links" in resp_json
assert set(resp_json["links"].keys()) == {"next", "last"}
assert set(resp_json["links"].keys()) == {"next", "last", "prev"}
def test_get_jobs_accepts_page_parameter(admin_request, sample_template):
+46
View File
@@ -1,4 +1,5 @@
import uuid
from unittest.mock import Mock
import pytest
from flask import current_app
@@ -12,6 +13,7 @@ from app.dao.organization_dao import (
from app.dao.services_dao import dao_archive_service
from app.enums import OrganizationType
from app.models import AnnualBilling, Organization
from app.organization.rest import check_request_args
from app.utils import utc_now
from tests.app.db import (
create_annual_billing,
@@ -928,3 +930,47 @@ def test_get_organization_services_usage_returns_400_if_year_is_empty(admin_requ
_expected_status=400,
)
assert response["message"] == "No valid year provided"
def test_valid_request_args():
request = Mock()
request.args = {"org_id": "123", "name": "Test Org"}
org_id, name = check_request_args(request)
assert org_id == "123"
assert name == "Test Org"
def test_missing_org_id():
request = Mock()
request.args = {"name": "Test Org"}
try:
check_request_args(request)
assert 1 == 0
except Exception as e:
assert e.status_code == 400
assert e.message == [{"org_id": ["Can't be empty"]}]
def test_missing_name():
request = Mock()
request.args = {"org_id": "123"}
try:
check_request_args(request)
assert 1 == 0
except Exception as e:
assert e.status_code == 400
assert e.message == [{"name": ["Can't be empty"]}]
def test_missing_both():
request = Mock()
request.args = {}
try:
check_request_args(request)
assert 1 == 0
except Exception as e:
assert e.status_code == 400
assert e.message == [
{"org_id": ["Can't be empty"]},
{"name": ["Can't be empty"]},
]
+80
View File
@@ -1,4 +1,5 @@
import collections
from collections import namedtuple
from datetime import datetime
from unittest.mock import Mock
@@ -12,6 +13,7 @@ from app.service.statistics import (
create_stats_dict,
create_zeroed_stats_dicts,
format_admin_stats,
format_monthly_template_notification_stats,
format_statistics,
)
@@ -337,3 +339,81 @@ def test_add_monthly_notification_status_stats():
},
"2018-06": {NotificationType.SMS: {}, NotificationType.EMAIL: {}},
}
def test_format_monthly_template_notification_stats():
Row = namedtuple(
"Row", ["month", "template_id", "name", "template_type", "status", "count"]
)
year = 2024
rows = [
Row(
datetime(2024, 4, 1), "1", "Template 1", "email", NotificationStatus.SENT, 5
),
Row(
datetime(2024, 4, 1),
"1",
"Template 1",
"email",
NotificationStatus.FAILED,
2,
),
Row(datetime(2024, 5, 1), "2", "Template 2", "sms", NotificationStatus.SENT, 3),
]
expected_output = {
"2024-04": {
"1": {
"name": "Template 1",
"type": "email",
"counts": {
NotificationStatus.CANCELLED: 0,
NotificationStatus.CREATED: 0,
NotificationStatus.DELIVERED: 0,
NotificationStatus.SENT: 5,
NotificationStatus.FAILED: 2,
NotificationStatus.PENDING: 0,
NotificationStatus.PENDING_VIRUS_CHECK: 0,
NotificationStatus.PERMANENT_FAILURE: 0,
NotificationStatus.SENDING: 0,
NotificationStatus.TECHNICAL_FAILURE: 0,
NotificationStatus.TEMPORARY_FAILURE: 0,
NotificationStatus.VALIDATION_FAILED: 0,
NotificationStatus.VIRUS_SCAN_FAILED: 0,
},
}
},
"2024-05": {
"2": {
"name": "Template 2",
"type": "sms",
"counts": {
NotificationStatus.CANCELLED: 0,
NotificationStatus.CREATED: 0,
NotificationStatus.DELIVERED: 0,
NotificationStatus.SENT: 3,
NotificationStatus.FAILED: 0,
NotificationStatus.PENDING: 0,
NotificationStatus.PENDING_VIRUS_CHECK: 0,
NotificationStatus.PERMANENT_FAILURE: 0,
NotificationStatus.SENDING: 0,
NotificationStatus.TECHNICAL_FAILURE: 0,
NotificationStatus.TEMPORARY_FAILURE: 0,
NotificationStatus.VALIDATION_FAILED: 0,
NotificationStatus.VIRUS_SCAN_FAILED: 0,
},
}
},
"2024-06": {},
"2024-07": {},
"2024-08": {},
"2024-09": {},
"2024-10": {},
"2024-11": {},
"2024-12": {},
"2025-01": {},
"2025-02": {},
"2025-03": {},
}
result = format_monthly_template_notification_stats(year, rows)
assert result == expected_output
+171 -3
View File
@@ -1,5 +1,6 @@
import datetime
import os
from datetime import datetime, timedelta
from unittest.mock import MagicMock, mock_open
import pytest
@@ -9,12 +10,16 @@ from app.commands import (
create_new_service,
create_test_user,
download_csv_file_by_name,
dump_sms_senders,
dump_user_info,
fix_billable_units,
insert_inbound_numbers_from_file,
populate_annual_billing_with_defaults,
populate_annual_billing_with_the_previous_years_allowance,
populate_go_live,
populate_organization_agreement_details_from_file,
populate_organizations_from_file,
process_row_from_job,
promote_user_to_platform_admin,
purge_functional_test_data,
update_jobs_archived_flag,
@@ -91,7 +96,8 @@ def test_purge_functional_test_data_bad_mobile(notify_db_session, notify_api):
"Fake Personson",
],
)
# The bad mobile phone number results in a bad parameter error, leading to a system exit 2 and no entry made in db
# The bad mobile phone number results in a bad parameter error,
# leading to a system exit 2 and no entry made in db
assert "SystemExit(2)" in str(command_response)
user_count = User.query.count()
assert user_count == 0
@@ -104,7 +110,7 @@ def test_update_jobs_archived_flag(notify_db_session, notify_api):
create_job(sms_template)
right_now = utc_now()
tomorrow = right_now + datetime.timedelta(days=1)
tomorrow = right_now + timedelta(days=1)
right_now = right_now.strftime("%Y-%m-%d")
tomorrow = tomorrow.strftime("%Y-%m-%d")
@@ -456,3 +462,165 @@ def test_promote_user_to_platform_admin_no_result_found(
)
assert "NoResultFound" in str(result)
assert sample_user.platform_admin is False
def test_populate_go_live_success(notify_api, mocker):
mock_csv_reader = mocker.patch("app.commands.csv.reader")
mocker.patch(
"app.commands.open",
new_callable=mock_open,
read_data="""count,Link,Service ID,DEPT,Service Name,Main contact,Contact detail,MOU,LIVE date,SMS,Email,Letters,CRM,Blue badge\n1,link,123,Dept A,Service A,Contact A,email@example.com,MOU,15/10/2024,Yes,Yes,Yes,Yes,No""", # noqa
)
mock_current_app = mocker.patch("app.commands.current_app")
mock_logger = mock_current_app.logger
mock_dao_update_service = mocker.patch("app.commands.dao_update_service")
mock_dao_fetch_service_by_id = mocker.patch("app.commands.dao_fetch_service_by_id")
mock_get_user_by_email = mocker.patch("app.commands.get_user_by_email")
mock_csv_reader.return_value = iter(
[
[
"count",
"Link",
"Service ID",
"DEPT",
"Service Name",
"Main contract",
"Contact detail",
"MOU",
"LIVE date",
"SMS",
"Email",
"Letters",
"CRM",
"Blue badge",
],
[
"1",
"link",
"123",
"Dept A",
"Service A",
"Contact A",
"email@example.com",
"MOU",
"15/10/2024",
"Yes",
"Yes",
"Yes",
"Yes",
"No",
],
]
)
mock_user = MagicMock()
mock_get_user_by_email.return_value = mock_user
mock_service = MagicMock()
mock_dao_fetch_service_by_id.return_value = mock_service
notify_api.test_cli_runner().invoke(
populate_go_live,
[
"-f",
"dummy_file.csv",
],
)
mock_get_user_by_email.assert_called_once_with("email@example.com")
mock_dao_fetch_service_by_id.assert_called_once_with("123")
mock_service.go_live_user = mock_user
mock_service.go_live_at = datetime.strptime("15/10/2024", "%d/%m/%Y") + timedelta(
hours=12
)
mock_dao_update_service.assert_called_once_with(mock_service)
mock_logger.info.assert_any_call("Populate go live user and date")
def test_process_row_from_job_success(notify_api, mocker):
mock_current_app = mocker.patch("app.commands.current_app")
mock_logger = mock_current_app.logger
mock_dao_get_job_by_id = mocker.patch("app.commands.dao_get_job_by_id")
mock_dao_get_template_by_id = mocker.patch("app.commands.dao_get_template_by_id")
mock_get_job_from_s3 = mocker.patch("app.commands.s3.get_job_from_s3")
mock_recipient_csv = mocker.patch("app.commands.RecipientCSV")
mock_process_row = mocker.patch("app.commands.process_row")
mock_job = MagicMock()
mock_job.service_id = "service_123"
mock_job.id = "job_456"
mock_job.template_id = "template_789"
mock_job.template_version = 1
mock_template = MagicMock()
mock_template._as_utils_template.return_value = MagicMock(
template_type="sms", placeholders=["name", "date"]
)
mock_row = MagicMock()
mock_row.index = 2
mock_recipient_csv.return_value.get_rows.return_value = [mock_row]
mock_dao_get_job_by_id.return_value = mock_job
mock_dao_get_template_by_id.return_value = mock_template
mock_get_job_from_s3.return_value = "some_csv_content"
mock_process_row.return_value = "notification_123"
notify_api.test_cli_runner().invoke(
process_row_from_job,
["-j", "job_456", "-n", "2"],
)
mock_dao_get_job_by_id.assert_called_once_with("job_456")
mock_dao_get_template_by_id.assert_called_once_with(
mock_job.template_id, mock_job.template_version
)
mock_get_job_from_s3.assert_called_once_with(
str(mock_job.service_id), str(mock_job.id)
)
mock_recipient_csv.assert_called_once_with(
"some_csv_content", template_type="sms", placeholders=["name", "date"]
)
mock_process_row.assert_called_once_with(
mock_row, mock_template._as_utils_template(), mock_job, mock_job.service
)
mock_logger.infoassert_called_once_with(
"Process row 2 for job job_456 created notification_id: notification_123"
)
def test_dump_sms_senders_single_service(notify_api, mocker):
mock_get_services_by_partial_name = mocker.patch(
"app.commands.get_services_by_partial_name"
)
mock_dao_get_sms_senders_by_service_id = mocker.patch(
"app.commands.dao_get_sms_senders_by_service_id"
)
mock_service = MagicMock()
mock_service.id = "service_123"
mock_get_services_by_partial_name.return_value = [mock_service]
mock_sender_1 = MagicMock()
mock_sender_1.serialize.return_value = {"name": "Sender 1", "id": "sender_1"}
mock_sender_2 = MagicMock()
mock_sender_2.serialize.return_value = {"name": "Sender 2", "id": "sender_2"}
mock_dao_get_sms_senders_by_service_id.return_value = [mock_sender_1, mock_sender_2]
notify_api.test_cli_runner().invoke(
dump_sms_senders,
["service_name"],
)
mock_get_services_by_partial_name.assert_called_once_with("service_name")
mock_dao_get_sms_senders_by_service_id.assert_called_once_with("service_123")
def test_dump_user_info(notify_api, mocker):
mock_open_file = mocker.patch("app.commands.open", new_callable=mock_open)
mock_get_user_by_email = mocker.patch("app.commands.get_user_by_email")
mock_user = MagicMock()
mock_user.serialize.return_value = {"name": "John Doe", "email": "john@example.com"}
mock_get_user_by_email.return_value = mock_user
notify_api.test_cli_runner().invoke(
dump_user_info,
["john@example.com"],
)
mock_get_user_by_email.assert_called_once_with("john@example.com")
mock_open_file.assert_called_once_with("user_download.json", "wb")
+38
View File
@@ -1,3 +1,5 @@
import datetime
import pytest
from marshmallow import ValidationError
from sqlalchemy import desc
@@ -7,6 +9,7 @@ from app.dao.provider_details_dao import (
get_provider_details_by_identifier,
)
from app.models import ProviderDetailsHistory
from app.schema_validation import validate_schema_date_with_hour
from tests.app.db import create_api_key
@@ -152,3 +155,38 @@ def test_provider_details_history_schema_returns_user_details(
data = provider_details_schema.dump(current_sms_provider_in_history)
assert sorted(data["created_by"].keys()) == sorted(["id", "email_address", "name"])
def test_valid_date_within_24_hours(mocker):
mocker.patch(
"app.schema_validation.utc_now",
return_value=datetime.datetime(2024, 10, 27, 15, 0, 0),
)
valid_datetime = "2024-10-28T14:00:00Z"
assert validate_schema_date_with_hour(valid_datetime)
def test_date_in_past(mocker):
mocker.patch(
"app.schema_validation.utc_now",
return_value=datetime.datetime(2024, 10, 27, 15, 0, 0),
)
past_datetime = "2024-10-26T14:00:00Z"
try:
validate_schema_date_with_hour(past_datetime)
assert 1 == 0
except Exception as e:
assert "datetime can not be in the past" in str(e)
def test_date_more_than_24_hours_in_future(mocker):
mocker.patch(
"app.schema_validation.utc_now",
return_value=datetime.datetime(2024, 10, 27, 15, 0, 0),
)
past_datetime = "2024-10-31T14:00:00Z"
try:
validate_schema_date_with_hour(past_datetime)
assert 1 == 0
except Exception as e:
assert "datetime can only be 24 hours in the future" in str(e)
+69
View File
@@ -0,0 +1,69 @@
from datetime import datetime
from unittest.mock import MagicMock
from app.upload.rest import get_paginated_uploads
def test_get_paginated_uploads(mocker):
mock_current_app = mocker.patch("app.upload.rest.current_app")
mock_dao_get_uploads = mocker.patch("app.upload.rest.dao_get_uploads_by_service_id")
mock_pagination_links = mocker.patch("app.upload.rest.pagination_links")
mock_fetch_notification_statuses = mocker.patch(
"app.upload.rest.fetch_notification_statuses_for_job"
)
mock_midnight_n_days_ago = mocker.patch("app.upload.rest.midnight_n_days_ago")
mock_dao_get_notification_outcomes = mocker.patch(
"app.upload.rest.dao_get_notification_outcomes_for_job"
)
mock_current_app.config = {"PAGE_SIZE": 10}
mock_pagination = MagicMock()
mock_pagination.items = [
MagicMock(
id="upload_1",
original_file_name="file1.csv",
notification_count=100,
scheduled_for=None,
created_at=datetime(2024, 10, 1, 12, 0, 0),
upload_type="job",
template_type="sms",
recipient="recipient@example.com",
processing_started=datetime(2024, 10, 2, 12, 0, 0),
),
MagicMock(
id="upload_2",
original_file_name="file2.csv",
notification_count=50,
scheduled_for=datetime(2024, 10, 3, 12, 0, 0),
created_at=None,
upload_type="letter",
template_type="letter",
recipient="recipient2@example.com",
processing_started=None,
),
]
mock_pagination.per_page = 10
mock_pagination.total = 2
mock_dao_get_uploads.return_value = mock_pagination
mock_midnight_n_days_ago.return_value = datetime(2024, 9, 30, 0, 0, 0)
mock_fetch_notification_statuses.return_value = [
MagicMock(status="delivered", count=90),
MagicMock(status="failed", count=10),
]
mock_dao_get_notification_outcomes.return_value = [
MagicMock(status="pending", count=40),
MagicMock(status="delivered", count=60),
]
mock_pagination_links.return_value = {"self": "/uploads?page=1"}
get_paginated_uploads("service_id_123", limit_days=7, page=1)
mock_dao_get_uploads.assert_called_once_with(
"service_id_123", limit_days=7, page=1, page_size=10
)
mock_midnight_n_days_ago.assert_called_once_with(3)
mock_dao_get_notification_outcomes.assert_called_once_with(
"service_id_123", "upload_1"
)
mock_pagination_links.assert_called_once_with(
mock_pagination, ".get_uploads_by_service", service_id="service_id_123"
)
+23
View File
@@ -64,3 +64,26 @@ def test_pii_filter():
pii_filter = logging.PIIFilter()
clean_msg = "phone1: 1XXXXXXXXXX, phone2: 1XXXXXXXXXX, email1: XXXXX@XXXXXXX, email2: XXXXX@XXXXXXX"
assert pii_filter.filter(record).msg == clean_msg
def test_process_log_record_successful(mocker):
mock_warning = mocker.patch("notifications_utils.logging.logger.warning")
log_record = {
"asctime": "2024-10-27 15:00:00",
"request_id": "12345",
"app_name": "test_app",
"service_id": "service_01",
"message": "Request 12345 received by test_app",
}
expected_output = {
"time": "2024-10-27 15:00:00",
"requestId": "12345",
"application": "test_app",
"service_id": "service_01",
"message": "Request 12345 received by test_app",
"logType": "application",
}
json_formatter = logging.JSONFormatter()
result = json_formatter.process_log_record(log_record)
assert result == expected_output
mock_warning.assert_not_called()