This commit is contained in:
Kenneth Kehl
2024-10-31 11:32:27 -07:00
parent 6b6bc2b4e7
commit bc7180185b
9 changed files with 74 additions and 28 deletions

View File

@@ -341,7 +341,7 @@
"filename": "tests/app/user/test_rest.py", "filename": "tests/app/user/test_rest.py",
"hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8", "hashed_secret": "5baa61e4c9b93f3f0682250b6cf8331b7ee68fd8",
"is_verified": false, "is_verified": false,
"line_number": 106, "line_number": 108,
"is_secret": false "is_secret": false
}, },
{ {
@@ -349,7 +349,7 @@
"filename": "tests/app/user/test_rest.py", "filename": "tests/app/user/test_rest.py",
"hashed_secret": "0beec7b5ea3f0fdbc95d0dd47f3c5bc275da8a33", "hashed_secret": "0beec7b5ea3f0fdbc95d0dd47f3c5bc275da8a33",
"is_verified": false, "is_verified": false,
"line_number": 810, "line_number": 822,
"is_secret": false "is_secret": false
} }
], ],
@@ -384,5 +384,5 @@
} }
] ]
}, },
"generated_at": "2024-10-30T18:15:03Z" "generated_at": "2024-10-31T18:32:23Z"
} }

View File

@@ -1,10 +1,10 @@
from datetime import timedelta from datetime import timedelta
from flask import current_app from flask import current_app
from sqlalchemy import between from sqlalchemy import between, select
from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.exc import SQLAlchemyError
from app import notify_celery, zendesk_client from app import db, notify_celery, zendesk_client
from app.celery.tasks import ( from app.celery.tasks import (
get_recipient_csv_and_template_and_sender_id, get_recipient_csv_and_template_and_sender_id,
process_incomplete_jobs, process_incomplete_jobs,
@@ -105,15 +105,18 @@ def check_job_status():
thirty_minutes_ago = utc_now() - timedelta(minutes=30) thirty_minutes_ago = utc_now() - timedelta(minutes=30)
thirty_five_minutes_ago = utc_now() - timedelta(minutes=35) thirty_five_minutes_ago = utc_now() - timedelta(minutes=35)
incomplete_in_progress_jobs = Job.query.filter( stmt = select(Job).where(
Job.job_status == JobStatus.IN_PROGRESS, Job.job_status == JobStatus.IN_PROGRESS,
between(Job.processing_started, thirty_five_minutes_ago, thirty_minutes_ago), between(Job.processing_started, thirty_five_minutes_ago, thirty_minutes_ago),
) )
incomplete_pending_jobs = Job.query.filter( incomplete_in_progress_jobs = db.session.execute(stmt).scalars().all()
stmt = select(Job).where(
Job.job_status == JobStatus.PENDING, Job.job_status == JobStatus.PENDING,
Job.scheduled_for.isnot(None), Job.scheduled_for.isnot(None),
between(Job.scheduled_for, thirty_five_minutes_ago, thirty_minutes_ago), between(Job.scheduled_for, thirty_five_minutes_ago, thirty_minutes_ago),
) )
incomplete_pending_jobs = db.session.execute(stmt).scalars().all()
jobs_not_complete_after_30_minutes = ( jobs_not_complete_after_30_minutes = (
incomplete_in_progress_jobs.union(incomplete_pending_jobs) incomplete_in_progress_jobs.union(incomplete_pending_jobs)

View File

@@ -646,8 +646,9 @@ def populate_annual_billing_with_defaults(year, missing_services_only):
This is useful to ensure all services start the new year with the correct annual billing. This is useful to ensure all services start the new year with the correct annual billing.
""" """
if missing_services_only: if missing_services_only:
active_services = ( stmt = (
Service.query.filter(Service.active) select(Service)
.where(Service.active)
.outerjoin( .outerjoin(
AnnualBilling, AnnualBilling,
and_( and_(
@@ -656,10 +657,11 @@ def populate_annual_billing_with_defaults(year, missing_services_only):
), ),
) )
.filter(AnnualBilling.id == None) # noqa .filter(AnnualBilling.id == None) # noqa
.all()
) )
active_services = db.session.execute(stmt).scalars().all()
else: else:
active_services = Service.query.filter(Service.active).all() stmt = select(Service).where(Service.active)
active_services = db.session.execute(stmt).scalars().all()
previous_year = year - 1 previous_year = year - 1
services_with_zero_free_allowance = ( services_with_zero_free_allowance = (
db.session.query(AnnualBilling.service_id) db.session.query(AnnualBilling.service_id)

View File

@@ -42,12 +42,17 @@ def test_move_notifications_does_nothing_if_notification_history_row_already_exi
1, 1,
) )
assert Notification.query.count() == 0 assert _get_notification_count() == 0
history = NotificationHistory.query.all() history = NotificationHistory.query.all()
assert len(history) == 1 assert len(history) == 1
assert history[0].status == NotificationStatus.DELIVERED assert history[0].status == NotificationStatus.DELIVERED
def _get_notification_count():
stmt = select(func.count()).select_from(Notification)
return db.session.execute(stmt).scalar() or 0
def test_move_notifications_only_moves_notifications_older_than_provided_timestamp( def test_move_notifications_only_moves_notifications_older_than_provided_timestamp(
sample_template, sample_template,
): ):
@@ -172,8 +177,10 @@ def test_move_notifications_just_deletes_test_key_notifications(sample_template)
assert result == 2 assert result == 2
assert Notification.query.count() == 0 assert _get_notification_count() == 0
assert NotificationHistory.query.count() == 2 stmt = select(func.count()).select_from(NotificationHistory)
count = db.session.execute(stmt).scalar() or 0
assert count == 2
stmt = ( stmt = (
select(func.count()) select(func.count())
.select_from(NotificationHistory) .select_from(NotificationHistory)

View File

@@ -1,9 +1,14 @@
from sqlalchemy import func, select
from app import db
from app.dao.events_dao import dao_create_event from app.dao.events_dao import dao_create_event
from app.models import Event from app.models import Event
def test_create_event(notify_db_session): def test_create_event(notify_db_session):
assert Event.query.count() == 0 stmt = select(func.count()).select_from(Event)
count = db.session.execute(stmt).scalar() or 0
assert count == 0
data = { data = {
"event_type": "sucessful_login", "event_type": "sucessful_login",
"data": {"something": "random", "in_fact": "could be anything"}, "data": {"something": "random", "in_fact": "could be anything"},
@@ -12,6 +17,8 @@ def test_create_event(notify_db_session):
event = Event(**data) event = Event(**data)
dao_create_event(event) dao_create_event(event)
assert Event.query.count() == 1 stmt = select(func.count()).select_from(Event)
count = db.session.execute(stmt).scalar() or 0
assert count == 1
event_from_db = Event.query.first() event_from_db = Event.query.first()
assert event == event_from_db assert event == event_from_db

View File

@@ -1,6 +1,8 @@
import pytest import pytest
from flask import current_app from flask import current_app
from sqlalchemy import func, select
from app import db
from app.dao.services_dao import dao_add_user_to_service from app.dao.services_dao import dao_add_user_to_service
from app.enums import NotificationType, TemplateType from app.enums import NotificationType, TemplateType
from app.models import Notification from app.models import Notification
@@ -23,7 +25,9 @@ def test_send_notification_to_service_users_persists_notifications_correctly(
notification = Notification.query.one() notification = Notification.query.one()
assert Notification.query.count() == 1 stmt = select(func.count()).select_from(Notification)
count = db.session.execute(stmt).scalar() or 0
assert count == 1
assert notification.to == "1" assert notification.to == "1"
assert str(notification.service_id) == current_app.config["NOTIFY_SERVICE_ID"] assert str(notification.service_id) == current_app.config["NOTIFY_SERVICE_ID"]
assert notification.template.id == template.id assert notification.template.id == template.id
@@ -89,4 +93,6 @@ def test_send_notification_to_service_users_sends_to_active_users_only(
send_notification_to_service_users(service_id=service.id, template_id=template.id) send_notification_to_service_users(service_id=service.id, template_id=template.id)
assert Notification.query.count() == 2 stmt = select(func.count()).select_from(Notification)
count = db.session.execute(stmt).scalar() or 0
assert count == 2

View File

@@ -1,7 +1,9 @@
import uuid import uuid
import pytest import pytest
from sqlalchemy import func, select
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.models import TemplateFolder from app.models import TemplateFolder
from tests.app.db import ( from tests.app.db import (
@@ -286,7 +288,9 @@ def test_delete_template_folder_fails_if_folder_has_subfolders(
assert resp == {"result": "error", "message": "Folder is not empty"} assert resp == {"result": "error", "message": "Folder is not empty"}
assert TemplateFolder.query.count() == 2 stmt = select(func.count()).select_from(TemplateFolder)
count = db.session.execute(stmt).scalar() or 0
assert count == 2
def test_delete_template_folder_fails_if_folder_contains_templates( def test_delete_template_folder_fails_if_folder_contains_templates(
@@ -304,7 +308,9 @@ def test_delete_template_folder_fails_if_folder_contains_templates(
assert resp == {"result": "error", "message": "Folder is not empty"} assert resp == {"result": "error", "message": "Folder is not empty"}
assert TemplateFolder.query.count() == 1 stmt = select(func.count()).select_from(TemplateFolder)
count = db.session.execute(stmt).scalar() or 0
assert count == 1
@pytest.mark.parametrize( @pytest.mark.parametrize(

View File

@@ -414,17 +414,20 @@ def test_create_service_command(notify_db_session, notify_api):
user = User.query.first() user = User.query.first()
service_count = Service.query.count() stmt = select(func.count()).select_from(Service)
service_count = db.session.execute(stmt).scalar() or 0
# run the command # run the command
result = notify_api.test_cli_runner().invoke( notify_api.test_cli_runner().invoke(
create_new_service, create_new_service,
["-e", "somebody@fake.gov", "-n", "Fake Service", "-c", user.id], ["-e", "somebody@fake.gov", "-n", "Fake Service", "-c", user.id],
) )
print(result)
# there should be one more service # there should be one more service
assert Service.query.count() == service_count + 1
stmt = select(func.count()).select_from(Service)
count = db.session.execute(stmt).scalar() or 0
assert count == service_count + 1
# that service should be the one we added # that service should be the one we added
stmt = select(Service).where(Service.name == "Fake Service") stmt = select(Service).where(Service.name == "Fake Service")

View File

@@ -6,7 +6,9 @@ from unittest import mock
import pytest import pytest
from flask import current_app from flask import current_app
from freezegun import freeze_time from freezegun import freeze_time
from sqlalchemy import func, select
from app import db
from app.dao.service_user_dao import dao_get_service_user, dao_update_service_user from app.dao.service_user_dao import dao_get_service_user, dao_update_service_user
from app.enums import AuthType, KeyType, NotificationType, PermissionType from app.enums import AuthType, KeyType, NotificationType, PermissionType
from app.models import Notification, Permission, User from app.models import Notification, Permission, User
@@ -153,12 +155,17 @@ def test_post_user_missing_attribute_email(admin_request, notify_db_session):
} }
json_resp = admin_request.post("user.create_user", _data=data, _expected_status=400) json_resp = admin_request.post("user.create_user", _data=data, _expected_status=400)
assert User.query.count() == 0 assert _get_user_count() == 0
assert {"email_address": ["Missing data for required field."]} == json_resp[ assert {"email_address": ["Missing data for required field."]} == json_resp[
"message" "message"
] ]
def _get_user_count():
stmt = select(func.count()).select_from(User)
return db.session.execute(stmt).scalar() or 0
def test_create_user_missing_attribute_password(admin_request, notify_db_session): def test_create_user_missing_attribute_password(admin_request, notify_db_session):
""" """
Tests POST endpoint '/' missing attribute password. Tests POST endpoint '/' missing attribute password.
@@ -174,7 +181,7 @@ def test_create_user_missing_attribute_password(admin_request, notify_db_session
"permissions": {}, "permissions": {},
} }
json_resp = admin_request.post("user.create_user", _data=data, _expected_status=400) json_resp = admin_request.post("user.create_user", _data=data, _expected_status=400)
assert User.query.count() == 0 assert _get_user_count() == 0
assert {"password": ["Missing data for required field."]} == json_resp["message"] assert {"password": ["Missing data for required field."]} == json_resp["message"]
@@ -512,8 +519,13 @@ def test_set_user_permissions_remove_old(admin_request, sample_user, sample_serv
_expected_status=204, _expected_status=204,
) )
query = Permission.query.filter_by(user=sample_user) query = (
assert query.count() == 1 select(func.count())
.select_from(Permission)
.where(Permission.user == sample_user)
)
count = db.session.execute(query).scalar() or 0
assert count == 1
assert query.first().permission == PermissionType.MANAGE_SETTINGS assert query.first().permission == PermissionType.MANAGE_SETTINGS