aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorstuppie2026-09-07 11:47:43 -0600
committerstuppie2026-09-07 11:47:43 -0600
commit092960233652cce1f4dc7841856034a6635e9cd9 (patch)
tree46e5fcd4d1e1b7ed0b987980c6c67ffa6e6b45c7
parent80fd8aab4c7271ddb619b0de18741d7ac77b490b (diff)
parent242579a44855873d5e054e375440e9d3492cd682 (diff)
downloadgeneralresearch-092960233652cce1f4dc7841856034a6635e9cd9.tar.gz
generalresearch-092960233652cce1f4dc7841856034a6635e9cd9.zip
Merge branch 'master' into dev-greg
-rw-r--r--Jenkinsfile237
-rw-r--r--generalresearch/__init__.py22
-rw-r--r--generalresearch/config.py21
-rw-r--r--generalresearch/currency.py38
-rw-r--r--generalresearch/grliq/managers/__init__.py34
-rw-r--r--generalresearch/grliq/managers/event_plotter.py16
-rw-r--r--generalresearch/grliq/managers/forensic_data.py147
-rw-r--r--generalresearch/grliq/managers/forensic_events.py132
-rw-r--r--generalresearch/grliq/managers/forensic_results.py30
-rw-r--r--generalresearch/grliq/managers/forensic_summary.py36
-rw-r--r--generalresearch/grliq/models/__init__.py4
-rw-r--r--generalresearch/grliq/models/custom_types.py3
-rw-r--r--generalresearch/grliq/models/decider.py10
-rw-r--r--generalresearch/grliq/models/events.py40
-rw-r--r--generalresearch/grliq/models/forensic_data.py49
-rw-r--r--generalresearch/grliq/models/forensic_result.py4
-rw-r--r--generalresearch/grliq/models/forensic_summary.py29
-rw-r--r--generalresearch/grliq/models/useragents.py16
-rw-r--r--generalresearch/grliq/utils.py5
-rw-r--r--generalresearch/grpc.py8
-rw-r--r--generalresearch/healing_ppe.py2
-rw-r--r--generalresearch/incite/__init__.py4
-rw-r--r--generalresearch/incite/base.py58
-rw-r--r--generalresearch/incite/collections/__init__.py757
-rw-r--r--generalresearch/incite/collections/base.py777
-rw-r--r--generalresearch/incite/collections/thl_marketplaces.py2
-rw-r--r--generalresearch/incite/collections/thl_web.py2
-rw-r--r--generalresearch/incite/defaults.py78
-rw-r--r--generalresearch/incite/exceptions.py19
-rw-r--r--generalresearch/incite/mergers/__init__.py305
-rw-r--r--generalresearch/incite/mergers/base.py311
-rw-r--r--generalresearch/incite/mergers/foundations/__init__.py20
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_session.py45
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_task_adjust.py18
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_wall.py21
-rw-r--r--generalresearch/incite/mergers/foundations/user_id_product.py6
-rw-r--r--generalresearch/incite/mergers/pop_ledger.py6
-rw-r--r--generalresearch/incite/mergers/ym_survey_wall.py22
-rw-r--r--generalresearch/incite/mergers/ym_wall_summary.py20
-rw-r--r--generalresearch/incite/schemas/__init__.py6
-rw-r--r--generalresearch/incite/schemas/admin_responses.py73
-rw-r--r--generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py6
-rw-r--r--generalresearch/incite/schemas/mergers/foundations/enriched_wall.py2
-rw-r--r--generalresearch/incite/schemas/mergers/foundations/user_id_product.py2
-rw-r--r--generalresearch/incite/schemas/mergers/pop_ledger.py10
-rw-r--r--generalresearch/incite/schemas/mergers/ym_wall_summary.py2
-rw-r--r--generalresearch/incite/schemas/thl_web.py16
-rw-r--r--generalresearch/locales/__init__.py9
-rw-r--r--generalresearch/locales/setup_json.py123
-rw-r--r--generalresearch/locales/timezone.py5
-rw-r--r--generalresearch/logging.py5
-rw-r--r--generalresearch/managers/__init__.py16
-rw-r--r--generalresearch/managers/base.py14
-rw-r--r--generalresearch/managers/cint/profiling.py7
-rw-r--r--generalresearch/managers/cint/survey.py6
-rw-r--r--generalresearch/managers/cint/user_pid.py2
-rw-r--r--generalresearch/managers/criteria.py31
-rw-r--r--generalresearch/managers/dynata/profiling.py7
-rw-r--r--generalresearch/managers/dynata/survey.py70
-rw-r--r--generalresearch/managers/dynata/user_pid.py2
-rw-r--r--generalresearch/managers/events.py115
-rw-r--r--generalresearch/managers/gr/authentication.py33
-rw-r--r--generalresearch/managers/gr/business.py24
-rw-r--r--generalresearch/managers/gr/team.py21
-rw-r--r--generalresearch/managers/innovate/profiling.py5
-rw-r--r--generalresearch/managers/innovate/survey.py91
-rw-r--r--generalresearch/managers/innovate/user_pid.py2
-rw-r--r--generalresearch/managers/leaderboard/__init__.py3
-rw-r--r--generalresearch/managers/leaderboard/manager.py8
-rw-r--r--generalresearch/managers/leaderboard/tasks.py5
-rw-r--r--generalresearch/managers/lucid/profiling.py8
-rw-r--r--generalresearch/managers/marketplace/__init__.py23
-rw-r--r--generalresearch/managers/marketplace/managers.py23
-rw-r--r--generalresearch/managers/marketplace/user_pid.py7
-rw-r--r--generalresearch/managers/morning/profiling.py5
-rw-r--r--generalresearch/managers/morning/survey.py97
-rw-r--r--generalresearch/managers/morning/user_pid.py2
-rw-r--r--generalresearch/managers/network/label.py147
-rw-r--r--generalresearch/managers/network/mtr.py49
-rw-r--r--generalresearch/managers/network/nmap.py57
-rw-r--r--generalresearch/managers/network/rdns.py33
-rw-r--r--generalresearch/managers/network/tool_run.py138
-rw-r--r--generalresearch/managers/pollfish/profiling.py5
-rw-r--r--generalresearch/managers/pollfish/user_pid.py2
-rw-r--r--generalresearch/managers/precision/profiling.py5
-rw-r--r--generalresearch/managers/precision/survey.py65
-rw-r--r--generalresearch/managers/precision/user_pid.py2
-rw-r--r--generalresearch/managers/prodege/profiling.py5
-rw-r--r--generalresearch/managers/prodege/survey.py59
-rw-r--r--generalresearch/managers/prodege/user_pid.py2
-rw-r--r--generalresearch/managers/repdata/profiling.py5
-rw-r--r--generalresearch/managers/repdata/survey.py94
-rw-r--r--generalresearch/managers/repdata/user_pid.py2
-rw-r--r--generalresearch/managers/sago/profiling.py5
-rw-r--r--generalresearch/managers/sago/survey.py75
-rw-r--r--generalresearch/managers/sago/user_pid.py2
-rw-r--r--generalresearch/managers/spectrum/profiling.py5
-rw-r--r--generalresearch/managers/spectrum/survey.py73
-rw-r--r--generalresearch/managers/spectrum/user_pid.py2
-rw-r--r--generalresearch/managers/survey.py7
-rw-r--r--generalresearch/managers/thl/buyer.py15
-rw-r--r--generalresearch/managers/thl/cashout_method.py51
-rw-r--r--generalresearch/managers/thl/category.py15
-rw-r--r--generalresearch/managers/thl/contest_manager.py48
-rw-r--r--generalresearch/managers/thl/ipinfo.py62
-rw-r--r--generalresearch/managers/thl/ledger_manager/conditions.py22
-rw-r--r--generalresearch/managers/thl/ledger_manager/exceptions.py5
-rw-r--r--generalresearch/managers/thl/ledger_manager/ledger.py35
-rw-r--r--generalresearch/managers/thl/ledger_manager/thl_ledger.py29
-rw-r--r--generalresearch/managers/thl/payout.py1039
-rw-r--r--generalresearch/managers/thl/product.py54
-rw-r--r--generalresearch/managers/thl/profiling/question.py2
-rw-r--r--generalresearch/managers/thl/profiling/uqa.py9
-rw-r--r--generalresearch/managers/thl/profiling/user_upk.py24
-rw-r--r--generalresearch/managers/thl/session.py65
-rw-r--r--generalresearch/managers/thl/survey.py158
-rw-r--r--generalresearch/managers/thl/survey_penalty.py23
-rw-r--r--generalresearch/managers/thl/tango_api.py2
-rw-r--r--generalresearch/managers/thl/task_adjustment.py31
-rw-r--r--generalresearch/managers/thl/user_compensate.py15
-rw-r--r--generalresearch/managers/thl/user_manager/__init__.py33
-rw-r--r--generalresearch/managers/thl/user_manager/exceptions.py6
-rw-r--r--generalresearch/managers/thl/user_manager/mysql_user_manager.py89
-rw-r--r--generalresearch/managers/thl/user_manager/rate_limit.py9
-rw-r--r--generalresearch/managers/thl/user_manager/user_manager.py31
-rw-r--r--generalresearch/managers/thl/userhealth.py39
-rw-r--r--generalresearch/managers/thl/wall.py39
-rw-r--r--generalresearch/managers/thl/wallet/__init__.py38
-rw-r--r--generalresearch/managers/thl/wallet/approve.py18
-rw-r--r--generalresearch/managers/thl/wallet/tango.py31
-rw-r--r--generalresearch/managers/utils.py16
-rw-r--r--generalresearch/mariadb.py8
-rw-r--r--generalresearch/models/__init__.py115
-rw-r--r--generalresearch/models/admin/__init__.py14
-rw-r--r--generalresearch/models/admin/request.py14
-rw-r--r--generalresearch/models/cint/__init__.py3
-rw-r--r--generalresearch/models/cint/question.py28
-rw-r--r--generalresearch/models/cint/survey.py26
-rw-r--r--generalresearch/models/cint/task_collection.py6
-rw-r--r--generalresearch/models/custom_types.py24
-rw-r--r--generalresearch/models/definitions.py114
-rw-r--r--generalresearch/models/device.py2
-rw-r--r--generalresearch/models/dynata/__init__.py4
-rw-r--r--generalresearch/models/dynata/question.py6
-rw-r--r--generalresearch/models/dynata/survey.py25
-rw-r--r--generalresearch/models/dynata/task_collection.py2
-rw-r--r--generalresearch/models/events.py79
-rw-r--r--generalresearch/models/gr/authentication.py31
-rw-r--r--generalresearch/models/gr/business.py200
-rw-r--r--generalresearch/models/gr/definitions.py13
-rw-r--r--generalresearch/models/gr/team.py108
-rw-r--r--generalresearch/models/innovate/__init__.py10
-rw-r--r--generalresearch/models/innovate/question.py26
-rw-r--r--generalresearch/models/innovate/survey.py64
-rw-r--r--generalresearch/models/legacy/bucket.py41
-rw-r--r--generalresearch/models/legacy/definitions.py4
-rw-r--r--generalresearch/models/legacy/questions.py42
-rw-r--r--generalresearch/models/lucid/__init__.py3
-rw-r--r--generalresearch/models/lucid/question.py17
-rw-r--r--generalresearch/models/lucid/survey.py2
-rw-r--r--generalresearch/models/marketplace/summary.py3
-rw-r--r--generalresearch/models/morning/__init__.py6
-rw-r--r--generalresearch/models/morning/question.py19
-rw-r--r--generalresearch/models/morning/survey.py102
-rw-r--r--generalresearch/models/morning/task_collection.py2
-rw-r--r--generalresearch/models/network/__init__.py0
-rw-r--r--generalresearch/models/network/definitions.py70
-rw-r--r--generalresearch/models/network/label.py125
-rw-r--r--generalresearch/models/network/mtr/__init__.py0
-rw-r--r--generalresearch/models/network/mtr/command.py73
-rw-r--r--generalresearch/models/network/mtr/execute.py55
-rw-r--r--generalresearch/models/network/mtr/parser.py18
-rw-r--r--generalresearch/models/network/mtr/result.py168
-rw-r--r--generalresearch/models/network/nmap/__init__.py0
-rw-r--r--generalresearch/models/network/nmap/command.py48
-rw-r--r--generalresearch/models/network/nmap/execute.py51
-rw-r--r--generalresearch/models/network/nmap/parser.py414
-rw-r--r--generalresearch/models/network/nmap/result.py434
-rw-r--r--generalresearch/models/network/rdns/__init__.py0
-rw-r--r--generalresearch/models/network/rdns/command.py35
-rw-r--r--generalresearch/models/network/rdns/execute.py44
-rw-r--r--generalresearch/models/network/rdns/parser.py21
-rw-r--r--generalresearch/models/network/rdns/result.py52
-rw-r--r--generalresearch/models/network/tool_run.py118
-rw-r--r--generalresearch/models/network/tool_run_command.py66
-rw-r--r--generalresearch/models/network/utils.py5
-rw-r--r--generalresearch/models/pollfish/question.py16
-rw-r--r--generalresearch/models/precision/__init__.py6
-rw-r--r--generalresearch/models/precision/question.py29
-rw-r--r--generalresearch/models/precision/survey.py90
-rw-r--r--generalresearch/models/precision/task_collection.py10
-rw-r--r--generalresearch/models/prodege/__init__.py9
-rw-r--r--generalresearch/models/prodege/question.py33
-rw-r--r--generalresearch/models/prodege/survey.py29
-rw-r--r--generalresearch/models/prodege/task_collection.py2
-rw-r--r--generalresearch/models/repdata/__init__.py4
-rw-r--r--generalresearch/models/repdata/question.py20
-rw-r--r--generalresearch/models/repdata/survey.py32
-rw-r--r--generalresearch/models/repdata/task_collection.py4
-rw-r--r--generalresearch/models/sago/__init__.py6
-rw-r--r--generalresearch/models/sago/question.py26
-rw-r--r--generalresearch/models/sago/survey.py39
-rw-r--r--generalresearch/models/spectrum/__init__.py2
-rw-r--r--generalresearch/models/spectrum/question.py45
-rw-r--r--generalresearch/models/spectrum/survey.py36
-rw-r--r--generalresearch/models/spectrum/task_collection.py4
-rw-r--r--generalresearch/models/string_utils.py3
-rw-r--r--generalresearch/models/thl/__init__.py24
-rw-r--r--generalresearch/models/thl/category.py3
-rw-r--r--generalresearch/models/thl/contest/__init__.py7
-rw-r--r--generalresearch/models/thl/contest/contest.py35
-rw-r--r--generalresearch/models/thl/contest/contest_entry.py30
-rw-r--r--generalresearch/models/thl/contest/definitions.py16
-rw-r--r--generalresearch/models/thl/contest/examples.py416
-rw-r--r--generalresearch/models/thl/contest/io.py4
-rw-r--r--generalresearch/models/thl/contest/leaderboard.py165
-rw-r--r--generalresearch/models/thl/contest/milestone.py132
-rw-r--r--generalresearch/models/thl/contest/raffle.py177
-rw-r--r--generalresearch/models/thl/definitions.py24
-rw-r--r--generalresearch/models/thl/demographics.py15
-rw-r--r--generalresearch/models/thl/finance.py57
-rw-r--r--generalresearch/models/thl/ipinfo.py43
-rw-r--r--generalresearch/models/thl/leaderboard.py18
-rw-r--r--generalresearch/models/thl/ledger.py108
-rw-r--r--generalresearch/models/thl/ledger_example.py64
-rw-r--r--generalresearch/models/thl/maxmind/__init__.py0
-rw-r--r--generalresearch/models/thl/maxmind/definitions.py22
-rw-r--r--generalresearch/models/thl/offerwall/__init__.py13
-rw-r--r--generalresearch/models/thl/offerwall/base.py21
-rw-r--r--generalresearch/models/thl/offerwall/cache.py12
-rw-r--r--generalresearch/models/thl/payout.py251
-rw-r--r--generalresearch/models/thl/payout_format.py14
-rw-r--r--generalresearch/models/thl/product.py134
-rw-r--r--generalresearch/models/thl/profiling/marketplace.py15
-rw-r--r--generalresearch/models/thl/profiling/other_option.py6
-rw-r--r--generalresearch/models/thl/profiling/upk_property.py6
-rw-r--r--generalresearch/models/thl/profiling/upk_question.py36
-rw-r--r--generalresearch/models/thl/profiling/upk_question_answer.py11
-rw-r--r--generalresearch/models/thl/profiling/user_info.py2
-rw-r--r--generalresearch/models/thl/profiling/user_question_answer.py68
-rw-r--r--generalresearch/models/thl/report_task.py6
-rw-r--r--generalresearch/models/thl/session.py265
-rw-r--r--generalresearch/models/thl/soft_pair.py15
-rw-r--r--generalresearch/models/thl/supplier_tag.py4
-rw-r--r--generalresearch/models/thl/survey/__init__.py8
-rw-r--r--generalresearch/models/thl/survey/buyer.py9
-rw-r--r--generalresearch/models/thl/survey/condition.py16
-rw-r--r--generalresearch/models/thl/survey/model.py31
-rw-r--r--generalresearch/models/thl/survey/penalty.py8
-rw-r--r--generalresearch/models/thl/survey/task_collection.py3
-rw-r--r--generalresearch/models/thl/task_adjustment.py8
-rw-r--r--generalresearch/models/thl/task_status.py36
-rw-r--r--generalresearch/models/thl/user.py84
-rw-r--r--generalresearch/models/thl/user_identifiers.py33
-rw-r--r--generalresearch/models/thl/user_iphistory.py30
-rw-r--r--generalresearch/models/thl/user_profile.py5
-rw-r--r--generalresearch/models/thl/user_quality_event.py14
-rw-r--r--generalresearch/models/thl/user_ref.py17
-rw-r--r--generalresearch/models/thl/user_streak.py12
-rw-r--r--generalresearch/models/thl/userhealth.py23
-rw-r--r--generalresearch/models/thl/utils.py11
-rw-r--r--generalresearch/models/thl/wallet/__init__.py87
-rw-r--r--generalresearch/models/thl/wallet/cashout_method.py46
-rw-r--r--generalresearch/models/thl/wallet/definitions.py87
-rw-r--r--generalresearch/models/thl/wallet/payout.py17
-rw-r--r--generalresearch/pg_helper.py31
-rw-r--r--generalresearch/schemas/survey_stats.py4
-rw-r--r--generalresearch/sql_helper.py36
-rw-r--r--generalresearch/thl_django/app/manage.py4
-rw-r--r--generalresearch/thl_django/app/test_settings.py17
-rw-r--r--generalresearch/thl_django/apps.py14
-rw-r--r--generalresearch/thl_django/event/models.py13
-rw-r--r--generalresearch/thl_django/fields.py3
-rw-r--r--generalresearch/thl_django/migrations/0001_initial.py3
-rw-r--r--generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py2
-rw-r--r--generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py2
-rw-r--r--generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py2
-rw-r--r--generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py5
-rw-r--r--generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py51
-rw-r--r--generalresearch/thl_django/network/models.py5
-rw-r--r--generalresearch/utils/aggregation.py4
-rw-r--r--generalresearch/utils/enum.py7
-rw-r--r--generalresearch/utils/grpc_logger.py6
-rw-r--r--generalresearch/wall_status_codes/__init__.py14
-rw-r--r--generalresearch/wall_status_codes/cint.py12
-rw-r--r--generalresearch/wall_status_codes/dynata.py29
-rw-r--r--generalresearch/wall_status_codes/fullcircle.py16
-rw-r--r--generalresearch/wall_status_codes/innovate.py16
-rw-r--r--generalresearch/wall_status_codes/lucid.py20
-rw-r--r--generalresearch/wall_status_codes/morning.py18
-rw-r--r--generalresearch/wall_status_codes/pollfish.py18
-rw-r--r--generalresearch/wall_status_codes/precision.py18
-rw-r--r--generalresearch/wall_status_codes/prodege.py14
-rw-r--r--generalresearch/wall_status_codes/repdata.py20
-rw-r--r--generalresearch/wall_status_codes/sago.py18
-rw-r--r--generalresearch/wall_status_codes/spectrum.py16
-rw-r--r--generalresearch/wall_status_codes/wxet.py19
-rw-r--r--generalresearch/wxet/models/definitions.py49
-rw-r--r--generalresearch/wxet/models/finish_type.py13
-rw-r--r--pyproject.toml16
-rw-r--r--requirements.txt112
-rw-r--r--test_utils/conftest.py266
-rw-r--r--test_utils/grliq/conftest.py214
-rw-r--r--test_utils/incite/collections/conftest.py9
-rw-r--r--test_utils/incite/conftest.py17
-rw-r--r--test_utils/incite/mergers/conftest.py73
-rw-r--r--test_utils/managers/cashout_methods.py75
-rw-r--r--test_utils/managers/conftest.py269
-rw-r--r--test_utils/managers/contest/conftest.py10
-rw-r--r--test_utils/managers/gr/conftest.py87
-rw-r--r--test_utils/managers/ledger/conftest.py14
-rw-r--r--test_utils/managers/network/__init__.py0
-rw-r--r--test_utils/managers/network/conftest.py0
-rw-r--r--test_utils/managers/thl/conftest.py189
-rw-r--r--test_utils/managers/upk/conftest.py3
-rw-r--r--test_utils/models/conftest.py472
-rw-r--r--test_utils/models/contest/conftest.py102
-rw-r--r--test_utils/models/gr/conftest.py337
-rw-r--r--test_utils/models/ledger/conftest.py274
-rw-r--r--test_utils/models/network/__init__.py0
-rw-r--r--test_utils/models/network/conftest.py144
-rw-r--r--test_utils/models/thl/conftest.py995
-rw-r--r--test_utils/models/upk/conftest.py16
-rw-r--r--test_utils/precision/__init__.py (renamed from generalresearch/managers/network/__init__.py)0
-rw-r--r--test_utils/precision/conftest.py129
-rw-r--r--test_utils/spectrum/conftest.py255
-rw-r--r--test_utils/spectrum/surveys_json.py140
-rw-r--r--tests/conftest.py5
-rw-r--r--tests/grliq/managers/test_forensic_data.py65
-rw-r--r--tests/grliq/managers/test_forensic_results.py11
-rw-r--r--tests/grliq/models/test_forensic_data.py8
-rw-r--r--tests/incite/collections/test_df_collection_base.py38
-rw-r--r--tests/incite/collections/test_df_collection_item_base.py53
-rw-r--r--tests/incite/collections/test_df_collection_item_thl_web.py454
-rw-r--r--tests/incite/collections/test_df_collection_thl_marketplaces.py30
-rw-r--r--tests/incite/collections/test_df_collection_thl_web.py131
-rw-r--r--tests/incite/mergers/foundations/test_enriched_session.py76
-rw-r--r--tests/incite/mergers/foundations/test_enriched_task_adjust.py49
-rw-r--r--tests/incite/mergers/foundations/test_enriched_wall.py119
-rw-r--r--tests/incite/mergers/foundations/test_user_id_product.py47
-rw-r--r--tests/incite/mergers/test_merge_collection.py73
-rw-r--r--tests/incite/mergers/test_merge_collection_item.py37
-rw-r--r--tests/incite/mergers/test_pop_ledger.py148
-rw-r--r--tests/incite/mergers/test_ym_survey_merge.py84
-rw-r--r--tests/incite/schemas/test_admin_responses.py62
-rw-r--r--tests/incite/schemas/test_thl_web.py8
-rw-r--r--tests/incite/test_collection_base.py91
-rw-r--r--tests/incite/test_collection_base_item.py82
-rw-r--r--tests/incite/test_grl_flow.py11
-rw-r--r--tests/incite/test_interval_idx.py9
-rw-r--r--tests/managers/gr/test_authentication.py122
-rw-r--r--tests/managers/gr/test_business.py139
-rw-r--r--tests/managers/gr/test_team.py130
-rw-r--r--tests/managers/leaderboard.py177
-rw-r--r--tests/managers/network/__init__.py0
-rw-r--r--tests/managers/network/test_label.py202
-rw-r--r--tests/managers/network/test_tool_run.py25
-rw-r--r--tests/managers/test_events.py146
-rw-r--r--tests/managers/test_lucid.py9
-rw-r--r--tests/managers/test_userpid.py6
-rw-r--r--tests/managers/thl/test_buyer.py16
-rw-r--r--tests/managers/thl/test_cashout_method.py76
-rw-r--r--tests/managers/thl/test_category.py112
-rw-r--r--tests/managers/thl/test_contest/test_leaderboard.py89
-rw-r--r--tests/managers/thl/test_contest/test_milestone.py136
-rw-r--r--tests/managers/thl/test_contest/test_raffle.py232
-rw-r--r--tests/managers/thl/test_harmonized_uqa.py21
-rw-r--r--tests/managers/thl/test_ipinfo.py85
-rw-r--r--tests/managers/thl/test_ledger/test_lm_accounts.py162
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx.py145
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_entries.py30
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_locks.py283
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_metadata.py43
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_accounts.py329
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py473
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx.py1213
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py404
-rw-r--r--tests/managers/thl/test_ledger/test_thl_pem.py139
-rw-r--r--tests/managers/thl/test_ledger/test_user_txs.py119
-rw-r--r--tests/managers/thl/test_ledger/test_wallet.py51
-rw-r--r--tests/managers/thl/test_maxmind.py500
-rw-r--r--tests/managers/thl/test_payout.py1073
-rw-r--r--tests/managers/thl/test_product.py119
-rw-r--r--tests/managers/thl/test_product_prod.py27
-rw-r--r--tests/managers/thl/test_profiling/test_question.py32
-rw-r--r--tests/managers/thl/test_profiling/test_schema.py18
-rw-r--r--tests/managers/thl/test_profiling/test_uqa.py1
-rw-r--r--tests/managers/thl/test_profiling/test_user_upk.py28
-rw-r--r--tests/managers/thl/test_session_manager.py94
-rw-r--r--tests/managers/thl/test_survey.py96
-rw-r--r--tests/managers/thl/test_survey_penalty.py27
-rw-r--r--tests/managers/thl/test_task_adjustment.py248
-rw-r--r--tests/managers/thl/test_task_status.py169
-rw-r--r--tests/managers/thl/test_user_manager/test_base.py76
-rw-r--r--tests/managers/thl/test_user_manager/test_mysql.py31
-rw-r--r--tests/managers/thl/test_user_manager/test_redis.py46
-rw-r--r--tests/managers/thl/test_user_manager/test_user_fetch.py19
-rw-r--r--tests/managers/thl/test_user_manager/test_user_metadata.py47
-rw-r--r--tests/managers/thl/test_user_streak.py73
-rw-r--r--tests/managers/thl/test_userhealth.py169
-rw-r--r--tests/managers/thl/test_wall_manager.py106
-rw-r--r--tests/models/admin/test_report_request.py20
-rw-r--r--tests/models/custom_types/test_aware_datetime.py8
-rw-r--r--tests/models/custom_types/test_dsn.py10
-rw-r--r--tests/models/custom_types/test_therest.py2
-rw-r--r--tests/models/dynata/test_eligbility.py8
-rw-r--r--tests/models/dynata/test_survey.py3
-rw-r--r--tests/models/gr/test_authentication.py222
-rw-r--r--tests/models/gr/test_base.py21
-rw-r--r--tests/models/gr/test_business.py1134
-rw-r--r--tests/models/gr/test_team.py356
-rw-r--r--tests/models/innovate/test_question.py10
-rw-r--r--tests/models/legacy/test_offerwall_parse_response.py4
-rw-r--r--tests/models/legacy/test_profiling_questions.py6
-rw-r--r--tests/models/legacy/test_user_question_answer_in.py69
-rw-r--r--tests/models/morning/test.py8
-rw-r--r--tests/models/network/__init__.py0
-rw-r--r--tests/models/network/test_mtr.py26
-rw-r--r--tests/models/network/test_nmap.py30
-rw-r--r--tests/models/network/test_nmap_parser.py22
-rw-r--r--tests/models/network/test_rdns.py34
-rw-r--r--tests/models/precision/__init__.py115
-rw-r--r--tests/models/precision/test_survey.py42
-rw-r--r--tests/models/prodege/test_survey_participation.py23
-rw-r--r--tests/models/spectrum/test_question.py21
-rw-r--r--tests/models/spectrum/test_survey.py86
-rw-r--r--tests/models/spectrum/test_survey_manager.py110
-rw-r--r--tests/models/test_currency.py126
-rw-r--r--tests/models/test_device.py8
-rw-r--r--tests/models/test_finance.py127
-rw-r--r--tests/models/thl/question/test_question_info.py139
-rw-r--r--tests/models/thl/question/test_user_info.py29
-rw-r--r--tests/models/thl/test_adjustments.py128
-rw-r--r--tests/models/thl/test_bucket.py8
-rw-r--r--tests/models/thl/test_buyer.py4
-rw-r--r--tests/models/thl/test_contest/test_contest.py10
-rw-r--r--tests/models/thl/test_contest/test_leaderboard_contest.py44
-rw-r--r--tests/models/thl/test_contest/test_raffle_contest.py44
-rw-r--r--tests/models/thl/test_ledger.py6
-rw-r--r--tests/models/thl/test_marketplace_condition.py42
-rw-r--r--tests/models/thl/test_payout.py120
-rw-r--r--tests/models/thl/test_payout_format.py2
-rw-r--r--tests/models/thl/test_product.py305
-rw-r--r--tests/models/thl/test_product_userwalletconfig.py10
-rw-r--r--tests/models/thl/test_soft_pair.py12
-rw-r--r--tests/models/thl/test_upkquestion.py86
-rw-r--r--tests/models/thl/test_user.py141
-rw-r--r--tests/models/thl/test_user_iphistory.py6
-rw-r--r--tests/models/thl/test_user_metadata.py4
-rw-r--r--tests/models/thl/test_user_streak.py4
-rw-r--r--tests/models/thl/test_wall.py42
-rw-r--r--tests/models/thl/test_wall_session.py24
-rw-r--r--tests/sql_helper.py4
-rw-r--r--tests/test_postgres.py39
-rw-r--r--tests/wall_status_codes/test_analyze.py2
-rw-r--r--tests/wxet/models/test_definitions.py37
-rw-r--r--tests/wxet/models/test_finish_type.py2
457 files changed, 15320 insertions, 16527 deletions
diff --git a/Jenkinsfile b/Jenkinsfile
index e829ba9..a646d22 100644
--- a/Jenkinsfile
+++ b/Jenkinsfile
@@ -12,243 +12,86 @@ pipeline {
environment {
VENV = "${env.WORKSPACE}/generalresearch-venv"
- SPECTRUM_CARER_VENV = "${env.WORKSPACE}/thl-spectrum-carer-venv"
- GRLIQ_CARER_VENV = "${env.WORKSPACE}/grliq-carer-venv"
- GR_CARER_VENV = "${env.WORKSPACE}/gr-carer-venv"
-
- INCITE_MOUNT_DIR = '/mnt/thl-incite'
- TMP_DIR = "${env.WORKSPACE}/tmp"
}
stages {
- stage('python versions') {
+
+ stage('Checkout') {
+ steps {
+ checkout scmGit(
+ branches: [[name: "*/${env.BRANCH_NAME}"]],
+ extensions: [ cloneOption(shallow: true) ],
+ userRemoteConfigs: [
+ [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822',
+ url: 'ssh://code.g-r-l.com:6611/generalresearch']
+ ],
+ )
+ stash name: 'source', useDefaultExcludes: false
+ }
+ }
+
+ stage('Python Versions') {
matrix {
axes {
axis {
- name 'PYTHON_VERSION'
- values 'python3.14' 'python3.13', 'python3.12', 'python3.11', 'python3.10'
+ name 'VER'
+ values 'python3.14', 'python3.13', 'python3.12'
}
}
stages {
- stage('Setup DB') {
- script {
- env.REDIS_DB = new Random().nextInt(1024).toString()
- env.REDIS = "${env.REDIS}:6379/${env.REDIS_DB}"
- env.THL_REDIS = "${env.THL_REDIS}:6379/${env.REDIS_DB}"
- echo "Using THL Redis: ${env.REDIS}"
- if (sh(script: "redis-cli -u ${env.REDIS} SET jenkins_lock 1 NX EX 3600", returnStdout: true).trim() != 'OK')
- error('Redis already locked... aborting.')
- }
- script {
- env.GR_REDIS_DB = new Random().nextInt(1024).toString()
- env.GR_REDIS = "redis://${env.REDIS}:6379/${env.GR_REDIS_DB}"
- echo "Using GR Redis: ${env.GR_REDIS}"
- if (sh(script: "redis-cli -u ${env.GR_REDIS} SET jenkins_lock 1 NX EX 3600", returnStdout: true).trim() != 'OK')
- error('Redis already locked... aborting.')
- }
- }
- }
-
- stage('Setup Git') {
+ stage('Setup') {
steps {
- cleanWs()
-
- dir('tmp') {
- sh 'pwd -P'
- }
-
- dir("generalresearch:$PYTHON_VERSION/") {
- checkout scmGit(
- branches: [[name: env.BRANCH_NAME]],
- extensions: [ cloneOption(shallow: true) ],
- userRemoteConfigs: [
- [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822',
- url: 'ssh://code.g-r-l.com:6611/generalresearch']
- ],
- )
- }
-
- dir("thl-spectrum:$PYTHON_VERSION/") {
- checkout scmGit(
- branches: [[name: env.BRANCH_NAME]],
- extensions: [ cloneOption(shallow: true) ],
- userRemoteConfigs: [
- [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822',
- url: 'ssh://code.g-r-l.com:6611/thl-marketplaces/thl-spectrum']
- ],
- )
- }
-
- dir("grliq:$PYTHON_VERSION/") {
- checkout scmGit(
- branches: [[name: env.BRANCH_NAME]],
- extensions: [ cloneOption(shallow: true) ],
- userRemoteConfigs: [
- [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822',
- url: 'ssh://code.g-r-l.com:6611/grl-iq']
- ],
- )
- }
-
- dir("gr:$PYTHON_VERSION/") {
- checkout scmGit(
- branches: [[name: env.BRANCH_NAME]],
- extensions: [ cloneOption(shallow: true) ],
- userRemoteConfigs: [
- [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822',
- url: 'ssh://code.g-r-l.com:6611/general-research/gr-carer']
- ],
- )
- }
- }
- }
-
- stage('Env & Migration') {
- steps {
- dir("generalresearch:$PYTHON_VERSION/") {
- sh "/usr/local/bin/$PYTHON_VERSION -m venv $VENV-$PYTHON_VERSION"
- sh "$VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip"
- sh "$VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt"
- sh "$VENV-$PYTHON_VERSION/bin/pip install '.[django]'"
- sh """
- export DB_NAME=${DB_NAME}
- export DB_USER=${env.DB_USER}
- export DB_PASSWORD=${env.DB_PASSWORD}
- export DB_HOST=${env.DB_POSTGRESQL_HOST}
- $VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION -m generalresearch.thl_django.app.manage migrate
- """
- }
-
- dir("thl-spectrum:$PYTHON_VERSION/") {
- dir('carer') {
- sh "/usr/local/bin/$PYTHON_VERSION -m venv $SPECTRUM_CARER_VENV-$PYTHON_VERSION"
- sh "$SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip"
- sh "$SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt"
-
- sh """
- export DB_NAME=${SPECTRUM_DB_NAME}
- $SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=carer.settings.unittest
- """
+ dir("generalresearch-${VER}") {
+ deleteDir()
+ unstash 'source'
+
+ withCredentials([file(
+ credentialsId: '971e1f48-09ce-4446-9155-a52c1adb6249',
+ variable: 'ENV_TEST_FILE')]) {
+ sh 'cp $ENV_TEST_FILE .env.test'
}
- }
-
- dir("grliq:$PYTHON_VERSION/") {
- dir('carer') {
- sh "/usr/local/bin/$PYTHON_VERSION -m venv $GRLIQ_CARER_VENV-$PYTHON_VERSION"
- sh "$GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip"
- sh "$GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt"
-
- sh """
- export DB_NAME=${GRLIQ_DB_NAME}
- $GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=carer.settings.unittest
- """
- }
- }
-
- dir("gr:$PYTHON_VERSION/") {
- sh "/usr/local/bin/$PYTHON_VERSION -m venv $GR_CARER_VENV-$PYTHON_VERSION"
- sh "$GR_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip"
- sh "$GR_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt"
-
- sh """
- export DB_NAME=${GR_DB_NAME}
- $GR_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=gr.settings.unittest
- """
+ sh "/usr/local/bin/${VER} -m venv ${VENV}-${VER}"
+ sh "${VENV}-${VER}/bin/pip install -U setuptools wheel pip"
+ sh "${VENV}-${VER}/bin/pip install '.'"
+ sh "${VENV}-${VER}/bin/pip install '.[django,dask]'"
}
}
}
stage('base') {
- when {
- expression { return true }
- }
steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/sql_helper.py"
+ dir("generalresearch-${VER}") {
+ sh "${VENV}-${VER}/bin/pytest tests/test_postgres.py -vs"
}
}
}
stage('models') {
- when {
- expression { return true }
- }
steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/models"
+ dir("generalresearch-${VER}") {
+ sh "${VENV}-${VER}/bin/pytest tests/models/gr/test_base.py -vs"
}
}
}
stage('managers') {
steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/managers"
- }
- }
- }
-
- stage('wall_status_codes') {
- steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/wall_status_codes"
- }
- }
- }
-
- stage('wxet') {
- steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/wxet"
- }
- }
- }
-
- stage('grliq') {
- steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/grliq"
+ dir("generalresearch-${VER}") {
+ sh "${VENV}-${VER}/bin/pytest tests/managers/gr/ -vs"
}
}
}
- stage('incite') {
- steps {
- dir("generalresearch:$PYTHON_VERSION") {
- sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/incite"
- }
- }
- }
}
}
}
}
+
post {
always {
echo 'One way or another, I have finished'
- deleteDir() /* clean up our workspace */
- sh """
- mariadb -h ${env.DB_MARIA_HOST} -u ${env.DB_USER} -p${env.DB_PASSWORD} --ssl=0 -e 'DROP DATABASE `${env.SPECTRUM_DB_NAME}`;'
- """
- sh """
- PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres <<EOF
- DROP DATABASE "${env.DB_NAME}";
- EOF
- """
- sh """
- PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres <<EOF
- DROP DATABASE "${env.GRLIQ_DB_NAME}";
- EOF
- """
- sh """
- PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres <<EOF
- DROP DATABASE "${env.GR_DB_NAME}";
- EOF
- """
-
- sh "redis-cli -u ${env.THL_REDIS} FLUSHDB"
- sh "redis-cli -u ${env.GR_REDIS} FLUSHDB"
+ deleteDir()
}
}
-}
+} \ No newline at end of file
diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py
index 2100d41..f27385c 100644
--- a/generalresearch/__init__.py
+++ b/generalresearch/__init__.py
@@ -1,19 +1,27 @@
+import logging
import threading
import time
+from collections.abc import Callable
from functools import wraps
-from typing import Any, Callable, Optional
+from typing import ParamSpec, TypeVar
from decorator import decorator
from wrapt import FunctionWrapper, ObjectProxy
+P = ParamSpec("P")
+R = TypeVar("R")
+
+ExceptionType = type[BaseException]
+ExceptionsArg = ExceptionType | tuple[ExceptionType, ...]
+
def retry(
- exceptions,
+ exceptions: ExceptionsArg,
tries: int = 4,
delay: float = 0.5,
backoff: int = 2,
- logger: Optional[Any] = None,
-) -> Callable:
+ logger: logging.Logger | None = None,
+) -> Callable[[Callable[P, R]], Callable[P, R]]:
"""
https://www.calazan.com/retry-decorator-for-python-3/
Retry calling the decorated function using an exponential backoff.
@@ -28,10 +36,10 @@ def retry(
logger: Logger to use. If None, print.
"""
- def deco_retry(f):
+ def deco_retry(f: Callable[P, R]) -> Callable[P, R]:
@wraps(f)
- def f_retry(*args, **kwargs):
+ def f_retry(*args: P.args, **kwargs: P.kwargs) -> R:
mtries, mdelay = tries, delay
while mtries > 1:
try:
@@ -128,7 +136,7 @@ def synchronized(wrapped):
if lock is None:
lock = threading.RLock()
- setattr(context, "_synchronized_lock", lock)
+ context._synchronized_lock = lock
return lock
diff --git a/generalresearch/config.py b/generalresearch/config.py
index 44f3db7..75b9901 100644
--- a/generalresearch/config.py
+++ b/generalresearch/config.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import os
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from pathlib import Path
from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn
@@ -16,9 +16,9 @@ def is_debug() -> bool:
import os
is_developer: bool = os.getenv("USER") in {"nanis", "gstupp"}
- is_pytest1: bool = bool(os.getenv("PYTEST_TEST", False))
- is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST", False))
- is_pytest3: bool = bool(os.getenv("PYTEST_VERSION", False))
+ is_pytest1: bool = bool(os.getenv("PYTEST_TEST"))
+ is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST"))
+ is_pytest3: bool = bool(os.getenv("PYTEST_VERSION"))
is_debugging1: bool = os.getenv("DEBUG", "").lower() in ("1", "true", "yes")
is_debugging2: bool = os.getenv("PYTHON_DEBUG", "").lower() in ("1", "true", "yes")
is_jenkins: bool = bool(os.getenv("JENKINS_HOME")) or bool(os.getenv("JENKINS_URL"))
@@ -53,6 +53,8 @@ class GRLBaseSettings(BaseSettings):
testing_postgres_user: str | None = Field(default=None)
testing_postgres_pass: str | None = Field(default=None)
+ testing_redis: InternalHostname | None = Field(default=None)
+
git_creds: str | None = Field(default=None)
# ---
@@ -112,12 +114,11 @@ class GRLBaseSettings(BaseSettings):
tango_customer_id: str | None = Field(default=None)
# --- Keeping this here as we use these ids regardless of the AMT account
- amt_bonus_cashout_method_id: str | None = Field(default=None)
- amt_assignment_cashout_method_id: str | None = Field(default=None)
+ amt_bonus_cashout_method_id: str | None = Field(default="1951a47541fb46519827b8783e2a53ab")
+ amt_assignment_cashout_method_id: str | None = Field(default="5b23e4df3e2c40609ca8edf40b13237f")
- # --- Maxmind Configuration ---
- maxmind_account_id: str | None = Field(default=None)
- maxmind_license_key: str | None = Field(default=None)
+ # --- GRIP Configuration ---
+ grip_token: str | None = Field(default=None)
EXAMPLE_PRODUCT_ID = "1108d053e4fa47c5b0dbdcd03a7981e7"
@@ -125,4 +126,4 @@ EXAMPLE_PRODUCT_ID = "1108d053e4fa47c5b0dbdcd03a7981e7"
# AMT accounting was changed many times and txs before this date
# are either missing AMT bonuses, or not accounting for hit rewards.
JAMES_BILLINGS_BPID = "888dbc589987425fa846d6e2a8daed04"
-JAMES_BILLINGS_TX_CUTOFF = datetime(2026, 1, 1, tzinfo=timezone.utc)
+JAMES_BILLINGS_TX_CUTOFF = datetime(2026, 1, 1, tzinfo=UTC)
diff --git a/generalresearch/currency.py b/generalresearch/currency.py
index c76b266..7a9d037 100644
--- a/generalresearch/currency.py
+++ b/generalresearch/currency.py
@@ -1,6 +1,6 @@
import warnings
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from typing import Any
from pydantic import GetCoreSchemaHandler, NonNegativeInt
@@ -9,7 +9,7 @@ from pydantic_core import CoreSchema, core_schema
from generalresearch.utils.enum import ReprEnumMeta
-class LedgerCurrency(str, Enum, metaclass=ReprEnumMeta):
+class LedgerCurrency(StrEnum, metaclass=ReprEnumMeta):
USD = "USD"
USDCent = "USDCent"
USDMill = "USDMill"
@@ -25,16 +25,16 @@ def format_usd_cent(usd_cent: int) -> str:
class USDCent(int):
- def __new__(cls, value, *args, **kwargs):
+ def __new__(cls, value: int, *args, **kwargs):
if isinstance(value, float):
warnings.warn(
- "USDCent init with a float. Rounding behavior may " "be unexpected"
+ "USDCent init with a float. Rounding behavior may be unexpected"
)
if isinstance(value, Decimal):
warnings.warn(
- "USDCent init with a Decimal. Rounding behavior may " "be unexpected"
+ "USDCent init with a Decimal. Rounding behavior may be unexpected"
)
if value < 0:
@@ -42,17 +42,17 @@ class USDCent(int):
return super(cls, cls).__new__(cls, value)
- def __add__(self, other):
+ def __add__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__add__(other)
return self.__class__(res)
- def __sub__(self, other):
+ def __sub__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__sub__(other)
return self.__class__(res)
- def __mul__(self, other):
+ def __mul__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__mul__(other)
return self.__class__(res)
@@ -61,14 +61,14 @@ class USDCent(int):
res = super().__abs__()
return self.__class__(res)
- def __truediv__(self, other):
+ def __truediv__(self, value):
raise ValueError("Division not allowed for USDCent")
def __str__(self):
- return "%d" % int(self)
+ return f"{int(self):d}"
def __repr__(self):
- return "USDCent(%d)" % int(self)
+ return f"USDCent({int(self)})"
@classmethod
def __get_pydantic_core_schema__(
@@ -97,12 +97,12 @@ class USDMill(int):
if isinstance(value, float):
warnings.warn(
- "USDMill init with a float. Rounding behavior " "may be unexpected"
+ "USDMill init with a float. Rounding behavior may be unexpected"
)
if isinstance(value, Decimal):
warnings.warn(
- "USDMill init with a Decimal. Rounding behavior " "may be unexpected"
+ "USDMill init with a Decimal. Rounding behavior may be unexpected"
)
if value < 0:
@@ -110,17 +110,17 @@ class USDMill(int):
return super(cls, cls).__new__(cls, value)
- def __add__(self, other):
+ def __add__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__add__(other)
return self.__class__(res)
- def __sub__(self, other):
+ def __sub__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__sub__(other)
return self.__class__(res)
- def __mul__(self, other):
+ def __mul__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__mul__(other)
return self.__class__(res)
@@ -129,14 +129,14 @@ class USDMill(int):
res = super().__abs__()
return self.__class__(res)
- def __truediv__(self, other):
+ def __truediv__(self, value):
raise ValueError("Division not allowed for USDMill")
def __str__(self):
- return "%d" % int(self)
+ return f"{int(self):d}"
def __repr__(self):
- return "USDMill(%d)" % int(self)
+ return f"USDMill({int(self)})"
@classmethod
def __get_pydantic_core_schema__(
diff --git a/generalresearch/grliq/managers/__init__.py b/generalresearch/grliq/managers/__init__.py
index 849b6c2..e69de29 100644
--- a/generalresearch/grliq/managers/__init__.py
+++ b/generalresearch/grliq/managers/__init__.py
@@ -1,34 +0,0 @@
-from generalresearch.grliq.models.forensic_data import GrlIqData
-from generalresearch.grliq.models.forensic_result import (
- GrlIqCheckerResults,
- GrlIqForensicCategoryResult,
-)
-
-DUMMY_GRLIQ_DATA = [
- {
- "data": GrlIqData.model_validate_json(
- """{"mid": "3722ed29314940fabd37b42d808dcf5a", "uuid": "b11441da5a854dfbb8401d4c32e56db5", "phase": "offerwall-enter", "events": null, "vendor": "Google Inc.", "app_name": "Netscape", "calendar": "gregory", "language": "en-US", "platform": "Linux x86_64", "timezone": "America/Mexico_City", "client_ip": "131.196.250.250", "timestamp": "2025-02-27T16:05:34-06:00", "webrtc_ip": "131.196.250.250", "created_at": "2025-02-27T22:05:35.370589Z", "language_2": "en-US", "language_3": null, "platform_2": "Linux x86_64", "platform_3": null, "prefetched": true, "product_id": "d0606a0b5d034a8d81b1e3579d1f76fd", "webgl_flag": true, "webgl_hash": "da27e1b9b660057a3f5e185d3f5deabe", "canvas_hash": "14ed764326ec454d976c322261d99f16", "color_gamut": "3", "country_iso": "mx", "inner_width": 612, "outer_width": 1813, "product_sub": "20030107", "audio_codecs": "1,1,1,1,1,3,1,3,1,3,3,1,1,3,3,3,3,1,3,3,3,2,1,1", "cookie_check": "", "graphics_api": "WebKit WebGL", "inner_height": 1174, "mouse_events": null, "ontouchstart": false, "outer_height": 1261, "plugins_hash": "4c05fa2f766a444d4f253ead792c8b0e|2", "screen_width": 2560, "video_codecs": "1,3,3,3,3,3,3,3,3,3,1,1,1,1,1,1,3,1,1,1,3,3,1", "webgl_hash_2": "fc73fd5db75e2c36222fe34251be3971", "webrtc_error": false, "window_opera": false, "battery_level": 0.9, "canvas_hash_2": "bd11ebbf5c26fd20e0217820b4159752", "dynamic_range": false, "error_message": "Cannot read", "forced_colors": false, "math_result_1": "1.9275814160560204e-50", "math_result_2": "1.6182817135715877", "screen_height": 1440, "webgl_check_1": true, "webgl_context": "webgl2", "window_chrome": true, "connection_rtt": 150, "history_length": 16, "user_agent_str": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "web_sql_exists": false, "calender_locale": "en-US", "connection_type": "", "inverted_colors": true, "navigator_brave": false, "product_user_id": "d1d55df1-959e-4740-b77c-fa1f4fc457ae", "request_headers": {"host": "test", "accept": "*/*", "connection": "keep-alive", "user-agent": "python-httpx/0.27.0", "content-length": "3646", "accept-encoding": "gzip, deflate", "x-forwarded-for": "131.196.250.250"}, "timezone_offset": 360, "webrtc_local_ip": "50486637-6b64-4812-b10a-0a75337c31bd.local", "battery_charging": true, "client_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "131.196.250.250", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "mx", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "max_touch_points": 0, "numbering_system": "latn", "path_fingerprint": 3252, "prefers_contrast": "0", "rendering_engine": "WebKit", "timezone_success": "pass", "user_agent_hints": {"model": null, "brands": [{"brand": "Google Chrome", "version": "131"}, {"brand": "Chromium", "version": "131"}, {"brand": "Not_A Brand", "version": "24"}], "mobile": false, "bitness": "64", "platform": "Linux", "brands_full": [{"brand": "Google Chrome", "version": "131.0.6778.204"}, {"brand": "Chromium", "version": "131.0.6778.204"}, {"brand": "Not_A Brand", "version": "24.0.0.0"}], "architecture": "x86", "platform_version": "6.2.0"}, "user_agent_str_2": null, "webgl_extensions": "EXT_clip_control|EXT_color_buffer_float|EXT_color_buffer_half_float|EXT_conservative_depth|EXT_depth_clamp|EXT_disjoint_timer_query_webgl2|EXT_float_blend|EXT_polygon_offset_clamp|EXT_render_snorm|EXT_texture_compression_bptc|EXT_texture_compression_rgtc|EXT_texture_filter_anisotropic|EXT_texture_mirror_clamp_to_edge|EXT_texture_norm16|KHR_parallel_shader_compile|NV_shader_noperspective_interpolation|OES_draw_buffers_indexed|OES_sample_variables|OES_shader_multisample_interpolation|OES_texture_float_linear|OVR_multiview2|WEBGL_blend_func_extended|WEBGL_clip_cull_distance|WEBGL_compressed_texture_astc|WEBGL_compressed_texture_etc|WEBGL_compressed_texture_etc1|WEBGL_compressed_texture_s3tc|WEBGL_compressed_texture_s3tc_srgb|WEBGL_debug_renderer_info|WEBGL_debug_shaders|WEBGL_lose_context|WEBGL_multi_draw|WEBGL_polygon_mode|WEBGL_provoking_vertex|WEBGL_stencil_texturing", "webrtc_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "131.196.250.250", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "mx", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "chrome_extensions": "", "execution_time_ms": 371.0999999642372, "graphics_renderer": "WebGL 2.0 (OpenGL ES 3.0 Chromium)", "keyboard_detected": true, "mime_types_length": 2, "request_fs_exists": true, "audio_context_flag": "pass", "audio_context_hash": "9307303774dec3248c18a939392090da", "canvas_fingerprint": 258, "canvas_pixel_check": false, "device_pixel_ratio": 1.0, "indexedDbData_blob": true, "navigator_keys_len": 79, "no_edge_pdf_plugin": false, "screen_avail_width": 2560, "webdriver_detected": false, "window_orientation": 0, "connection_downlink": 10.0, "navigator_webdriver": false, "non_native_function": false, "screen_avail_height": 1400, "supported_fonts_str": "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0|2147483648|109543424|1880099872|268435471", "text_2d_fingerprint": "bfcce91c9e71d11af7b14dbee4c75f83", "webrtc_is_supported": "pass", "canvas_support_level": "full", "do_not_track_enabled": "1", "hardware_concurrency": 12, "keyboard_layout_size": 48, "prefers_color_scheme": false, "webgl_max_anisotropy": 16, "battery_charging_time": 0.0, "browser_by_properties": "c", "eval_to_string_length": 33, "performance_loop_time": 0.09999996423721313, "session_storage_check": "pass", "unmasked_vendor_webgl": "Google Inc. (Intel)", "hardware_concurrency_2": 12, "hardware_concurrency_3": null, "localStorage_available": true, "memory_jsHeapSizeLimit": 4294705152, "mozilla_web_app_exists": false, "navigator_deviceMemory": 8.0, "navigator_java_enabled": false, "prefers_reduced_motion": false, "storage_estimate_quota": 1178717110272, "webdriver_detected_msg": "", "window_active_x_object": false, "window_external_exists": true, "color_depth_pixel_depth": "24-24", "indexedDbData_available": true, "navigator_cookieEnabled": true, "unmasked_renderer_webgl": "ANGLE (Intel, Mesa Intel(R) Graphics (RPL-P), OpenGL 4.6)", "battery_discharging_time": 0.0, "connection_effectiveType": "4g", "non_native_function_flag": "", "speech_synthesis_voice_1": "Google Bahasa Indonesia", "window_client_information": true, "audio_compressor_reduction": 20.538288116455078, "navigator_mediaDevices_len": 3, "audio_intensity_fingerprint": 124.04347527516074, "speech_synthesis_voice_hash": "8010ee3313813de521e48e63bd5a6f13", "microsoft_credentials_exists": false, "window_installTrigger_exists": false, "speech_synthesis_voices_count": 19, "webgl_shading_language_version": "WebGL GLSL ES 3.00 (OpenGL ES GLSL ES 3.0 Chromium)", "error_message_stack_access_count": 0, "speech_synthesis_avail_voices_count": 19, "error_message_stack_access_count_worker": 0}"""
- ),
- "result_data": GrlIqCheckerResults.model_validate_json(
- """{"uuid": "b11441da5a854dfbb8401d4c32e56db5", "check_codecs": {"score": 0}, "check_timezone": {"score": 0}, "check_timestamp": {"score": 0}, "check_user_type": {"score": 0}, "check_ip_changes": {"score": 0}, "check_ip_country": {"score": 0}, "check_environment": {"score": 0}, "check_ip_timezone": {"score": 0}, "check_isp_changes": {"score": 0}, "check_useragent_js": {"score": 0}, "check_required_fonts": {"score": 0}, "check_user_anonymous": {"score": 0}, "check_webrtc_success": {"score": 0}, "check_seen_timestamps": {"msg": "duplicate timestamp", "score": 100}, "check_country_timezone": {"score": 0}, "check_prohibited_fonts": {"score": 0}, "check_timezone_changes": {"score": 0}, "check_execution_time_ms": {"msg": "duplicate execution_time_ms", "score": 100}, "check_fingerprint_reuse": {"score": 0}, "check_fingerprint_cycling": {"score": 0}, "check_ip_webrtc_ip_detail": {"score": 0}, "check_environment_critical": {"score": 0}, "check_useragent_other_enums": {"score": 0}, "check_useragent_ip_properties": {"score": 0}, "check_useragent_data_properties": {"score": 0}, "check_useragent_device_family_brand": {"score": 0}}"""
- ),
- "category_result": GrlIqForensicCategoryResult.model_validate_json(
- """{"uuid": "b11441da5a854dfbb8401d4c32e56db5", "is_bot": 0, "is_tampered": 100, "is_velocity": 0, "is_anonymous": 0, "suspicious_ip": 0, "is_oscillating": 0, "is_teleporting": 0, "is_inconsistent": 0, "platform_ip_inconsistent": 0}"""
- ),
- "fraud_score": 100,
- "is_attempt_allowed": False,
- },
- {
- "data": GrlIqData.model_validate_json(
- """{"mid": "35f6f5c30bc74ea7ac4aca7b40a02352", "uuid": "d54509f2f310499f8ab74839b10b2a41", "phase": "offerwall-enter", "events": null, "vendor": "Google Inc.", "app_name": "Netscape", "calendar": "gregory", "language": "en-US", "platform": "Linux x86_64", "timezone": "America/Los_Angeles", "client_ip": "104.9.125.144", "timestamp": "2025-02-28T11:34:39-08:00", "webrtc_ip": "172.56.209.195", "created_at": "2025-02-28T19:34:39.681872Z", "language_2": "en-US", "language_3": null, "platform_2": "Linux x86_64", "platform_3": null, "prefetched": true, "product_id": "d0606a0b5d034a8d81b1e3579d1f76fd", "webgl_flag": true, "webgl_hash": "da27e1b9b660057a3f5e185d3f5deabe", "canvas_hash": "e6e4d17da26050ce85ad00d3c6ea999e", "color_gamut": "3", "country_iso": "us", "inner_width": 841, "outer_width": 1680, "product_sub": "20030107", "audio_codecs": "1,1,1,1,1,3,1,3,1,3,3,1,1,3,3,3,3,1,3,3,3,2,1,1", "cookie_check": "", "graphics_api": "WebKit WebGL", "inner_height": 891, "mouse_events": null, "ontouchstart": false, "outer_height": 978, "plugins_hash": "4c05fa2f766a444d4f253ead792c8b0e|2", "screen_width": 1680, "video_codecs": "1,3,3,3,3,3,3,3,3,3,1,1,1,1,1,1,3,1,1,1,3,3,1", "webgl_hash_2": "fc73fd5db75e2c36222fe34251be3971", "webrtc_error": false, "window_opera": false, "battery_level": 0.41, "canvas_hash_2": "e0559d49b1864985cafc0d1c3a6b053c", "dynamic_range": false, "error_message": "Cannot read", "forced_colors": false, "math_result_1": "1.9275814160560204e-50", "math_result_2": "1.6182817135715877", "screen_height": 1050, "webgl_check_1": true, "webgl_context": "webgl2", "window_chrome": true, "connection_rtt": 100, "history_length": 11, "user_agent_str": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "web_sql_exists": false, "calender_locale": "en-US", "connection_type": "", "inverted_colors": true, "navigator_brave": false, "product_user_id": "test-unit", "request_headers": {"dnt": "1", "host": "127.0.0.1:8081", "accept": "application/json, lk/null q=0.1", "origin": "http://127.0.0.1:8080", "referer": "http://127.0.0.1:8080/", "sec-ch-ua": "\\"Google Chrome\\";v=\\"131\\", \\"Chromium\\";v=\\"131\\", \\"Not_A Brand\\";v=\\"24\\"", "connection": "keep-alive", "user-agent": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "content-type": "application/json", "content-length": "3313", "sec-fetch-dest": "empty", "sec-fetch-mode": "cors", "sec-fetch-site": "same-site", "accept-encoding": "gzip, deflate, br, zstd", "accept-language": "en-US,en;q=0.9", "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": "\\"Linux\\""}, "timezone_offset": 480, "webrtc_local_ip": "10.253.217.45,[2607:fb91:20c5:c6af:cda0:10b4:830a:a85e]", "battery_charging": false, "client_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "104.9.125.144", "isp": "AT&T Internet", "latitude": 37.3897, "city_name": "Mountain View", "longitude": -122.083, "time_zone": "America/Los_Angeles", "user_type": "residential", "country_iso": "us", "postal_code": "94041", "is_anonymous": false, "accuracy_radius": 5, "static_ip_score": 40.3, "subdivision_1_iso": "CA", "subdivision_2_iso": null, "subdivision_1_name": "California", "subdivision_2_name": null, "registered_country_iso": "us"}, "max_touch_points": 0, "numbering_system": "latn", "path_fingerprint": 3252, "prefers_contrast": "0", "rendering_engine": "WebKit", "timezone_success": "pass", "user_agent_hints": {"model": null, "brands": [{"brand": "Google Chrome", "version": "131"}, {"brand": "Chromium", "version": "131"}, {"brand": "Not_A Brand", "version": "24"}], "mobile": false, "bitness": "64", "platform": "Linux", "brands_full": [{"brand": "Google Chrome", "version": "131.0.6778.204"}, {"brand": "Chromium", "version": "131.0.6778.204"}, {"brand": "Not_A Brand", "version": "24.0.0.0"}], "architecture": "x86", "platform_version": "6.2.0"}, "user_agent_str_2": null, "webgl_extensions": "EXT_clip_control|EXT_color_buffer_float|EXT_color_buffer_half_float|EXT_conservative_depth|EXT_depth_clamp|EXT_disjoint_timer_query_webgl2|EXT_float_blend|EXT_polygon_offset_clamp|EXT_render_snorm|EXT_texture_compression_bptc|EXT_texture_compression_rgtc|EXT_texture_filter_anisotropic|EXT_texture_mirror_clamp_to_edge|EXT_texture_norm16|KHR_parallel_shader_compile|NV_shader_noperspective_interpolation|OES_draw_buffers_indexed|OES_sample_variables|OES_shader_multisample_interpolation|OES_texture_float_linear|OVR_multiview2|WEBGL_blend_func_extended|WEBGL_clip_cull_distance|WEBGL_compressed_texture_astc|WEBGL_compressed_texture_etc|WEBGL_compressed_texture_etc1|WEBGL_compressed_texture_s3tc|WEBGL_compressed_texture_s3tc_srgb|WEBGL_debug_renderer_info|WEBGL_debug_shaders|WEBGL_lose_context|WEBGL_multi_draw|WEBGL_polygon_mode|WEBGL_provoking_vertex|WEBGL_stencil_texturing", "webrtc_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "172.56.209.195", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "us", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "chrome_extensions": "", "execution_time_ms": 924.5, "graphics_renderer": "WebGL 2.0 (OpenGL ES 3.0 Chromium)", "keyboard_detected": true, "mime_types_length": 2, "request_fs_exists": true, "audio_context_flag": "pass", "audio_context_hash": "9307303774dec3248c18a939392090da", "canvas_fingerprint": 258, "canvas_pixel_check": false, "device_pixel_ratio": 1.0, "indexedDbData_blob": true, "navigator_keys_len": 79, "no_edge_pdf_plugin": false, "screen_avail_width": 1680, "webdriver_detected": false, "window_orientation": 0, "connection_downlink": 10.0, "navigator_webdriver": false, "non_native_function": false, "screen_avail_height": 1010, "supported_fonts_str": "72|17152|327680|1073741824|0|0|540736|73728|7340032|1342177280|117446657|256|16|0|262687|4290797636|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0|2147483648|109543680|1880099888|301989903", "text_2d_fingerprint": "bfcce91c9e71d11af7b14dbee4c75f83", "webrtc_is_supported": "pass", "canvas_support_level": "full", "do_not_track_enabled": "1", "hardware_concurrency": 12, "keyboard_layout_size": 48, "prefers_color_scheme": false, "webgl_max_anisotropy": 16, "battery_charging_time": 0.0, "browser_by_properties": "c", "eval_to_string_length": 33, "performance_loop_time": 0.09999999962747097, "session_storage_check": "pass", "unmasked_vendor_webgl": "Google Inc. (Intel)", "hardware_concurrency_2": 12, "hardware_concurrency_3": null, "localStorage_available": true, "memory_jsHeapSizeLimit": 4294705152, "mozilla_web_app_exists": false, "navigator_deviceMemory": 8.0, "navigator_java_enabled": false, "prefers_reduced_motion": false, "storage_estimate_quota": 1178717110272, "webdriver_detected_msg": "", "window_active_x_object": false, "window_external_exists": true, "color_depth_pixel_depth": "24-24", "indexedDbData_available": true, "navigator_cookieEnabled": true, "unmasked_renderer_webgl": "ANGLE (Intel, Mesa Intel(R) Graphics (RPL-P), OpenGL 4.6)", "battery_discharging_time": 4844.0, "connection_effectiveType": "4g", "non_native_function_flag": "", "speech_synthesis_voice_1": "Google Bahasa Indonesia", "window_client_information": true, "audio_compressor_reduction": 20.538288116455078, "navigator_mediaDevices_len": 8, "audio_intensity_fingerprint": 124.04347527516074, "speech_synthesis_voice_hash": "8010ee3313813de521e48e63bd5a6f13", "microsoft_credentials_exists": false, "window_installTrigger_exists": false, "speech_synthesis_voices_count": 19, "webgl_shading_language_version": "WebGL GLSL ES 3.00 (OpenGL ES GLSL ES 3.0 Chromium)", "error_message_stack_access_count": 2, "speech_synthesis_avail_voices_count": 19, "error_message_stack_access_count_worker": 2}"""
- ),
- "result_data": GrlIqCheckerResults.model_validate_json(
- """{"uuid": "d54509f2f310499f8ab74839b10b2a41", "check_codecs": {"score": 0}, "check_timezone": {"score": 0}, "check_timestamp": {"score": 0}, "check_user_type": {"score": 0}, "check_ip_changes": {"score": 0}, "check_ip_country": {"score": 0}, "check_environment": {"msg": "error_message_stack_access_count: 2", "score": 100}, "check_ip_timezone": {"score": 0}, "check_isp_changes": {"score": 0}, "check_useragent_js": {"score": 0}, "check_required_fonts": {"score": 0}, "check_user_anonymous": {"score": 0}, "check_webrtc_success": {"score": 0}, "check_seen_timestamps": {"score": 0}, "check_country_timezone": {"score": 0}, "check_prohibited_fonts": {"score": 0}, "check_timezone_changes": {"score": 0}, "check_execution_time_ms": {"score": 0}, "check_fingerprint_reuse": {"score": 0}, "check_fingerprint_cycling": {"score": 0}, "check_ip_webrtc_ip_detail": {"score": 0}, "check_environment_critical": {"score": 0}, "check_useragent_other_enums": {"score": 0}, "check_useragent_ip_properties": {"score": 0}, "check_useragent_data_properties": {"score": 0}, "check_useragent_device_family_brand": {"score": 0}}"""
- ),
- "category_result": GrlIqForensicCategoryResult.model_validate_json(
- """{"uuid": "d54509f2f310499f8ab74839b10b2a41", "is_bot": 0, "is_tampered": 0, "is_velocity": 0, "is_anonymous": 0, "suspicious_ip": 0, "is_oscillating": 0, "is_teleporting": 0, "is_inconsistent": 10, "platform_ip_inconsistent": 0}"""
- ),
- "fraud_score": 10,
- "is_attempt_allowed": True,
- },
-]
diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py
index 54105ce..61cc52c 100644
--- a/generalresearch/grliq/managers/event_plotter.py
+++ b/generalresearch/grliq/managers/event_plotter.py
@@ -1,20 +1,22 @@
import html
import webbrowser
-from typing import List
+from typing import TYPE_CHECKING
import numpy as np
from more_itertools import windowed
from scipy.spatial.distance import euclidean
from generalresearch.grliq.managers.colormap import turbo_colormap_data
-from generalresearch.grliq.models.events import KeyboardEvent, MouseEvent
+
+if TYPE_CHECKING:
+ from generalresearch.grliq.models.events import KeyboardEvent, MouseEvent
def make_events_svg(
- mouse_events: List[MouseEvent], keyboard_events: List[KeyboardEvent]
+ mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent]
) -> str:
if len(mouse_events) + len(keyboard_events) == 0:
- return f'<svg xmlns="http://www.w3.org/2000/svg">\n' + "\n</svg>"
+ return '<svg xmlns="http://www.w3.org/2000/svg">\n' + "\n</svg>"
t = np.array([pm.timeStamp for pm in mouse_events])
t_diff = t.max() - t.min()
@@ -89,7 +91,7 @@ def make_events_svg(
svg_elements.append(svg_multiline_text(text, cx + 5, cy - 5, font_size))
svg = (
- f'<svg xmlns="http://www.w3.org/2000/svg">'
+ '<svg xmlns="http://www.w3.org/2000/svg">'
+ "\n".join(svg_elements)
+ "\n</svg>"
)
@@ -119,8 +121,8 @@ def svg_multiline_text(
def group_input_events_by_xy(
- mouse_events: List[MouseEvent], keyboard_events: List[KeyboardEvent]
-) -> List[tuple[tuple[float, float], List[str]]]:
+ mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent]
+) -> list[tuple[tuple[float, float], list[str]]]:
"""
Each keypress is its own event. For plotting, we want to group together
all keypresses that were made when the mouse was at the same position,
diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py
index 739c520..0810723 100644
--- a/generalresearch/grliq/managers/forensic_data.py
+++ b/generalresearch/grliq/managers/forensic_data.py
@@ -1,7 +1,8 @@
from __future__ import annotations
+from collections.abc import Collection
from datetime import datetime
-from typing import Any, Collection
+from typing import TYPE_CHECKING, Any
from psycopg import sql
from pydantic import NonNegativeInt, PositiveInt
@@ -14,12 +15,13 @@ from generalresearch.grliq.models.forensic_result import (
Phase,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.user import User
-from generalresearch.pg_helper import PostgresConfig
+if TYPE_CHECKING:
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
-class GrlIqDataManager:
+class GrlIqDataManager:
def __init__(self, postgres_config: PostgresConfig):
self.postgres_config = postgres_config
@@ -32,7 +34,15 @@ class GrlIqDataManager:
is_attempt_allowed: bool | None = None,
) -> GrlIqData:
- data = iq_data.model_dump_sql(exclude={"events", "mouse_events", "timing_data"})
+ data = iq_data.model_dump_sql(
+ exclude={
+ "events",
+ "mouse_events",
+ "timing_data",
+ "results",
+ "category_result",
+ }
+ )
data["result_data"] = None
if result_data:
@@ -102,14 +112,13 @@ class GrlIqDataManager:
is_attempt_allowed = %(is_attempt_allowed)s
WHERE uuid = %(uuid)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def update_fingerprint(self, iq_data: GrlIqData) -> None:
# We should only run this if we modified the fingerprint algorithm
@@ -122,14 +131,13 @@ class GrlIqDataManager:
SET fingerprint = %(fingerprint)s
WHERE uuid = %(uuid)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def update_data(self, iq_data: GrlIqData) -> None:
# We should only run this if we structured new fields and want to
@@ -140,14 +148,13 @@ class GrlIqDataManager:
SET data = %(data)s
WHERE id = %(id)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def get_data_if_exists(
self, forensic_uuid: UUIDStr, load_events: bool = False
@@ -471,20 +478,20 @@ class GrlIqDataManager:
filters.append("d.fingerprint = ANY(%(fingerprints)s)")
if product_ids and len(product_ids) == 1:
- product_id = list(product_ids)[0]
+ product_id = next(iter(product_ids))
product_ids = None
if product_ids:
- assert (
- users is None and user is None and product_id is None
- ), "user, users, product_id, and product_ids are mutually exclusive"
+ assert users is None and user is None and product_id is None, (
+ "user, users, product_id, and product_ids are mutually exclusive"
+ )
params["product_ids"] = list(set(product_ids))
filters.append("d.product_id = ANY(%(product_ids)s::UUID[])")
if product_id:
- assert (
- users is None and user is None and product_ids is None
- ), "user, users, product_id, and product_ids are mutually exclusive"
+ assert users is None and user is None and product_ids is None, (
+ "user, users, product_id, and product_ids are mutually exclusive"
+ )
params["product_id"] = product_id
filters.append("d.product_id = %(product_id)s")
@@ -505,12 +512,12 @@ class GrlIqDataManager:
)
if created_between:
- assert (
- created_after is None
- ), "Cannot pass both created_after and created_between"
- assert (
- created_before is None
- ), "Cannot pass both created_before and created_between"
+ assert created_after is None, (
+ "Cannot pass both created_after and created_between"
+ )
+ assert created_before is None, (
+ "Cannot pass both created_before and created_between"
+ )
params["created_after"] = created_between[0]
params["created_before"] = created_between[1]
filters.append(
@@ -518,9 +525,9 @@ class GrlIqDataManager:
)
if user:
- assert (
- product_ids is None and users is None
- ), "user, users, and product_ids are mutually exclusive"
+ assert product_ids is None and users is None, (
+ "user, users, and product_ids are mutually exclusive"
+ )
params["product_id"] = user.product_id
params["product_user_id"] = user.product_user_id
filters.append(
@@ -528,16 +535,16 @@ class GrlIqDataManager:
)
if users:
- assert (
- product_ids is None and user is None
- ), "user, users, and product_ids are mutually exclusive"
+ assert product_ids is None and user is None, (
+ "user, users, and product_ids are mutually exclusive"
+ )
user_args = ", ".join(
[f"(%(bp_{i})s, %(bpuid_{i})s)" for i in range(len(users))]
)
filters.append(f"(d.product_id, d.product_user_id) IN ({user_args})")
- for i, user in enumerate(users):
- params[f"bp_{i}"] = user.product_id
- params[f"bpuid_{i}"] = user.product_user_id
+ for i, _user in enumerate(users):
+ params[f"bp_{i}"] = _user.product_id
+ params[f"bpuid_{i}"] = _user.product_user_id
if phase:
params["phase"] = phase.value
@@ -594,24 +601,19 @@ class GrlIqDataManager:
)
if only_product_id:
- try:
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(
- query="""
- SELECT count AS c
- FROM grliq_forensicdata_product_counts
- WHERE product_id = %s
- LIMIT 1
- """,
- params=(product_id,),
- )
- res = c.fetchone()
- if res and res["c"] >= 0:
- return int(res["c"])
-
- except (Exception,) as e:
- pass
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ query="""
+ SELECT count AS c
+ FROM grliq_forensicdata_product_counts
+ WHERE product_id = %s
+ LIMIT 1
+ """,
+ params=(product_id,),
+ )
+ res = c.fetchone()
+ if res and res["c"] >= 0:
+ return int(res["c"])
query = f"""
SELECT COUNT(1) AS c
@@ -653,9 +655,9 @@ class GrlIqDataManager:
if product_ids:
# It doesn't use the (product_id, created_at) index with multiple product_ids
- assert (
- offset == 0
- ), "Cannot paginate using product_ids, use product_id instead"
+ assert offset == 0, (
+ "Cannot paginate using product_ids, use product_id instead"
+ )
filter_str, params = self.make_filter_str(
session_uuid=session_uuid,
@@ -686,7 +688,6 @@ class GrlIqDataManager:
res: list[dict[str, Any]] = c.fetchall() # type: ignore
for x in res:
-
if "data" in x:
self.temporary_add_missing_fields(x["data"])
x["data"]["id"] = x["id"]
diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py
index bbc6b6d..fc1ae3e 100644
--- a/generalresearch/grliq/managers/forensic_events.py
+++ b/generalresearch/grliq/managers/forensic_events.py
@@ -1,6 +1,7 @@
import json
+from collections.abc import Collection
from datetime import datetime
-from typing import Any, Collection, Dict, List, Optional
+from typing import TYPE_CHECKING, Any
from uuid import uuid4
from psycopg import sql
@@ -14,7 +15,10 @@ from generalresearch.grliq.models.events import (
TimingData,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+
+ from generalresearch.pg_helper import PostgresConfig
class GrlIqEventManager:
@@ -25,7 +29,7 @@ class GrlIqEventManager:
def update_or_create_timing(
self,
session_uuid: UUIDStr,
- timing_data: Optional[TimingData] = None,
+ timing_data: TimingData | None = None,
) -> PositiveInt:
data = {
"session_uuid": session_uuid,
@@ -35,40 +39,35 @@ class GrlIqEventManager:
"uuid": uuid4().hex,
}
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
- # Try to update first
- update_query = sql.SQL(
- """
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
+ # Try to update first
+ update_query = sql.SQL("""
UPDATE grliq_forensicevents
SET timing_data = %(timing_data)s
WHERE session_uuid = %(session_uuid)s
AND timing_data IS NULL
RETURNING id
- """
- )
- c.execute(update_query, data)
- result = c.fetchone()
+ """)
+ c.execute(update_query, data)
+ result = c.fetchone()
- if result:
- pk = result["id"]
- conn.commit()
- return pk
+ if result:
+ pk = result["id"]
+ conn.commit()
+ return pk
- # No matching row to update. Do an insert
- insert_query = sql.SQL(
- """
+ # No matching row to update. Do an insert
+ insert_query = sql.SQL("""
INSERT INTO grliq_forensicevents
(uuid, session_uuid, timing_data)
VALUES
(%(uuid)s, %(session_uuid)s, %(timing_data)s)
RETURNING id
- """
- )
- c.execute(insert_query, data)
- pk = c.fetchone()["id"]
- conn.commit()
+ """)
+ c.execute(insert_query, data)
+ pk = c.fetchone()["id"]
+ conn.commit()
return int(pk)
@@ -77,8 +76,8 @@ class GrlIqEventManager:
session_uuid: UUIDStr,
event_start: datetime,
event_end: datetime,
- events: Optional[List[Dict]] = None,
- mouse_events: Optional[List[Dict]] = None,
+ events: list[dict] | None = None,
+ mouse_events: list[dict] | None = None,
) -> PositiveInt:
data = {
"uuid": uuid4().hex,
@@ -91,12 +90,10 @@ class GrlIqEventManager:
"event_end": event_end,
}
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
- # Try to update first
- update_query = sql.SQL(
- """
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
+ # Try to update first
+ update_query = sql.SQL("""
UPDATE grliq_forensicevents
SET events = %(events)s,
mouse_events = %(mouse_events)s,
@@ -105,19 +102,17 @@ class GrlIqEventManager:
WHERE session_uuid = %(session_uuid)s
AND events IS NULL
RETURNING id
- """
- )
- c.execute(update_query, data)
- result = c.fetchone()
+ """)
+ c.execute(update_query, data)
+ result = c.fetchone()
- if result:
- pk = result["id"]
- conn.commit()
- return pk
+ if result:
+ pk = result["id"]
+ conn.commit()
+ return pk
- # No matching row to update. Do an insert
- insert_query = sql.SQL(
- """
+ # No matching row to update. Do an insert
+ insert_query = sql.SQL("""
INSERT INTO grliq_forensicevents
(uuid, session_uuid, events, mouse_events,
event_start, event_end)
@@ -125,24 +120,23 @@ class GrlIqEventManager:
(%(uuid)s, %(session_uuid)s, %(events)s, %(mouse_events)s,
%(event_start)s, %(event_end)s)
RETURNING id
- """
- )
- c.execute(insert_query, data)
- pk = c.fetchone()["id"]
- conn.commit()
+ """)
+ c.execute(insert_query, data)
+ pk = c.fetchone()["id"]
+ conn.commit()
return int(pk)
def filter(
self,
- select_str: Optional[str] = None,
- session_uuid: Optional[str] = None,
- session_uuids: Optional[Collection[str]] = None,
- uuids: Optional[Collection[str]] = None,
- started_since: Optional[datetime] = None,
- limit: Optional[int] = None,
+ select_str: str | None = None,
+ session_uuid: str | None = None,
+ session_uuids: Collection[str] | None = None,
+ uuids: Collection[str] | None = None,
+ started_since: datetime | None = None,
+ limit: int | None = None,
order_by: str = "event_start DESC",
- ) -> List[Dict[str, Any]]:
+ ) -> list[dict[str, Any]]:
if not limit:
limit = 100
@@ -174,10 +168,9 @@ class GrlIqEventManager:
{filter_str}
ORDER BY {order_by} LIMIT {limit}
"""
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query=query, params=params)
- res = c.fetchall()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query=query, params=params)
+ res = c.fetchall()
for x in res:
if x.get("mouse_events"):
@@ -199,10 +192,9 @@ class GrlIqEventManager:
def filter_distinct_timing(
self,
session_uuids: Collection[str],
- ) -> List[Dict[str, Any]]:
+ ) -> list[dict[str, Any]]:
params = {"session_uuids": list(session_uuids)}
- query = sql.SQL(
- """
+ query = sql.SQL("""
SELECT DISTINCT ON (fe.session_uuid)
timing_data,
fe.session_uuid,
@@ -213,12 +205,10 @@ class GrlIqEventManager:
WHERE fe.session_uuid = ANY(%(session_uuids)s)
AND timing_data IS NOT NULL
ORDER BY session_uuid, fe.id DESC;
- """
- )
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
- res = c.fetchall()
+ """)
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, params)
+ res = c.fetchall()
for x in res:
x["timing_data"] = TimingData.model_validate(x["timing_data"])
@@ -229,7 +219,7 @@ class GrlIqEventManager:
return res
@staticmethod
- def process_mouse_events(pointer_moves: List[PointerMove], events: List[Dict]):
+ def process_mouse_events(pointer_moves: list[PointerMove], events: list[dict]):
"""
In the db column 'mouse_events' we put all 'pointermove' events. Pull
those out, and then any 'pointerdown' and 'pointerup' events from the
@@ -274,7 +264,7 @@ class GrlIqEventManager:
return mouse_events
@staticmethod
- def process_keyboard_events(events: List[Dict]):
+ def process_keyboard_events(events: list[dict]):
res = [
KeyboardEvent(
type=x["type"],
diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py
index 30db53d..158e582 100644
--- a/generalresearch/grliq/managers/forensic_results.py
+++ b/generalresearch/grliq/managers/forensic_results.py
@@ -1,5 +1,6 @@
+from collections.abc import Collection
from datetime import datetime
-from typing import Any, Collection, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.grliq.models.forensic_result import (
GrlIqForensicCategoryResult,
@@ -16,16 +17,16 @@ class GrlIqCategoryResultsReader:
def filter_category_results(
self,
- session_uuid: Optional[str] = None,
- fingerprint: Optional[str] = None,
- phase: Optional[Phase] = None,
- uuids: Optional[Collection[str]] = None,
- product_ids: Optional[Collection[str]] = None,
- created_since: Optional[datetime] = None,
- created_between: Optional[Tuple[datetime, datetime]] = None,
- user: Optional[User] = None,
- limit: Optional[int] = None,
- ) -> List[Dict[str, Any]]:
+ session_uuid: str | None = None,
+ fingerprint: str | None = None,
+ phase: Phase | None = None,
+ uuids: Collection[str] | None = None,
+ product_ids: Collection[str] | None = None,
+ created_since: datetime | None = None,
+ created_between: tuple[datetime, datetime] | None = None,
+ user: User | None = None,
+ limit: int | None = None,
+ ) -> list[dict[str, Any]]:
"""
For retrieving GrlIqForensicCategoryResult objects from db.
@@ -89,10 +90,9 @@ class GrlIqCategoryResultsReader:
ORDER BY created_at DESC LIMIT {limit}
"""
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
- res = c.fetchall()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, params)
+ res = c.fetchall()
for x in res:
x["client_ip"] = str(x["client_ip"])
diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py
index 21b7e4b..c222075 100644
--- a/generalresearch/grliq/managers/forensic_summary.py
+++ b/generalresearch/grliq/managers/forensic_summary.py
@@ -2,15 +2,11 @@ from __future__ import annotations
import statistics
from collections import defaultdict
-from datetime import datetime, timedelta, timezone
-from typing import Any, Dict, List
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING, Any
import numpy as np
-from generalresearch.grliq.managers.forensic_data import GrlIqDataManager
-from generalresearch.grliq.managers.forensic_events import (
- GrlIqEventManager,
-)
from generalresearch.grliq.models.forensic_result import (
GrlIqCheckerResults,
GrlIqForensicCategoryResult,
@@ -22,12 +18,18 @@ from generalresearch.grliq.models.forensic_summary import (
TimingDataCountrySummary,
UserForensicSummary,
)
-from generalresearch.models.thl.user import User
-from generalresearch.redis_helper import RedisConfig
+
+if TYPE_CHECKING:
+ from generalresearch.grliq.managers.forensic_data import GrlIqDataManager
+ from generalresearch.grliq.managers.forensic_events import (
+ GrlIqEventManager,
+ )
+ from generalresearch.models.thl.user import User
+ from generalresearch.redis_helper import RedisConfig
def calculate_category_summary(
- res: List[GrlIqForensicCategoryResult],
+ res: list[GrlIqForensicCategoryResult],
) -> GrlIqForensicCategorySummary:
totals = defaultdict(int)
is_complete_count = 0
@@ -55,7 +57,7 @@ def calculate_category_summary(
def calculate_checker_summary(
- res: List[GrlIqCheckerResults],
+ res: list[GrlIqCheckerResults],
) -> GrlIqCheckerResultsSummary:
totals = defaultdict(list)
none_totals = defaultdict(int)
@@ -85,8 +87,8 @@ def calculate_checker_summary(
def calculate_timing_summary(
- redis_config: RedisConfig, timing_res: List[Dict[str, Any]]
-) -> Dict[str, TimingDataCountrySummary]:
+ redis_config: RedisConfig, timing_res: list[dict[str, Any]]
+) -> dict[str, TimingDataCountrySummary]:
country_median_rtts = defaultdict(list)
for x in timing_res:
@@ -109,7 +111,7 @@ def calculate_timing_summary(
for k, v in country_distributions.items()
}
- out = dict()
+ out = {}
for country_iso, median_rtts in country_median_rtts.items():
country_stats = country_distributions[country_iso]
z_scores = [
@@ -137,7 +139,7 @@ def run_user_forensic_summary(
user: User,
) -> UserForensicSummary:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
created_between = (now - timedelta(days=90), now)
select_str = "id, session_uuid, product_id, product_user_id, created_at, result_data, category_result"
res = iq_dm.filter(
@@ -158,12 +160,14 @@ def run_user_forensic_summary(
)
session_uuids = {x["session_uuid"] for x in res}
- timing_res: List[Dict] = iq_em.filter_distinct_timing(session_uuids=session_uuids)
+ timing_res: list[dict[str, Any]] = iq_em.filter_distinct_timing(
+ session_uuids=session_uuids
+ )
country_timing_data_summary = (
calculate_timing_summary(redis_config=redis_config, timing_res=timing_res)
if timing_res
- else dict()
+ else {}
)
s = UserForensicSummary(
diff --git a/generalresearch/grliq/models/__init__.py b/generalresearch/grliq/models/__init__.py
index 998de2d..fabbca2 100644
--- a/generalresearch/grliq/models/__init__.py
+++ b/generalresearch/grliq/models/__init__.py
@@ -1,10 +1,10 @@
from __future__ import annotations
import json
-from enum import Enum
+from enum import StrEnum
-class RiskWeighting(str, Enum):
+class RiskWeighting(StrEnum):
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
diff --git a/generalresearch/grliq/models/custom_types.py b/generalresearch/grliq/models/custom_types.py
index c5eb93c..903e230 100644
--- a/generalresearch/grliq/models/custom_types.py
+++ b/generalresearch/grliq/models/custom_types.py
@@ -1,5 +1,6 @@
+from typing import Annotated
+
import annotated_types
-from typing_extensions import Annotated
GrlIqScore = Annotated[int, annotated_types.Ge(0), annotated_types.Le(100)]
GrlIqAvgScore = Annotated[float, annotated_types.Ge(0), annotated_types.Le(100)]
diff --git a/generalresearch/grliq/models/decider.py b/generalresearch/grliq/models/decider.py
index 4464a7f..579c39f 100644
--- a/generalresearch/grliq/models/decider.py
+++ b/generalresearch/grliq/models/decider.py
@@ -1,14 +1,14 @@
from __future__ import annotations
-from datetime import datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field
from generalresearch.models.custom_types import AwareDatetimeISO
-class Decider(str, Enum):
+class Decider(StrEnum):
# This decision was made in the thl-core: pre-offerwall-entry view
PRE_ENTRY = "pre_entry"
# This decision made by grl-iq (synchronously)
@@ -17,7 +17,7 @@ class Decider(str, Enum):
YM_USER = "ym_user"
-class AttemptDecision(str, Enum):
+class AttemptDecision(StrEnum):
# This attempt should be allowed to continue
PASS = "pass"
# This attempt is deemed fraudulent
@@ -35,7 +35,7 @@ class GrlIqAttemptResult(BaseModel):
timestamp: AwareDatetimeISO = Field(
description="When this decision was made",
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
)
decider: Decider = Field(description="Where this decision was made")
decision: AttemptDecision = Field(
diff --git a/generalresearch/grliq/models/events.py b/generalresearch/grliq/models/events.py
index 7e9ebee..995a6fa 100644
--- a/generalresearch/grliq/models/events.py
+++ b/generalresearch/grliq/models/events.py
@@ -3,7 +3,7 @@ from __future__ import annotations
from collections import namedtuple
from dataclasses import dataclass, fields
from functools import cached_property
-from typing import Any, Dict, List, Optional
+from typing import Any, Self
import numpy as np
from pydantic import (
@@ -14,7 +14,6 @@ from pydantic import (
NonNegativeInt,
PositiveFloat,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr
@@ -37,14 +36,14 @@ class Event:
# in microseconds, since page load (?)
timeStamp: float
# optional ID of the event target (e.g.: where the mouse is hovering)
- _elementId: Optional[str] = None
+ _elementId: str | None = None
# optional tag name of the event target
- _elementTagName: Optional[str] = None
+ _elementTagName: str | None = None
# extracted coordinates for the element being interacted with
- _elementBounds: Optional[Bounds] = None
+ _elementBounds: Bounds | None = None
@classmethod
- def from_dict(cls, data: Dict[str, Any]) -> Self:
+ def from_dict(cls, data: dict[str, Any]) -> Self:
data = {k: v for k, v in data.items() if k in cls.__dataclass_fields__}
bounds = data.get("_elementBounds")
if bounds is not None and not isinstance(bounds, Bounds):
@@ -104,13 +103,13 @@ class KeyboardEvent(Event):
# "insertText", "insertCompositionText", "deleteCompositionText",
# "insertFromComposition", "deleteContentBackward"
- inputType: Optional[str]
+ inputType: str | None
# e.g., 'Enter', 'a', 'Backspace'
- key: Optional[str] = None
+ key: str | None = None
# This is the actual text, if applicable
- data: Optional[str] = None
+ data: str | None = None
@property
def key_text(self):
@@ -159,18 +158,18 @@ class TimingData(BaseModel):
"""
model_config = ConfigDict(extra="forbid", validate_assignment=True)
- client_rtts: List[float] = Field()
- server_rtts: List[float] = Field()
+ client_rtts: list[float] = Field()
+ server_rtts: list[float] = Field()
# Have to be optional for backwards-compatibility, but should always be set.
- started_at: Optional[AwareDatetimeISO] = Field(default=None)
- ended_at: Optional[AwareDatetimeISO] = Field(default=None)
- client_ip: Optional[IPvAnyAddressStr] = Field(
+ started_at: AwareDatetimeISO | None = Field(default=None)
+ ended_at: AwareDatetimeISO | None = Field(default=None)
+ client_ip: IPvAnyAddressStr | None = Field(
description="This comes from the websocket request's headers",
examples=["72.39.217.116"],
default=None,
)
- server_hostname: Optional[str] = Field(
+ server_hostname: str | None = Field(
description="The hostname of the server that handled this request",
examples=["grliq-web-0"],
default=None,
@@ -178,18 +177,13 @@ class TimingData(BaseModel):
@property
def server_location(self) -> str:
- # TODO: when we have more locations ...
- return (
- "fremont_ca"
- if self.server_hostname in {"grliq-web-0", "grliq-web-1"}
- else "fremont_ca"
- )
+ return "fremont_ca"
@property
def has_data(self):
return len(self.client_rtts) > 0 and len(self.server_rtts) > 0
- def filter_rtts(self, rtts: List[float]) -> List[float]:
+ def filter_rtts(self, rtts: list[float]) -> list[float]:
# Skip the first 5 pings, unless we have <10 pings, then get the last
# 5 instead.
# The first couple pings are usually outliers as they are running
@@ -234,7 +228,7 @@ class TimingData(BaseModel):
return rtts
@property
- def summarize(self) -> Optional[TimingDataSummary]:
+ def summarize(self) -> TimingDataSummary | None:
if len(self.filtered_rtts) < 5:
return None
diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py
index f8bdd98..9d69e41 100644
--- a/generalresearch/grliq/models/forensic_data.py
+++ b/generalresearch/grliq/models/forensic_data.py
@@ -3,10 +3,10 @@ from __future__ import annotations
import hashlib
import re
from collections import Counter
-from datetime import datetime, timedelta, timezone
-from enum import Enum
+from datetime import UTC, datetime, timedelta
+from enum import StrEnum
from functools import cached_property
-from typing import Any, Literal
+from typing import TYPE_CHECKING, Annotated, Any, Literal, Self
from uuid import uuid4
import pycountry
@@ -23,7 +23,6 @@ from pydantic import (
)
from pydantic.json_schema import SkipJsonSchema
from pydantic_extra_types.timezone_name import TimeZoneName
-from typing_extensions import Annotated, Self
from generalresearch.grliq.models import (
AUDIO_CODEC_NAMES,
@@ -55,12 +54,14 @@ from generalresearch.models.custom_types import (
UUIDStr,
)
from generalresearch.models.thl.ipinfo import GeoIPInformation
-from generalresearch.models.thl.session import Session
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.session import Session
fake = Faker()
-class Platform(str, Enum):
+class Platform(StrEnum):
MAC_INTEL = "MacIntel"
ARM = "ARM"
IPAD = "iPad"
@@ -75,7 +76,7 @@ class Platform(str, Enum):
OTHER = "Other"
-class PassFailError(str, Enum):
+class PassFailError(StrEnum):
PASS = "pass"
FAIL = "fail"
ERROR = "error"
@@ -96,7 +97,7 @@ class PassFailError(str, Enum):
return {2: cls.PASS, 1: cls.FAIL, 0: cls.ERROR, -1: cls.ERROR}[int(v)]
-class SupportLevel(str, Enum):
+class SupportLevel(StrEnum):
# Used for checking if certain features are available in the browser
FULL = "full"
PARTIAL = "partial"
@@ -475,9 +476,11 @@ class GrlIqData(BaseModel):
description="Bit-packed string for font support. Each element is 32 bits, with each bit representing T/F for "
"font support.",
examples=[
- "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636"
- "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0"
- "|2147483648|109543424|1880099872|268435471"
+ (
+ "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636"
+ "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0"
+ "|2147483648|109543424|1880099872|268435471"
+ )
],
)
@@ -559,19 +562,21 @@ class GrlIqData(BaseModel):
@cached_property
def audio_codecs_named(self) -> dict[str, bool]:
+ assert self.audio_codecs
return dict(
zip(
AUDIO_CODEC_NAMES,
- [True if x == "3" else False for x in self.audio_codecs.split(",")],
+ [x == "3" for x in self.audio_codecs.split(",")],
)
)
@cached_property
def video_codecs_named(self) -> dict[str, bool]:
+ assert self.video_codecs
return dict(
zip(
VIDEO_CODEC_NAMES,
- [True if x == "3" else False for x in self.video_codecs.split(",")],
+ [x == "3" for x in self.video_codecs.split(",")],
)
)
@@ -771,19 +776,17 @@ class GrlIqData(BaseModel):
# product_id and product_user_id are parsed from the post body. make sure
# they match the session whose mid was specified
assert self.product_id == session.user.product_id, "product_id mismatch"
- assert (
- self.product_user_id == session.user.product_user_id
- ), "product_user_id mismatch"
+ assert self.product_user_id == session.user.product_user_id, (
+ "product_user_id mismatch"
+ )
# validate the Session's mid is "recent"
- assert (datetime.now(tz=timezone.utc) - session.started) < timedelta(
- minutes=90
- ), "expired session"
-
- return None
+ assert (datetime.now(tz=UTC) - session.started) < timedelta(minutes=90), (
+ "expired session"
+ )
def model_dump_sql(self, **kwargs) -> dict[str, Any]:
- d = dict()
+ d = {}
d["uuid"] = self.uuid
d["session_uuid"] = self.mid
d["created_at"] = self.created_at
@@ -803,7 +806,7 @@ class GrlIqData(BaseModel):
return d
@classmethod
- def from_db(cls, d: dict[str, Any]) -> Self:
+ def from_db(cls, d: dict[str, Any]) -> GrlIqData:
res = GrlIqData.model_validate(d["data"])
if d.get("category_result"):
diff --git a/generalresearch/grliq/models/forensic_result.py b/generalresearch/grliq/models/forensic_result.py
index 3fe6481..1fd99cc 100644
--- a/generalresearch/grliq/models/forensic_result.py
+++ b/generalresearch/grliq/models/forensic_result.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from uuid import uuid4
from pydantic import BaseModel, ConfigDict, Field, computed_field
@@ -14,7 +14,7 @@ from generalresearch.grliq.models.decider import (
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-class Phase(str, Enum):
+class Phase(StrEnum):
# The 'phase' of a THL-Session experience. grliq may be collected in
# multiple places multiple times within one session
diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py
index d6f46f8..b3bbd55 100644
--- a/generalresearch/grliq/models/forensic_summary.py
+++ b/generalresearch/grliq/models/forensic_summary.py
@@ -2,16 +2,15 @@ from __future__ import annotations
import random
from typing import (
- List,
Literal,
Union,
get_args,
get_origin,
get_type_hints,
- Optional,
)
import numpy as np
+from grip_client.enums import AccessType
from pydantic import (
BaseModel,
ConfigDict,
@@ -29,7 +28,6 @@ from generalresearch.grliq.models.forensic_result import (
)
from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr
from generalresearch.models.thl.locales import CountryISO
-from generalresearch.models.thl.maxmind.definitions import UserType
example_rtt_percentiles = (
[133.332]
@@ -57,9 +55,7 @@ class UserForensicSummary(BaseModel):
)
# These must be nullable in case a user has 0 attempts!
- category_result_summary: GrlIqForensicCategorySummary | None= Field(
- default=None
- )
+ category_result_summary: GrlIqForensicCategorySummary | None = Field(default=None)
checker_result_summary: GrlIqCheckerResultsSummary | None = Field(default=None)
country_timing_data_summary: dict[CountryISO, TimingDataCountrySummary] = Field(
@@ -128,7 +124,7 @@ def generate_GrlIqCheckerResultsSummary():
if base_type == GrlIqCheckerResult:
if is_opt:
fields[f"{field_name}_avg"] = (
- Optional[GrlIqAvgScore],
+ GrlIqAvgScore | None,
Field(default=None, examples=[random.randint(0, 100)]),
)
fields[f"{field_name}_pct_none"] = (
@@ -189,7 +185,9 @@ class IPTimingDataSummary(BaseModel):
client_ip: IPvAnyAddressStr = Field(examples=["123.123.123.123"])
country_iso: CountryISO = Field(examples=["us"])
server_location: Literal["fremont_ca"] = Field(default="fremont_ca")
- user_type: UserType | None = Field(default=None, examples=[UserType.RESIDENTIAL])
+ user_type: AccessType | None = Field(
+ default=None, examples=[AccessType.RESIDENTIAL]
+ )
expected_rtt_range: tuple[float, float] = Field(
description="The expected rtt range for this IP (based on country_iso/user_type) to server_location",
examples=[(45.193, 120.841)],
@@ -211,16 +209,16 @@ class CountryRTTDistribution(BaseModel):
description="Country client_ip is located in", examples=["fr"]
)
# For users marked as fraud or not
- is_fraud: bool| None = Field(
+ is_fraud: bool | None = Field(
default=None,
description="If timing data from sessions determined to be fraud are included",
)
# we could split by this optionally
- user_type: UserType|None = Field(
+ user_type: AccessType | None = Field(
default=None,
description="user_type of the client_ip as determined by MaxMind",
- examples=[UserType.RESIDENTIAL],
+ examples=[AccessType.RESIDENTIAL],
)
rtt_min: float = Field(gt=0, examples=[133.332])
@@ -228,7 +226,7 @@ class CountryRTTDistribution(BaseModel):
rtt_mean: float = Field(gt=0, examples=[179.302])
rtt_max: float = Field(gt=0, examples=[890.006])
rtt_std: float = Field(gt=0, examples=[46.831])
- rtt_percentiles: List[float] = Field(
+ rtt_percentiles: list[float] = Field(
min_length=101, max_length=101, examples=[example_rtt_percentiles]
)
@@ -255,10 +253,9 @@ class CountryRTTDistribution(BaseModel):
Render a boxplot from the RTT percentiles.
"""
try:
- # annoying pycharm error
import matplotlib.pyplot as plt
- except ImportError as e:
- raise e
+ except ImportError:
+ return
p = self.rtt_percentiles
data = {
@@ -270,7 +267,7 @@ class CountryRTTDistribution(BaseModel):
"fliers": [p[0]] + ([p[100]] if p[100] > p[95] else []),
}
- fig, ax = plt.subplots(figsize=(4, 1.5))
+ _, ax = plt.subplots(figsize=(4, 1.5))
ax.bxp([data], showfliers=True, vert=False)
ax.set_title(f"RTT Boxplot for {self.country_iso}")
ax.set_xlabel("RTT (ms)")
diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py
index 1953f6d..3d5e5ce 100644
--- a/generalresearch/grliq/models/useragents.py
+++ b/generalresearch/grliq/models/useragents.py
@@ -1,15 +1,15 @@
from __future__ import annotations
import hashlib
-from enum import Enum
+from enum import StrEnum
+from typing import Self
from pydantic import BaseModel, ConfigDict, Field, field_validator
-from typing_extensions import Self
from user_agents import parse as ua_parse
from user_agents.parsers import UserAgent
-class BrowserFamily(str, Enum):
+class BrowserFamily(StrEnum):
CHROME_MOBILE = "Chrome Mobile"
CHROME = "Chrome"
CHROME_MOBILE_WEBVIEW = "Chrome Mobile WebView"
@@ -32,7 +32,7 @@ class BrowserFamily(str, Enum):
OTHER = "Other"
-class OSFamily(str, Enum):
+class OSFamily(StrEnum):
ANDROID = "Android"
WINDOWS = "Windows"
IOS = "iOS"
@@ -43,7 +43,7 @@ class OSFamily(str, Enum):
OTHER = "Other"
-class DeviceBrand(str, Enum):
+class DeviceBrand(StrEnum):
GENERIC_ANDROID = "Generic_Android"
NONE = "None"
APPLE = "Apple"
@@ -65,7 +65,7 @@ class DeviceBrand(str, Enum):
OTHER = "Other"
-class DeviceModelFamily(str, Enum):
+class DeviceModelFamily(StrEnum):
NONE = "None"
OTHER = "Other"
K = "K"
@@ -132,7 +132,7 @@ class BrowserInfo(BaseModel):
class DeviceInfo(BaseModel):
family: DeviceModelFamily = Field()
- brand: DeviceBrand = Field()
+ brand: DeviceBrand | None = Field(default=None)
model: DeviceModelFamily = Field()
@field_validator("family", "model", mode="before")
@@ -186,7 +186,7 @@ class GrlUserAgent(BaseModel):
def ua_string_values(self) -> dict[str, str]:
# Returns the raw parsed string values for each of these. To be used
# for db filtering, identifying trends, etc.
- d = dict()
+ d = {}
d["ua_browser_family"] = self.ua_parsed.browser.family
d["ua_browser_version"] = self.ua_parsed.browser.version_string
d["ua_os_family"] = self.ua_parsed.os.family
diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py
index 95390a8..711e562 100644
--- a/generalresearch/grliq/utils.py
+++ b/generalresearch/grliq/utils.py
@@ -1,11 +1,10 @@
from __future__ import annotations
import os
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from pathlib import Path
from uuid import UUID
-# from generalresearch.config import
from generalresearch.models.custom_types import UUIDStr
@@ -16,7 +15,7 @@ def get_screenshot_fp(
grliq_ss_dir_name: str = "canvas2html",
create_dir_if_not_exists: bool = True,
) -> Path | None:
- assert created_at.tzinfo == timezone.utc
+ assert created_at.tzinfo == UTC
if isinstance(forensic_uuid, UUID):
forensic_uuid = forensic_uuid.hex
diff --git a/generalresearch/grpc.py b/generalresearch/grpc.py
index 040fd26..f1b5611 100644
--- a/generalresearch/grpc.py
+++ b/generalresearch/grpc.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from google.protobuf.duration_pb2 import Duration
from google.protobuf.timestamp_pb2 import Timestamp
@@ -20,13 +20,13 @@ def timestamp_from_datetime_nullable(dt: datetime | None) -> Timestamp:
def timestamp_to_datetime(ts: Timestamp) -> datetime:
- return datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc)
+ return datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=UTC)
def timestamp_to_datetime_nullable(ts: Timestamp) -> datetime | None:
# grpc has no None. If a google.protobuf.Timestamp field is not set, it gets interpreted as timestamp 0
- default = datetime.fromtimestamp(0, tz=timezone.utc)
- d = datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc)
+ default = datetime.fromtimestamp(0, tz=UTC)
+ d = datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=UTC)
return None if d == default else d
diff --git a/generalresearch/healing_ppe.py b/generalresearch/healing_ppe.py
index 254a893..dfee689 100644
--- a/generalresearch/healing_ppe.py
+++ b/generalresearch/healing_ppe.py
@@ -73,7 +73,7 @@ def test():
time.sleep(0.5)
# Kill a process in the pool
- pid = list(pool._processes.keys())[0]
+ pid = next(iter(pool._processes.keys()))
os.kill(pid, signal.SIGKILL)
time.sleep(0.5)
diff --git a/generalresearch/incite/__init__.py b/generalresearch/incite/__init__.py
index e69de29..8b60e4b 100644
--- a/generalresearch/incite/__init__.py
+++ b/generalresearch/incite/__init__.py
@@ -0,0 +1,4 @@
+import logging
+
+logging.basicConfig()
+LOG = logging.getLogger(f"{__name__}.incite")
diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py
index 54d1565..d83886b 100644
--- a/generalresearch/incite/base.py
+++ b/generalresearch/incite/base.py
@@ -1,14 +1,13 @@
from __future__ import annotations
import glob
-import logging
import os
import re
import shutil
import subprocess
import warnings
from collections.abc import Callable, Sequence
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from os import R_OK, access, listdir
from os.path import isdir
from os.path import join as pjoin
@@ -17,11 +16,13 @@ from sys import platform
from typing import (
TYPE_CHECKING,
Any,
+ Self,
)
from uuid import uuid4
import dask.dataframe as dd
import pandas as pd
+import pandera as pa
import pyarrow.parquet as pq
from distributed import Client as DaskClient
from pandera.pandas import DataFrameSchema
@@ -34,15 +35,14 @@ from pydantic import (
PositiveInt,
PrivateAttr,
TypeAdapter,
- ValidationInfo,
field_validator,
model_validator,
)
from pydantic.json_schema import SkipJsonSchema
from sentry_sdk import capture_exception
-from typing_extensions import Self
from generalresearch.config import is_debug
+from generalresearch.incite import LOG
from generalresearch.incite.schemas import (
ARCHIVE_AFTER,
empty_dataframe_from_schema,
@@ -51,15 +51,11 @@ from generalresearch.models.custom_types import AwareDatetimeISO
if TYPE_CHECKING:
from generalresearch.incite.collections import DFCollection, DFCollectionItem
- from generalresearch.incite.collections.thl_marketplaces import (
- DFCollectionType,
- )
- from generalresearch.incite.mergers import MergeCollection, MergeType
+ from generalresearch.incite.collections.thl_marketplaces import DFCollectionType
+ from generalresearch.incite.mergers.base import MergeCollection, MergeType
Collection = DFCollection | MergeCollection
-logging.basicConfig()
-LOG = logging.getLogger()
# Item = Union["DFCollectionItem", "MergeCollectionItem"]
Item = Any
@@ -67,7 +63,6 @@ Items = Sequence[Item]
DT_STR = "%Y-%m-%d %H:%M:%S"
_dir_adapter = TypeAdapter(DirectoryPath)
-_filepath_adapter = TypeAdapter(FilePath)
class NFSMount(BaseModel):
@@ -95,7 +90,7 @@ class GRLDatasets(BaseModel):
from generalresearch.incite.collections.thl_marketplaces import (
DFCollectionType,
)
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
assert self.data_src, "data src must be defined"
@@ -128,9 +123,10 @@ class GRLDatasets(BaseModel):
type..
"""
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
folder = "mergers" if isinstance(enum_type, MergeType) else "raw/df-collections"
+ assert self.incite is not None
return Path(
pjoin(self.data_src, self.incite.point, folder, str(enum_type.value))
)
@@ -163,7 +159,7 @@ class CollectionBase(BaseModel):
offset: str = Field(default="72h", max_length=5)
start: AwareDatetimeISO = Field(
- default=datetime(year=2018, month=1, day=1, tzinfo=timezone.utc),
+ default=datetime(year=2018, month=1, day=1, tzinfo=UTC),
description="This is the starting point in which data will be retrieved"
"in chunks from.",
frozen=True,
@@ -201,11 +197,8 @@ class CollectionBase(BaseModel):
@model_validator(mode="after")
def check_model_after(self) -> Self:
- if self.offset is None or self.start is None:
- return self
-
offset_total_sec = pd.Timedelta(self.offset).total_seconds()
- start_total_sec = (datetime.now(tz=timezone.utc) - self.start).total_seconds()
+ start_total_sec = (datetime.now(tz=UTC) - self.start).total_seconds()
if offset_total_sec > start_total_sec:
raise ValueError("Offset must be equal to, or smaller the start timestamp")
@@ -213,22 +206,20 @@ class CollectionBase(BaseModel):
return self
@field_validator("start")
- def check_start(
- cls, start: datetime | None, info: ValidationInfo
- ) -> datetime | None:
+ def check_start(cls, start: datetime | None) -> datetime | None:
if start and start.microsecond != 0:
raise ValueError("Collection.start must not have microseconds")
return start
@field_validator("offset")
- def check_offset(cls, v: str | None, info: ValidationInfo):
+ def check_offset(cls, v: str | None):
# pd.offsets.__all__
if v is None:
# In MergeCollections, offset can be None
return v
try:
pd.Timedelta(v)
- except Exception as e:
+ except (ValueError, TypeError) as e:
capture_exception(error=e)
raise ValueError(
"Invalid offset alias provided. Please review: "
@@ -291,14 +282,14 @@ class CollectionBase(BaseModel):
@property
def interval_range(self) -> list[tuple[datetime, datetime]]:
"""closed='left', so 0 <= x < 5"""
- end = self.finished or datetime.now(tz=timezone.utc).replace(microsecond=0)
+ end = self.finished or datetime.now(tz=UTC).replace(microsecond=0)
iv_r = self._interval_range(end)
return [(iv.left.to_pydatetime(), iv.right.to_pydatetime()) for iv in iv_r]
@property
def progress(self) -> pd.DataFrame:
records = [i.to_dict() for i in self.items]
- end = self.finished if self.finished else datetime.now(tz=timezone.utc)
+ end = self.finished if self.finished else datetime.now(tz=UTC)
return pd.DataFrame.from_records(records, index=self._interval_range(end))
@property
@@ -601,15 +592,15 @@ class CollectionBase(BaseModel):
return res
def get_items_from_year(self, year: int) -> Items:
- ts = datetime(year=year, month=1, day=1, tzinfo=timezone.utc)
+ ts = datetime(year=year, month=1, day=1, tzinfo=UTC)
return self.get_items(since=ts)
def get_items_last90(self) -> Items:
- ts = datetime.now(tz=timezone.utc) - timedelta(days=90)
+ ts = datetime.now(tz=UTC) - timedelta(days=90)
return self.get_items(since=ts)
def get_items_last365(self) -> Items:
- ts = datetime.now(tz=timezone.utc) - timedelta(days=365)
+ ts = datetime.now(tz=UTC) - timedelta(days=365)
return self.get_items(since=ts)
@@ -617,7 +608,7 @@ class CollectionItemBase(BaseModel):
# I want to intentionally keep these as native python types, and not
# pandas specific types.
start: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc).replace(microsecond=0)
+ default_factory=lambda: datetime.now(tz=UTC).replace(microsecond=0)
)
# --- Private attrs ---
@@ -807,7 +798,7 @@ class CollectionItemBase(BaseModel):
if archive_after is None:
return False
- return datetime.now(tz=timezone.utc) > self.finish + archive_after
+ return datetime.now(tz=UTC) > self.finish + archive_after
def set_empty(self):
assert self.should_archive(), (
@@ -841,8 +832,8 @@ class CollectionItemBase(BaseModel):
raise ValueError("Unknown path type.")
df = parquet.read().to_pandas()
- except Exception:
- LOG.warning(f"Invalid archive {path=}")
+ except Exception as e:
+ LOG.warning(f"Invalid archive {path=} {e=}")
df = None
# Check if it's None or a totally empty pd.DataFrame before we waste
@@ -860,8 +851,7 @@ class CollectionItemBase(BaseModel):
try:
schema: DataFrameSchema = self._collection._schema
return schema.validate(check_obj=df, lazy=True, sample=sample)
- except Exception as e:
- LOG.exception(e)
+ except pa.errors.SchemaErrors as e:
capture_exception(error=e)
return None
diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py
index 17c119a..e69de29 100644
--- a/generalresearch/incite/collections/__init__.py
+++ b/generalresearch/incite/collections/__init__.py
@@ -1,757 +0,0 @@
-from __future__ import annotations
-
-import logging
-import os
-import subprocess
-import time
-from datetime import datetime
-from enum import Enum
-from sys import platform
-from typing import Any
-
-import dask
-import dask.dataframe as dd
-import pandas as pd
-import pyarrow.parquet as pq
-from dask.distributed import Future
-from distributed import Client, as_completed
-from more_itertools import chunked
-from pandera.pandas import DataFrameSchema
-from psycopg import Cursor
-from pydantic import Field, FilePath, ValidationInfo, field_validator
-from sentry_sdk import capture_exception
-
-from generalresearch.incite.base import CollectionBase, CollectionItemBase
-from generalresearch.incite.schemas import (
- ARCHIVE_AFTER,
- ORDER_KEY,
- PARTITION_ON,
- empty_dataframe_from_schema,
-)
-from generalresearch.incite.schemas.thl_marketplaces import (
- InnovateSurveyHistorySchema,
- MorningSurveyTimeseriesSchema,
- SagoSurveyHistorySchema,
- SpectrumSurveyTimeseriesSchema,
-)
-from generalresearch.incite.schemas.thl_web import (
- LedgerSchema,
- THLIPInfoSchema,
- THLSessionSchema,
- THLTaskAdjustmentSchema,
- THLUserSchema,
- THLWallSchema,
- TransactionMetadataColumns,
- TxMetaSchema,
- TxSchema,
- UserHealthAuditLogSchema,
- UserHealthIPHistorySchema,
- UserHealthIPHistoryWSSchema,
-)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.sql_helper import SqlHelper
-
-LOG = logging.getLogger("incite")
-
-DT_STR = "%Y-%m-%d %H:%M:%S"
-
-
-class DFCollectionType(str, Enum):
- TEST = "test"
-
- USER = "thl_user"
- SESSION = "thl_session"
- WALL = "thl_wall"
- TASK_ADJUSTMENT = "thl_taskadjustment"
- IP_INFO = "thl_ipinformation"
-
- AUDIT_LOG = "userhealth_auditlog"
- IP_HISTORY = "userhealth_iphistory"
- IP_HISTORY_WS = "userhealth_iphistory_ws"
-
- LEDGER = "ledger"
-
- INNOVATE_SURVEY_HISTORY = "innovate_surveyhistory"
- MORNING_SURVEY_TIMESERIES = "morning_surveytimeseries"
- SAGO_SURVEY_HISTORY = "sago_surveyhistory"
- SPECTRUM_SURVEY_TIMESERIES = "spectrum_surveytimeseries"
-
-
-DFCollectionTypeSchemas = {
- DFCollectionType.USER: THLUserSchema,
- DFCollectionType.WALL: THLWallSchema,
- DFCollectionType.SESSION: THLSessionSchema,
- DFCollectionType.IP_INFO: THLIPInfoSchema,
- DFCollectionType.TASK_ADJUSTMENT: THLTaskAdjustmentSchema,
- DFCollectionType.IP_HISTORY: UserHealthIPHistorySchema,
- DFCollectionType.IP_HISTORY_WS: UserHealthIPHistoryWSSchema,
- DFCollectionType.AUDIT_LOG: UserHealthAuditLogSchema,
- DFCollectionType.LEDGER: LedgerSchema,
- DFCollectionType.INNOVATE_SURVEY_HISTORY: InnovateSurveyHistorySchema,
- DFCollectionType.MORNING_SURVEY_TIMESERIES: MorningSurveyTimeseriesSchema,
- DFCollectionType.SAGO_SURVEY_HISTORY: SagoSurveyHistorySchema,
- DFCollectionType.SPECTRUM_SURVEY_TIMESERIES: SpectrumSurveyTimeseriesSchema,
-}
-
-
-class DFCollectionItem(CollectionItemBase):
-
- # --- Properties ---
- @property
- def filename(self) -> str:
- return (
- f"{self._collection.data_type.name.lower()}-{self._collection.offset}"
- f"-{self.start.strftime('%Y-%m-%d-%H-%M-%S')}.parquet"
- )
-
- # --- Methods ---
-
- def has_mysql(self) -> bool:
- if self._collection.sql_helper is None:
- return False
-
- connected = True
- try:
- self._collection.sql_helper.execute_sql_query("""SELECT 1;""")
- except:
- connected = False
-
- return connected
-
- def has_postgres(self) -> bool:
- if self._collection.pg_config is None:
- return False
-
- connected = True
- try:
- self._collection.pg_config.execute_sql_query("""SELECT 1;""")
- except:
- connected = False
-
- return connected
-
- def has_db(self) -> bool:
- return self.has_mysql() or self.has_postgres()
-
- def update_partial_archive(self) -> bool:
- if not self.valid_archive(self.partial_path, sample=1000):
- LOG.error(f"invalid partial archive: {self.partial_path}")
- return self.create_partial_archive()
- df = pq.ParquetDataset(self.partial_path).read().to_pandas()
-
- order_key = self._collection._schema.metadata[ORDER_KEY]
- archive_after = self._collection._schema.metadata[ARCHIVE_AFTER]
-
- partial_max = df[order_key].max().to_pydatetime()
-
- since = partial_max - archive_after
- since = max([since, self.start]) # don't allow to query before the item's start
- df = df[df[order_key] < since].copy()
-
- _df = self.from_mysql(since=since)
-
- if _df is not None:
- df = pd.concat([df, _df])
- self.to_archive(ddf=dd.from_pandas(df, npartitions=1), is_partial=True)
- else:
- # The update to the partial returned no rows, but the partial
- # still exists, so we'll continue with whatever was calling this.
- # We don't need to re-write the partial or really do anything.
- pass
- return True
-
- def create_partial_archive(self) -> bool:
- _df = self.from_mysql()
- if _df is None:
- # Returned no rows, but the period is not closed, so we
- # don't want to mark as empty. Do nothing.
- return False
- return self.to_archive(ddf=dd.from_pandas(_df, npartitions=1), is_partial=True)
-
- # --- ORM / Data handlers---
- def to_dict(self) -> dict[str, Any]:
- return self._to_dict()
-
- def from_mysql(self, since: datetime | None = None) -> pd.DataFrame | None:
- if self._collection.data_type == DFCollectionType.LEDGER:
- assert since is None, "Shouldn't pass since for Ledger item"
- assert self._collection.pg_config is not None
- return self.from_postgres_ledger()
- else:
- if self._collection.sql_helper:
- return self.from_mysql_standard(since=since)
- else:
- return self.from_postgres_standard(since=since)
-
- def from_mysql_standard(self, since: datetime | None = None) -> pd.DataFrame | None:
-
- assert (
- self._collection.data_type != DFCollectionType.LEDGER
- ), "Can't call from_mysql_standard for Ledger DFCollectionItem"
-
- start, finish = self.start, self.finish
- LOG.debug(
- f"{self._collection.data_type.value}.from_mysql("
- f"start={start.strftime(DT_STR)}, "
- f"finish={finish.strftime(DT_STR)})"
- )
- coll = self._collection
- schema = coll._schema
- sql_helper = coll.sql_helper
-
- start = since or start
- order_key = schema.metadata[ORDER_KEY]
- cols = list(schema.columns.keys()) + [schema.index.name]
- cols_str = ",".join(map(sql_helper._quote, cols))
- db_name = sql_helper.db
-
- try:
- res = sql_helper.execute_sql_query(
- query=f"""
- SELECT {cols_str}
- FROM `{db_name}`.`{coll.data_type.value}`
- WHERE `{order_key}` >= %s AND `{order_key}` < %s;
- """,
- params=[start, finish],
- )
- except (Exception,) as e:
- capture_exception(error=e)
- LOG.error(f"_from_mysql Exception: {e}")
- return None
-
- if not res:
- LOG.warning(f"_from_mysql query returned nothing")
- # Return an empty df.DataFrame with the correct columns
- return empty_dataframe_from_schema(coll._schema)
-
- df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name)
- df = self.validate_df(df=df)
-
- if df is None:
- LOG.warning(f"_from_mysql query results failed validation")
- # Schema validation can fail...
- return None
-
- return df
-
- def from_postgres_standard(
- self, since: datetime | None = None
- ) -> pd.DataFrame | None:
- assert (
- self._collection.data_type != DFCollectionType.LEDGER
- ), "Can't call from_postgres_standard for Ledger DFCollectionItem"
-
- start, finish = self.start, self.finish
- LOG.debug(
- f"{self._collection.data_type.value}.from_postgres("
- f"start={start.strftime(DT_STR)}, "
- f"finish={finish.strftime(DT_STR)})"
- )
- coll = self._collection
- schema = coll._schema
- pg_config = coll.pg_config
-
- start = since or start
- order_key = schema.metadata[ORDER_KEY]
- cols = list(schema.columns.keys()) + [schema.index.name]
- cols_str = ", ".join(cols)
-
- try:
- res = pg_config.execute_sql_query(
- query=f"""
- SELECT {cols_str}
- FROM {coll.data_type.value}
- WHERE {order_key} >= %s AND {order_key} < %s;
- """,
- params=[start, finish],
- )
- except (Exception,) as e:
- capture_exception(error=e)
- LOG.error(f"_from_postgres Exception: {e}")
- return None
-
- if not res:
- LOG.warning(f"_from_postgres query returned nothing")
- # Return an empty df.DataFrame with the correct columns
- return empty_dataframe_from_schema(coll._schema)
-
- df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name)
- df = self.validate_df(df=df)
-
- if df is None:
- LOG.warning(f"_from_postgres query results failed validation")
- # Schema validation can fail...
- return None
-
- return df
-
- def from_postgres_ledger(self) -> pd.DataFrame | None:
- assert (
- self._collection.data_type == DFCollectionType.LEDGER
- ), "Can only call from_postgres_ledger on Ledger DFCollectionItem"
-
- start, finish = self.start, self.finish
- LOG.info(
- f"{self._collection.data_type.value}.from_postgres_ledger("
- f"start={start.strftime(DT_STR)}, "
- f"finish={finish.strftime(DT_STR)})"
- )
-
- coll = self._collection
- pg_config: PostgresConfig = coll.pg_config
-
- limit = 20000
- offset = 0
- res = []
- while True:
- logging.info(
- f"{self._collection.data_type.value}.from_postgres_ledger({limit=}, {offset=})"
- )
- chunk = pg_config.execute_sql_query(
- query=f"""
- SELECT lt.id AS tx_id, lt.created, lt.ext_description, lt.tag,
- le.id AS entry_id, le.direction, le.amount, le.account_id,
- la.display_name, la.qualified_name, la.account_type,
- la.normal_balance, la.reference_type, la.reference_uuid,
- la.currency
- FROM ledger_transaction AS lt
- LEFT JOIN ledger_entry AS le
- ON lt.id = le.transaction_id
- LEFT JOIN ledger_account AS la
- ON la.uuid = le.account_id
- WHERE lt.created >= %s AND lt.created < %s
- AND le.id IS NOT NULL
- ORDER BY lt.created
- LIMIT {limit} OFFSET {offset};
- """,
- params=[start, finish],
- )
- res.extend(chunk)
- if not chunk:
- break
- offset += limit
-
- if len(res) == 0:
- return None
-
- # Note (AND le.id IS NOT NULL): It is possible we have transactions with
- # no ledger entries. This is because the transaction creation failed
- # for some reason. The ledger is not unbalanced, it is just an orphan
- # transaction. Just skip those here.
-
- tx_df = TxSchema.validate(
- check_obj=pd.DataFrame.from_records(res).set_index("entry_id"),
- lazy=True,
- )
-
- tx_ids = list(tx_df["tx_id"].unique())
- metadata_res = []
- # "MySQL server has gone away" if this is too big
- conn = pg_config.make_connection()
- c: Cursor = conn.cursor()
- for chunk in chunked(tx_ids, n=5_000):
- c.execute(
- query=f"""
- SELECT ltm.transaction_id AS tx_id,
- ltm.id AS tx_metadata_id,
- ltm.key, ltm.value
- FROM ledger_transactionmetadata AS ltm
- WHERE ltm.transaction_id = ANY(%s);
- """,
- params=[chunk],
- )
- metadata_res += c.fetchall()
-
- conn.close()
-
- tx_meta = (
- pd.DataFrame(
- TxMetaSchema.validate(
- check_obj=pd.DataFrame.from_records(metadata_res).set_index(
- ["tx_id", "tx_metadata_id"]
- ),
- lazy=True,
- ).pivot(columns="key", values="value"),
- # This makes sure we expand to have all the possible columns
- columns=[e.value for e in TransactionMetadataColumns],
- )
- .groupby("tx_id")
- .first()
- )
-
- df = tx_df.merge(tx_meta, how="left", left_on="tx_id", right_index=True)
- df = self.validate_df(df=df)
-
- if df is None:
- # Schema validation can fail...
- return None
-
- return df
-
- def to_archive(
- self,
- ddf: dd.DataFrame,
- is_partial: bool = False,
- overwrite: bool = False,
- ) -> bool:
- """
- :returns: bool (saved_successful)
- """
- assert isinstance(ddf, dd.DataFrame), "must pass dask df"
-
- client: Client | None = self._collection._client
- # client = None
-
- if client:
- row_len = client.compute(collections=ddf.shape[0], sync=True)
- else:
- row_len = len(ddf.index)
- is_empty = row_len == 0
-
- if is_partial:
- return self.to_archive_numbered_partial(ddf=ddf)
- else:
- return self._to_archive(
- ddf=ddf,
- is_empty=is_empty,
- overwrite=overwrite,
- )
-
- def _to_archive(
- self,
- ddf: dd.DataFrame | None,
- is_empty: bool,
- overwrite: bool = False,
- ) -> bool:
- """
- For archiving an item. Will write an empty file if ddf is empty. This
- is NOT for writing partials.
-
- :returns: bool (saved_successful)
- """
-
- if ddf is None:
- return False
-
- should_archive = self.should_archive()
- if not should_archive:
- LOG.warning(f"Cannot create archive for such new data: {self.path}")
- return False
-
- if overwrite is False:
- has_archive = self.has_archive(include_empty=True)
- if has_archive:
- LOG.warning(f"archive already exists: {self.path}")
- return False
-
- if is_empty:
- # Create an .empty only if the Item is "archiveable" (which we checked above)
- self.set_empty()
- return True
-
- # Incase the file saving is interrupted, or otherwise fails
- # save it to a tmp file first, then rename once we can confirm
- # that it successfully loads
- tmp_path = self.tmp_path()
- try:
- schema = self._collection._schema
- partition = schema.metadata.get(PARTITION_ON, None)
-
- ddf.to_parquet(
- path=tmp_path,
- partition_on=partition,
- engine="pyarrow",
- overwrite=True,
- write_metadata_file=True,
- compression="brotli",
- )
-
- except (Exception,) as e:
- LOG.exception(e)
- self.delete_archive(tmp_path)
- return False
-
- # It was saved, but the file seems to be corrupt
- if not self.valid_archive(tmp_path):
- LOG.error(f"not valid archive: {tmp_path}")
- self.delete_archive(tmp_path)
- # File did not save correctly so return it as saved=False
- return False
-
- # To debug, just set this key to auto expire in 5 seconds
- # RC.set(name=f"_to_archive:{self.path.as_posix()}", value=1, ex=15)
- # with RC.lock(f"_to_archive:{self.path.as_posix()}:lock", timeout=15):
-
- if os.path.isfile(tmp_path):
- # If the file was saved okay, seems okay, rename it
- os.replace(tmp_path, self.path)
- os.remove(tmp_path)
-
- if os.path.isdir(tmp_path):
- if os.path.exists(self.path.as_posix()):
- if overwrite:
- subprocess.call(["rm", "-r", self.path.as_posix()])
- time.sleep(1)
- else:
- LOG.error(f"already exists: {self.path.as_posix()}")
- return False
-
- if platform == "darwin":
- subprocess.call(["mv", tmp_path.as_posix(), self.path.as_posix()])
- else:
- # -T will (should) cause the mv to fail if path wasn't successfully deleted
- subprocess.call(["mv", "-T", tmp_path.as_posix(), self.path.as_posix()])
- return True
-
- def to_archive_numbered_partial(self, ddf: dd.DataFrame | None = None) -> bool:
- """
- For partial files/dirs only. Writes the .partial file with a number
- at the end (.partial.####) and then creates a symlink
- from .partial -> .partial.####
-
- :returns: bool (saved_successful)
- """
- if ddf is None:
- return False
-
- collection = self._collection
- schema = collection._schema
- client: Client | None = collection._client
-
- next_numbered_path = self.next_numbered_path(self.partial_path)
- partial_path = self.partial_path
- # finish = self.finish
-
- # Make sure these are in the same dir. b/c the symlink has to be
- # relative, not an absolute path
- assert (
- partial_path.parent == next_numbered_path.parent
- ), "Can't have numbered_path in a different directory"
- target = (
- next_numbered_path.name
- ) # this is the symlink's target. it is a relative path (only the name)
-
- should_archive = self.should_archive()
- assert should_archive is False, "Don't write partial if the item is archiveable"
-
- if client:
- row_len = client.compute(collections=ddf.shape[0], sync=True)
- else:
- row_len = len(ddf.index)
-
- if row_len == 0:
- LOG.warning("Skipping, don't partial save an empty dd.DataFrame")
- return False
-
- try:
- partition = schema.metadata.get(PARTITION_ON, None)
- ddf.to_parquet(
- path=next_numbered_path,
- partition_on=partition,
- engine="pyarrow",
- overwrite=True,
- write_metadata_file=True,
- compression="brotli",
- )
- except (Exception,) as e:
- LOG.exception(e)
- self.delete_archive(next_numbered_path)
- return False
-
- if platform == "darwin":
- subprocess.call(["ln", "-sfn", target, partial_path])
- else:
- subprocess.call(["ln", "-sfnT", target, partial_path])
-
- return True
-
- def initial_load(self, overwrite: bool = False) -> bool:
-
- if overwrite is False:
- assert not self.has_archive(include_empty=True), "already archived"
-
- assert self.should_archive(), "not ready to archive!"
-
- df: pd.DataFrame | None = self.from_mysql()
-
- if df is None:
- self.set_empty()
- return False
-
- ddf = dd.from_pandas(df, npartitions=1)
- return self.to_archive(ddf=ddf, is_partial=False, overwrite=overwrite)
-
- def clear_corrupt_archive(self):
- if self.has_archive(include_empty=False) and not self.valid_archive(self.path):
- LOG.warning(f"invalid archive, deleting: {self.path}")
- self.delete_archive(self.path)
-
-
-class DFCollection(CollectionBase):
- data_type: DFCollectionType | None = Field(default=None)
-
- # --- Private ---
- pg_config: PostgresConfig | None = Field(default=None)
- sql_helper: SqlHelper | None = Field(default=None)
-
- def __repr__(self):
- res = self.signature() + "\n"
- if len(self.items) > 6:
- items = self.items[:3] + ["..."] + self.items[-3:]
- else:
- items = self.items
-
- for i in items:
- res += f" – {repr(i) if isinstance(i, DFCollectionItem) else i}\n"
-
- return res
-
- def signature(self):
- arr = [
- 1 if i.has_archive(include_empty=True) else 0
- for i in self.items
- if i.should_archive()
- ]
- repr_str = (
- f"items={len(self.items)}; start={self.start} @ {self.offset}; {int(sum(arr) / len(arr) * 100)}% "
- f"archived"
- )
- res = f"{self.__repr_name__()}({repr_str})"
- return res
-
- @field_validator("data_type")
- def check_data_type(cls, data_type, info: ValidationInfo):
- if data_type is None:
- raise ValueError("Must explicitly provide a data_type")
-
- if data_type not in DFCollectionTypeSchemas:
- raise ValueError("Must provide a supported data_type")
-
- return data_type
-
- # --- Properties ---
- @property
- def items(self) -> list[DFCollectionItem]:
- items = []
- for iv in self.interval_range:
- cm = DFCollectionItem(start=iv[0])
- cm._collection = self
- items.append(cm)
- return items
-
- @property
- def _schema(self) -> DataFrameSchema:
- return DFCollectionTypeSchemas[self.data_type]
-
- # --- Methods ---
-
- def initial_load(
- self,
- client: Client | None = None,
- sync: bool = True,
- since: datetime | None = None,
- client_resources: dict[str, Any] | None = None,
- timeout: float | None = None,
- ) -> list[Future]:
- # This can be used to just build all local archive files
- # We typically want to go backwards first, so we can most quickly
- # populate the last 90 days for example
-
- client = client or self._client
-
- LOG.info(f"{self.data_type.value}.initial_load({since=}, {sync=})")
-
- items = self.items
- if since:
- items = self.get_items(since=since)
-
- if client is None:
- for item in reversed(items):
- if item.has_archive(include_empty=True):
- continue
- if not item.should_archive():
- continue
- item.initial_load()
- return []
-
- fs = []
- for item in items:
- if item.has_archive(include_empty=True):
- continue
- if not item.should_archive():
- continue
- f = dask.delayed(item.initial_load)()
- fs.append(f)
-
- if sync:
- fs = client.compute(fs, sync=False, priority=2, resources=client_resources)
- ac = as_completed(fs, timeout=timeout)
- return fs
-
- else:
- return client.compute(fs, sync=True, priority=2, resources=client_resources)
-
- def fetch_force_rr_latest(self, sources) -> list[FilePath]:
- LOG.info(
- f"{self.data_type.value}.fetch_force_rr_latest(sources={len(sources)})"
- )
-
- # We only want 'partial-able' items (those that can not yet be archived).
- rr_items = [
- i for i in self.items if not i.should_archive() and not i.is_empty()
- ]
- if rr_items:
- # If the ARCHIVE_AFTER time is > the collection offset (which it is always currently),
- # then there typically wouldn't be more than 1 un-archivable item.
- _start = rr_items[0].start
- _end = rr_items[-1].finish
- rr_duration = (_end - _start).total_seconds()
-
- # TODO: Do we want to be smarter about any rr selects max durations?
- # allowing 2x the length of the offset. If we have more than this not archived,
- # we want to run the archive first, not fetch from rr
- archive_after = self._schema.metadata[ARCHIVE_AFTER]
- allowed_rr_duration = (
- (pd.Timedelta(self.offset) * 2) + archive_after
- ).total_seconds()
- if rr_duration > allowed_rr_duration:
- raise ValueError(
- f"rr select duration exceeds {pd.Timedelta(allowed_rr_duration)}"
- )
-
- for rr_item in rr_items:
- if (
- rr_item.has_partial_archive()
- and self.data_type != DFCollectionType.LEDGER
- ):
- saved = rr_item.update_partial_archive()
- else:
- saved = rr_item.create_partial_archive()
- if saved:
- sources.append(rr_item.partial_path)
-
- return sources
-
- def force_rr_latest(
- self,
- client: Client,
- client_resources: dict[str, Any] | None = None,
- sync: bool = True,
- ) -> list[Future]:
-
- # For forcing update of any partials asynchronously if desired
- LOG.info(f"{self.data_type.value}.force_rr_latest({client=})")
-
- rr_items = [
- i for i in self.items if not i.should_archive() and not i.is_empty()
- ]
- fs = []
- for rr_item in rr_items:
- if (
- rr_item.has_partial_archive()
- and self.data_type != DFCollectionType.LEDGER
- ):
- fs.append(dask.delayed(rr_item.update_partial_archive)())
- else:
- fs.append(dask.delayed(rr_item.create_partial_archive)())
- return client.compute(fs, sync=sync, priority=2, resources=client_resources)
diff --git a/generalresearch/incite/collections/base.py b/generalresearch/incite/collections/base.py
new file mode 100644
index 0000000..04dd97e
--- /dev/null
+++ b/generalresearch/incite/collections/base.py
@@ -0,0 +1,777 @@
+from __future__ import annotations
+
+import logging
+import os
+import subprocess
+import time
+import warnings
+from datetime import datetime
+from enum import StrEnum
+from sys import platform
+from typing import Any
+
+import dask
+import dask.dataframe as dd
+import pandas as pd
+import pyarrow as pa
+import pyarrow.parquet as pq
+from dask.distributed import Client as DaskClient
+from dask.distributed import Future
+from distributed import as_completed
+from more_itertools import chunked
+from pandera.pandas import DataFrameSchema
+from psycopg import Cursor
+from pydantic import Field, FilePath, ValidationInfo, field_validator
+from sentry_sdk import capture_exception
+
+from generalresearch.incite import LOG
+from generalresearch.incite.base import CollectionBase, CollectionItemBase
+from generalresearch.incite.schemas import (
+ ARCHIVE_AFTER,
+ ORDER_KEY,
+ PARTITION_ON,
+ empty_dataframe_from_schema,
+)
+from generalresearch.incite.schemas.thl_marketplaces import (
+ InnovateSurveyHistorySchema,
+ MorningSurveyTimeseriesSchema,
+ SagoSurveyHistorySchema,
+ SpectrumSurveyTimeseriesSchema,
+)
+from generalresearch.incite.schemas.thl_web import (
+ LedgerSchema,
+ THLIPInfoSchema,
+ THLSessionSchema,
+ THLTaskAdjustmentSchema,
+ THLUserSchema,
+ THLWallSchema,
+ TransactionMetadataColumns,
+ TxMetaSchema,
+ TxSchema,
+ UserHealthAuditLogSchema,
+ UserHealthIPHistorySchema,
+ UserHealthIPHistoryWSSchema,
+)
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.sql_helper import SqlHelper
+
+DT_STR = "%Y-%m-%d %H:%M:%S"
+
+
+class DFCollectionType(StrEnum):
+ TEST = "test"
+
+ USER = "thl_user"
+ SESSION = "thl_session"
+ WALL = "thl_wall"
+ TASK_ADJUSTMENT = "thl_taskadjustment"
+ IP_INFO = "thl_ipinformation"
+
+ AUDIT_LOG = "userhealth_auditlog"
+ IP_HISTORY = "userhealth_iphistory"
+ IP_HISTORY_WS = "userhealth_iphistory_ws"
+
+ LEDGER = "ledger"
+
+ INNOVATE_SURVEY_HISTORY = "innovate_surveyhistory"
+ MORNING_SURVEY_TIMESERIES = "morning_surveytimeseries"
+ SAGO_SURVEY_HISTORY = "sago_surveyhistory"
+ SPECTRUM_SURVEY_TIMESERIES = "spectrum_surveytimeseries"
+
+
+DFCollectionTypeSchemas = {
+ DFCollectionType.USER: THLUserSchema,
+ DFCollectionType.WALL: THLWallSchema,
+ DFCollectionType.SESSION: THLSessionSchema,
+ DFCollectionType.IP_INFO: THLIPInfoSchema,
+ DFCollectionType.TASK_ADJUSTMENT: THLTaskAdjustmentSchema,
+ DFCollectionType.IP_HISTORY: UserHealthIPHistorySchema,
+ DFCollectionType.IP_HISTORY_WS: UserHealthIPHistoryWSSchema,
+ DFCollectionType.AUDIT_LOG: UserHealthAuditLogSchema,
+ DFCollectionType.LEDGER: LedgerSchema,
+ DFCollectionType.INNOVATE_SURVEY_HISTORY: InnovateSurveyHistorySchema,
+ DFCollectionType.MORNING_SURVEY_TIMESERIES: MorningSurveyTimeseriesSchema,
+ DFCollectionType.SAGO_SURVEY_HISTORY: SagoSurveyHistorySchema,
+ DFCollectionType.SPECTRUM_SURVEY_TIMESERIES: SpectrumSurveyTimeseriesSchema,
+}
+
+# This is not a technical limitation, it is just a double check
+# since this should never happen since thl is now and forever
+# forward on postgres
+MYSQL_ALLOWED_COLL_TYPES = {
+ DFCollectionType.INNOVATE_SURVEY_HISTORY,
+ DFCollectionType.MORNING_SURVEY_TIMESERIES,
+ DFCollectionType.SAGO_SURVEY_HISTORY,
+ DFCollectionType.SPECTRUM_SURVEY_TIMESERIES,
+}
+
+
+class DFCollectionItem(CollectionItemBase):
+ # --- Properties ---
+ @property
+ def filename(self) -> str:
+ return (
+ f"{self._collection.data_type.name.lower()}-{self._collection.offset}"
+ f"-{self.start.strftime('%Y-%m-%d-%H-%M-%S')}.parquet"
+ )
+
+ # --- Methods ---
+
+ def has_mysql(self) -> bool:
+ if self._collection.sql_helper is None:
+ return False
+
+ connected = True
+ try:
+ self._collection.sql_helper.execute_sql_query("""SELECT 1;""")
+ except:
+ connected = False
+
+ return connected
+
+ def has_postgres(self) -> bool:
+ if self._collection.pg_config is None:
+ return False
+
+ connected = True
+ try:
+ self._collection.pg_config.execute_sql_query("""SELECT 1;""")
+ except Exception:
+ connected = False
+
+ return connected
+
+ def has_db(self) -> bool:
+ return self.has_mysql() or self.has_postgres()
+
+ def update_partial_archive(self) -> bool:
+ if not self.valid_archive(self.partial_path, sample=1000):
+ LOG.error(f"invalid partial archive: {self.partial_path}")
+ return self.create_partial_archive()
+ df = pq.ParquetDataset(self.partial_path).read().to_pandas()
+
+ order_key = self._collection._schema.metadata[ORDER_KEY]
+ archive_after = self._collection._schema.metadata[ARCHIVE_AFTER]
+
+ partial_max = df[order_key].max().to_pydatetime()
+
+ since = partial_max - archive_after
+ since = max([since, self.start]) # don't allow to query before the item's start
+ df = df[df[order_key] < since].copy()
+
+ _df = self.from_db(since=since)
+
+ if _df is not None:
+ df = pd.concat([df, _df])
+ self.to_archive(ddf=dd.from_pandas(df, npartitions=1), is_partial=True)
+ else:
+ # The update to the partial returned no rows, but the partial
+ # still exists, so we'll continue with whatever was calling this.
+ # We don't need to re-write the partial or really do anything.
+ pass
+ return True
+
+ def create_partial_archive(self) -> bool:
+ _df = self.from_db()
+ if _df is None:
+ # Returned no rows, but the period is not closed, so we
+ # don't want to mark as empty. Do nothing.
+ return False
+ return self.to_archive(ddf=dd.from_pandas(_df, npartitions=1), is_partial=True)
+
+ # --- ORM / Data handlers---
+ def to_dict(self) -> dict[str, Any]:
+ return self._to_dict()
+
+ def from_db(self, since: datetime | None = None) -> pd.DataFrame | None:
+ if self._collection.data_type == DFCollectionType.LEDGER:
+ assert since is None, "Shouldn't pass since for Ledger item"
+ assert self._collection.pg_config is not None
+ return self.from_postgres_ledger()
+ else:
+ if self._collection.sql_helper:
+ return self.from_mysql_standard(since=since)
+ else:
+ return self.from_postgres_standard(since=since)
+
+ def from_mysql_standard(self, since: datetime | None = None) -> pd.DataFrame | None:
+ data_type = self._collection.data_type
+ assert data_type in MYSQL_ALLOWED_COLL_TYPES, (
+ f"Unsupported {data_type=} for mysql"
+ )
+
+ start, finish = self.start, self.finish
+ LOG.debug(
+ f"{data_type.value}.from_mysql("
+ f"start={start.strftime(DT_STR)}, "
+ f"finish={finish.strftime(DT_STR)})"
+ )
+ coll = self._collection
+ schema = coll._schema
+ sql_helper = coll.sql_helper
+
+ start = since or start
+ order_key = schema.metadata[ORDER_KEY]
+ cols = list(schema.columns.keys()) + [schema.index.name]
+ cols_str = ",".join(map(sql_helper._quote, cols))
+ db_name = sql_helper.db
+
+ try:
+ res = sql_helper.execute_sql_query(
+ query=f"""
+ SELECT {cols_str}
+ FROM `{db_name}`.`{coll.data_type.value}`
+ WHERE `{order_key}` >= %s AND `{order_key}` < %s;
+ """,
+ params=[start, finish],
+ )
+ except Exception as e:
+ capture_exception(error=e)
+ LOG.error(f"_from_mysql Exception: {e}")
+ return None
+
+ if not res:
+ LOG.warning("_from_mysql query returned nothing")
+ # Return an empty df.DataFrame with the correct columns
+ return empty_dataframe_from_schema(coll._schema)
+
+ df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name)
+ df = self.validate_df(df=df)
+
+ if df is None:
+ LOG.warning("_from_mysql query results failed validation")
+ # Schema validation can fail...
+ return None
+
+ return df
+
+ def from_postgres_standard(
+ self, since: datetime | None = None
+ ) -> pd.DataFrame | None:
+ data_type = self._collection.data_type
+ assert data_type != DFCollectionType.LEDGER, (
+ "Can't call from_postgres_standard for Ledger DFCollectionItem"
+ )
+
+ start, finish = self.start, self.finish
+ LOG.debug(
+ f"{data_type.value}.from_postgres("
+ f"start={start.strftime(DT_STR)}, "
+ f"finish={finish.strftime(DT_STR)})"
+ )
+ coll = self._collection
+ schema = coll._schema
+ pg_config = coll.pg_config
+
+ start = since or start
+ order_key = schema.metadata[ORDER_KEY]
+ cols = list(schema.columns.keys()) + [schema.index.name]
+ cols_str = ", ".join(cols)
+
+ try:
+ res = pg_config.execute_sql_query(
+ query=f"""
+ SELECT {cols_str}
+ FROM {data_type.value}
+ WHERE {order_key} >= %s AND {order_key} < %s;
+ """,
+ params=[start, finish],
+ )
+ except Exception as e:
+ capture_exception(error=e)
+ LOG.error(f"_from_postgres Exception: {e}")
+ return None
+
+ if not res:
+ LOG.warning("_from_postgres query returned nothing")
+ # Return an empty df.DataFrame with the correct columns
+ return empty_dataframe_from_schema(coll._schema)
+
+ df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name)
+ df = self.validate_df(df=df)
+
+ if df is None:
+ LOG.warning("_from_postgres query results failed validation")
+ # Schema validation can fail...
+ return None
+
+ return df
+
+ def from_postgres_ledger(self) -> pd.DataFrame | None:
+ data_type = self._collection.data_type
+ assert data_type == DFCollectionType.LEDGER, (
+ "Can only call from_postgres_ledger on Ledger DFCollectionItem"
+ )
+
+ start, finish = self.start, self.finish
+ LOG.info(
+ f"{data_type.value}.from_postgres_ledger("
+ f"start={start.strftime(DT_STR)}, "
+ f"finish={finish.strftime(DT_STR)})"
+ )
+
+ coll = self._collection
+ pg_config: PostgresConfig = coll.pg_config
+
+ limit = 20000
+ offset = 0
+ res = []
+ while True:
+ logging.info(f"{data_type.value}.from_postgres_ledger({limit=}, {offset=})")
+ chunk = pg_config.execute_sql_query(
+ query=f"""
+ SELECT lt.id AS tx_id, lt.created, lt.ext_description, lt.tag,
+ le.id AS entry_id, le.direction, le.amount, le.account_id,
+ la.display_name, la.qualified_name, la.account_type,
+ la.normal_balance, la.reference_type, la.reference_uuid,
+ la.currency
+ FROM ledger_transaction AS lt
+ LEFT JOIN ledger_entry AS le
+ ON lt.id = le.transaction_id
+ LEFT JOIN ledger_account AS la
+ ON la.uuid = le.account_id
+ WHERE lt.created >= %s AND lt.created < %s
+ AND le.id IS NOT NULL
+ ORDER BY lt.created
+ LIMIT {limit} OFFSET {offset};
+ """,
+ params=[start, finish],
+ )
+ res.extend(chunk)
+ if not chunk:
+ break
+ offset += limit
+
+ if len(res) == 0:
+ return None
+
+ # Note (AND le.id IS NOT NULL): It is possible we have transactions with
+ # no ledger entries. This is because the transaction creation failed
+ # for some reason. The ledger is not unbalanced, it is just an orphan
+ # transaction. Just skip those here.
+
+ tx_df = TxSchema.validate(
+ check_obj=pd.DataFrame.from_records(res).set_index("entry_id"),
+ lazy=True,
+ )
+
+ tx_ids = list(tx_df["tx_id"].unique())
+ metadata_res = []
+ # "MySQL server has gone away" if this is too big
+ conn = pg_config.make_connection()
+ c: Cursor = conn.cursor()
+ for chunk in chunked(tx_ids, n=5_000):
+ c.execute(
+ query="""
+ SELECT ltm.transaction_id AS tx_id,
+ ltm.id AS tx_metadata_id,
+ ltm.key, ltm.value
+ FROM ledger_transactionmetadata AS ltm
+ WHERE ltm.transaction_id = ANY(%s);
+ """,
+ params=[chunk],
+ )
+ metadata_res += c.fetchall()
+
+ conn.close()
+
+ tx_meta = (
+ pd.DataFrame(
+ TxMetaSchema.validate(
+ check_obj=pd.DataFrame.from_records(metadata_res).set_index(
+ ["tx_id", "tx_metadata_id"]
+ ),
+ lazy=True,
+ ).pivot(columns="key", values="value"),
+ # This makes sure we expand to have all the possible columns
+ columns=[e.value for e in TransactionMetadataColumns],
+ )
+ .groupby("tx_id")
+ .first()
+ )
+
+ df = tx_df.merge(tx_meta, how="left", left_on="tx_id", right_index=True)
+ df = self.validate_df(df=df)
+
+ if df is None:
+ # Schema validation can fail...
+ return None
+
+ return df
+
+ def to_archive(
+ self,
+ ddf: dd.DataFrame,
+ is_partial: bool = False,
+ overwrite: bool = False,
+ ) -> bool:
+ """
+ :returns: bool (saved_successful)
+ """
+ assert isinstance(ddf, dd.DataFrame), "must pass dask df"
+
+ client: DaskClient | None = self._collection._client
+ # client = None
+
+ if client:
+ row_len = client.compute(collections=ddf.shape[0], sync=True)
+ else:
+ row_len = len(ddf.index)
+ is_empty = row_len == 0
+
+ if is_partial:
+ return self.to_archive_numbered_partial(ddf=ddf)
+ else:
+ return self._to_archive(
+ ddf=ddf,
+ is_empty=is_empty,
+ overwrite=overwrite,
+ )
+
+ def _to_archive(
+ self,
+ ddf: dd.DataFrame | None,
+ is_empty: bool,
+ overwrite: bool = False,
+ ) -> bool:
+ """
+ For archiving an item. Will write an empty file if ddf is empty. This
+ is NOT for writing partials.
+
+ :returns: bool (saved_successful)
+ """
+
+ if ddf is None:
+ return False
+
+ should_archive = self.should_archive()
+ if not should_archive:
+ LOG.warning(f"Cannot create archive for such new data: {self.path}")
+ return False
+
+ if overwrite is False:
+ has_archive = self.has_archive(include_empty=True)
+ if has_archive:
+ LOG.warning(f"archive already exists: {self.path}")
+ return False
+
+ if is_empty:
+ # Create an .empty only if the Item is "archiveable" (which we checked above)
+ self.set_empty()
+ return True
+
+ # Incase the file saving is interrupted, or otherwise fails
+ # save it to a tmp file first, then rename once we can confirm
+ # that it successfully loads
+ tmp_path = self.tmp_path()
+ try:
+ schema = self._collection._schema
+ assert schema
+ assert schema.metadata
+ partition = schema.metadata.get(PARTITION_ON)
+
+ ddf.to_parquet(
+ path=tmp_path,
+ partition_on=partition,
+ engine="pyarrow",
+ overwrite=True,
+ write_metadata_file=True,
+ compression="brotli",
+ )
+
+ except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e:
+ LOG.exception(e)
+ self.delete_archive(tmp_path)
+ return False
+
+ # It was saved, but the file seems to be corrupt
+ if not self.valid_archive(tmp_path):
+ LOG.error(f"not valid archive: {tmp_path}")
+ self.delete_archive(tmp_path)
+ # File did not save correctly so return it as saved=False
+ return False
+
+ # To debug, just set this key to auto expire in 5 seconds
+ # RC.set(name=f"_to_archive:{self.path.as_posix()}", value=1, ex=15)
+ # with RC.lock(f"_to_archive:{self.path.as_posix()}:lock", timeout=15):
+
+ if os.path.isfile(tmp_path):
+ # If the file was saved okay, seems okay, rename it
+ os.replace(tmp_path, self.path)
+ os.remove(tmp_path)
+
+ if os.path.isdir(tmp_path):
+ if os.path.exists(self.path.as_posix()):
+ if overwrite:
+ subprocess.call(["rm", "-r", self.path.as_posix()])
+ time.sleep(1)
+ else:
+ LOG.error(f"already exists: {self.path.as_posix()}")
+ return False
+
+ if platform == "darwin":
+ subprocess.call(["mv", tmp_path.as_posix(), self.path.as_posix()])
+ else:
+ # -T will (should) cause the mv to fail if path wasn't successfully deleted
+ subprocess.call(["mv", "-T", tmp_path.as_posix(), self.path.as_posix()])
+ return True
+
+ def to_archive_numbered_partial(self, ddf: dd.DataFrame | None = None) -> bool:
+ """
+ For partial files/dirs only. Writes the .partial file with a number
+ at the end (.partial.####) and then creates a symlink
+ from .partial -> .partial.####
+
+ :returns: bool (saved_successful)
+ """
+ if ddf is None:
+ return False
+
+ collection = self._collection
+ schema = collection._schema
+ assert schema
+ client: DaskClient | None = collection._client
+
+ next_numbered_path = self.next_numbered_path(self.partial_path)
+ partial_path = self.partial_path
+ # finish = self.finish
+
+ # Make sure these are in the same dir. b/c the symlink has to be
+ # relative, not an absolute path
+ assert partial_path.parent == next_numbered_path.parent, (
+ "Can't have numbered_path in a different directory"
+ )
+ target = (
+ next_numbered_path.name
+ ) # this is the symlink's target. it is a relative path (only the name)
+
+ should_archive = self.should_archive()
+ assert should_archive is False, "Don't write partial if the item is archiveable"
+
+ if client:
+ row_len = client.compute(collections=ddf.shape[0], sync=True)
+ else:
+ row_len = len(ddf.index)
+
+ if row_len == 0:
+ LOG.warning("Skipping, don't partial save an empty dd.DataFrame")
+ return False
+
+ try:
+ assert schema.metadata
+ partition = schema.metadata.get(PARTITION_ON)
+ ddf.to_parquet(
+ path=next_numbered_path,
+ partition_on=partition,
+ engine="pyarrow",
+ overwrite=True,
+ write_metadata_file=True,
+ compression="brotli",
+ )
+ except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e:
+ LOG.exception(e)
+ self.delete_archive(next_numbered_path)
+ return False
+
+ if platform == "darwin":
+ subprocess.call(["ln", "-sfn", target, partial_path])
+ else:
+ subprocess.call(["ln", "-sfnT", target, partial_path])
+
+ return True
+
+ def initial_load(self, overwrite: bool = False) -> bool:
+
+ if overwrite is False:
+ assert not self.has_archive(include_empty=True), "already archived"
+
+ assert self.should_archive(), "not ready to archive!"
+
+ df: pd.DataFrame | None = self.from_db()
+
+ if df is None:
+ self.set_empty()
+ return False
+
+ ddf = dd.from_pandas(df, npartitions=1)
+ return self.to_archive(ddf=ddf, is_partial=False, overwrite=overwrite)
+
+ def clear_corrupt_archive(self):
+ if self.has_archive(include_empty=False) and not self.valid_archive(self.path):
+ LOG.warning(f"invalid archive, deleting: {self.path}")
+ self.delete_archive(self.path)
+
+
+class DFCollection(CollectionBase):
+ data_type: DFCollectionType | None = Field(default=None)
+
+ # --- Private ---
+ pg_config: PostgresConfig | None = Field(default=None)
+ sql_helper: SqlHelper | None = Field(default=None)
+
+ def __repr__(self):
+ res = self.signature() + "\n"
+ if len(self.items) > 6:
+ items = self.items[:3] + ["..."] + self.items[-3:]
+ else:
+ items = self.items
+
+ for i in items:
+ res += f" – {repr(i) if isinstance(i, DFCollectionItem) else i}\n"
+
+ return res
+
+ def signature(self):
+ arr = [
+ 1 if i.has_archive(include_empty=True) else 0
+ for i in self.items
+ if i.should_archive()
+ ]
+ repr_str = (
+ f"items={len(self.items)}; start={self.start} @ {self.offset}; {int(sum(arr) / len(arr) * 100)}% "
+ f"archived"
+ )
+ res = f"{self.__repr_name__()}({repr_str})"
+ return res
+
+ @field_validator("data_type")
+ def check_data_type(cls, data_type: DFCollectionType | None, info: ValidationInfo):
+ if data_type is None:
+ raise ValueError("Must explicitly provide a data_type")
+
+ if data_type not in DFCollectionTypeSchemas:
+ raise ValueError("Must provide a supported data_type")
+
+ return data_type
+
+ # --- Properties ---
+ @property
+ def items(self) -> list[DFCollectionItem]:
+ items = []
+ for iv in self.interval_range:
+ cm = DFCollectionItem(start=iv[0])
+ cm._collection = self
+ items.append(cm)
+ return items
+
+ @property
+ def type_schema(self) -> DataFrameSchema:
+ return DFCollectionTypeSchemas[self.data_type]
+
+ @property
+ def _schema(self) -> DataFrameSchema:
+ warnings.deprecated("The _schema attribute on DFCollection is Deprecated")
+ return self.type_schema
+
+ # --- Methods ---
+
+ def initial_load(
+ self,
+ client: DaskClient | None = None,
+ sync: bool = True,
+ since: datetime | None = None,
+ client_resources: dict[str, Any] | None = None,
+ timeout: float | None = None,
+ ) -> list[Future]:
+ # This can be used to just build all local archive files
+ # We typically want to go backwards first, so we can most quickly
+ # populate the last 90 days for example
+
+ client = client or self._client
+
+ LOG.info(f"{self.data_type.value}.initial_load({since=}, {sync=})")
+
+ items = self.items
+ if since:
+ items = self.get_items(since=since)
+
+ if client is None:
+ for item in reversed(items):
+ if item.has_archive(include_empty=True):
+ continue
+ if not item.should_archive():
+ continue
+ item.initial_load()
+ return []
+
+ fs = []
+ for item in items:
+ if item.has_archive(include_empty=True):
+ continue
+ if not item.should_archive():
+ continue
+ f = dask.delayed(item.initial_load)()
+ fs.append(f)
+
+ if sync:
+ fs = client.compute(fs, sync=False, priority=2, resources=client_resources)
+ _ = as_completed(fs, timeout=timeout)
+ return fs
+
+ else:
+ return client.compute(fs, sync=True, priority=2, resources=client_resources)
+
+ def fetch_force_rr_latest(self, sources) -> list[FilePath]:
+ LOG.info(
+ f"{self.data_type.value}.fetch_force_rr_latest(sources={len(sources)})"
+ )
+
+ # We only want 'partial-able' items (those that can not yet be archived).
+ rr_items = [
+ i for i in self.items if not i.should_archive() and not i.is_empty()
+ ]
+ if rr_items:
+ # If the ARCHIVE_AFTER time is > the collection offset (which it is always currently),
+ # then there typically wouldn't be more than 1 un-archivable item.
+ _start = rr_items[0].start
+ _end = rr_items[-1].finish
+ rr_duration = (_end - _start).total_seconds()
+
+ # TODO: Do we want to be smarter about any rr selects max durations?
+ # allowing 2x the length of the offset. If we have more than this not archived,
+ # we want to run the archive first, not fetch from rr
+ archive_after = self._schema.metadata[ARCHIVE_AFTER]
+ allowed_rr_duration = (
+ (pd.Timedelta(self.offset) * 2) + archive_after
+ ).total_seconds()
+ if rr_duration > allowed_rr_duration:
+ raise ValueError(
+ f"rr select duration exceeds {pd.Timedelta(allowed_rr_duration)}"
+ )
+
+ for rr_item in rr_items:
+ if (
+ rr_item.has_partial_archive()
+ and self.data_type != DFCollectionType.LEDGER
+ ):
+ saved = rr_item.update_partial_archive()
+ else:
+ saved = rr_item.create_partial_archive()
+ if saved:
+ sources.append(rr_item.partial_path)
+
+ return sources
+
+ def force_rr_latest(
+ self,
+ client: DaskClient,
+ client_resources: dict[str, Any] | None = None,
+ sync: bool = True,
+ ) -> list[Future]:
+
+ # For forcing update of any partials asynchronously if desired
+ LOG.info(f"{self.data_type.value}.force_rr_latest({client=})")
+
+ rr_items = [
+ i for i in self.items if not i.should_archive() and not i.is_empty()
+ ]
+ fs = []
+ for rr_item in rr_items:
+ if (
+ rr_item.has_partial_archive()
+ and self.data_type != DFCollectionType.LEDGER
+ ):
+ fs.append(dask.delayed(rr_item.update_partial_archive)())
+ else:
+ fs.append(dask.delayed(rr_item.create_partial_archive)())
+ return client.compute(fs, sync=sync, priority=2, resources=client_resources)
diff --git a/generalresearch/incite/collections/thl_marketplaces.py b/generalresearch/incite/collections/thl_marketplaces.py
index fe2b01f..246fe87 100644
--- a/generalresearch/incite/collections/thl_marketplaces.py
+++ b/generalresearch/incite/collections/thl_marketplaces.py
@@ -1,6 +1,6 @@
from typing import Literal
-from generalresearch.incite.collections import DFCollection, DFCollectionType
+from generalresearch.incite.collections.base import DFCollection, DFCollectionType
from generalresearch.incite.schemas.thl_marketplaces import (
InnovateSurveyHistorySchema,
MorningSurveyTimeseriesSchema,
diff --git a/generalresearch/incite/collections/thl_web.py b/generalresearch/incite/collections/thl_web.py
index 951406c..d60c77f 100644
--- a/generalresearch/incite/collections/thl_web.py
+++ b/generalresearch/incite/collections/thl_web.py
@@ -1,6 +1,6 @@
from typing import Literal
-from generalresearch.incite.collections import DFCollection, DFCollectionType
+from generalresearch.incite.collections.base import DFCollection, DFCollectionType
class UserDFCollection(DFCollection):
diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py
index 421710e..5ee305b 100644
--- a/generalresearch/incite/defaults.py
+++ b/generalresearch/incite/defaults.py
@@ -1,7 +1,9 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from generalresearch.incite.base import GRLDatasets
-from generalresearch.incite.collections import DFCollectionType
+from generalresearch.incite.collections.base import DFCollectionType
from generalresearch.incite.collections.thl_marketplaces import (
InnovateSurveyHistoryCollection,
MorningSurveyTimeseriesCollection,
@@ -15,7 +17,7 @@ from generalresearch.incite.collections.thl_web import (
UserDFCollection,
WallDFCollection,
)
-from generalresearch.incite.mergers import MergeType
+from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.foundations.enriched_session import (
EnrichedSessionMerge,
)
@@ -37,69 +39,65 @@ from generalresearch.sql_helper import SqlHelper
def session_df_collection(
- ds: "GRLDatasets", pg_config: PostgresConfig
+ ds: GRLDatasets, pg_config: PostgresConfig
) -> SessionDFCollection:
return SessionDFCollection(
offset="37h",
pg_config=pg_config,
- start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=timezone.utc),
+ start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.SESSION),
)
-def wall_df_collection(
- ds: "GRLDatasets", pg_config: PostgresConfig
-) -> WallDFCollection:
+def wall_df_collection(ds: GRLDatasets, pg_config: PostgresConfig) -> WallDFCollection:
return WallDFCollection(
offset="49h",
pg_config=pg_config,
- start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=timezone.utc),
+ start=datetime(year=2022, month=5, day=3, hour=12, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.WALL),
)
-def user_df_collection(
- ds: "GRLDatasets", pg_config: PostgresConfig
-) -> UserDFCollection:
+def user_df_collection(ds: GRLDatasets, pg_config: PostgresConfig) -> UserDFCollection:
return UserDFCollection(
offset="73h",
pg_config=pg_config,
- start=datetime(year=2016, month=7, day=13, hour=1, tzinfo=timezone.utc),
+ start=datetime(year=2016, month=7, day=13, hour=1, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.USER),
)
def task_df_collection(
- ds: "GRLDatasets", pg_config: PostgresConfig
+ ds: GRLDatasets, pg_config: PostgresConfig
) -> TaskAdjustmentDFCollection:
return TaskAdjustmentDFCollection(
offset="48h",
pg_config=pg_config,
- start=datetime(year=2022, month=7, day=16, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2022, month=7, day=16, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.TASK_ADJUSTMENT),
)
def ledger_df_collection(
- ds: "GRLDatasets", pg_config: PostgresConfig
+ ds: GRLDatasets, pg_config: PostgresConfig
) -> LedgerDFCollection:
return LedgerDFCollection(
- offset="12d",
+ offset="12D",
pg_config=pg_config,
# thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232
- start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.LEDGER),
)
# --- Marketplace Specifics --- #
def innovate_survey_history_collection(
- ds: "GRLDatasets", sql_helper: SqlHelper
+ ds: GRLDatasets, sql_helper: SqlHelper
) -> InnovateSurveyHistoryCollection:
return InnovateSurveyHistoryCollection(
offset="12h",
sql_helper=sql_helper,
- start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(
enum_type=DFCollectionType.INNOVATE_SURVEY_HISTORY
),
@@ -107,12 +105,12 @@ def innovate_survey_history_collection(
def morning_survey_ts_collection(
- ds: "GRLDatasets", sql_helper: SqlHelper
+ ds: GRLDatasets, sql_helper: SqlHelper
) -> MorningSurveyTimeseriesCollection:
return MorningSurveyTimeseriesCollection(
offset="12h",
sql_helper=sql_helper,
- start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(
enum_type=DFCollectionType.MORNING_SURVEY_TIMESERIES
),
@@ -120,23 +118,23 @@ def morning_survey_ts_collection(
def sago_survey_history_collection(
- ds: "GRLDatasets", sql_helper: SqlHelper
+ ds: GRLDatasets, sql_helper: SqlHelper
) -> SagoSurveyHistoryCollection:
return SagoSurveyHistoryCollection(
offset="12h",
sql_helper=sql_helper,
- start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(enum_type=DFCollectionType.SAGO_SURVEY_HISTORY),
)
def spectrum_survey_ts_collection(
- ds: "GRLDatasets", sql_helper: SqlHelper
+ ds: GRLDatasets, sql_helper: SqlHelper
) -> SpectrumSurveyTimeseriesCollection:
return SpectrumSurveyTimeseriesCollection(
offset="12h",
sql_helper=sql_helper,
- start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=timezone.utc),
+ start=datetime(year=2024, month=3, day=1, hour=0, tzinfo=UTC),
archive_path=ds.archive_path(
enum_type=DFCollectionType.SPECTRUM_SURVEY_TIMESERIES
),
@@ -144,50 +142,50 @@ def spectrum_survey_ts_collection(
# --- Mergers: Foundations --- #
-def user_id_product(ds: "GRLDatasets") -> UserIdProductMerge:
+def user_id_product(ds: GRLDatasets) -> UserIdProductMerge:
return UserIdProductMerge(
- start=datetime(year=2010, month=1, day=1, tzinfo=timezone.utc),
+ start=datetime(year=2010, month=1, day=1, tzinfo=UTC),
offset=None,
archive_path=ds.archive_path(enum_type=MergeType.USER_ID_PRODUCT),
)
-def enriched_session(ds: "GRLDatasets") -> EnrichedSessionMerge:
+def enriched_session(ds: GRLDatasets) -> EnrichedSessionMerge:
return EnrichedSessionMerge(
- start=datetime(year=2023, month=5, day=1, tzinfo=timezone.utc),
- offset="14d",
+ start=datetime(year=2023, month=5, day=1, tzinfo=UTC),
+ offset="14D",
archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_SESSION),
)
-def enriched_wall(ds: "GRLDatasets") -> EnrichedWallMerge:
+def enriched_wall(ds: GRLDatasets) -> EnrichedWallMerge:
return EnrichedWallMerge(
# start=datetime(year=2022, month=5, day=1, tzinfo=timezone.utc),
- start=datetime(year=2023, month=7, day=23, tzinfo=timezone.utc),
- offset="14d",
+ start=datetime(year=2023, month=7, day=23, tzinfo=UTC),
+ offset="14D",
archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_WALL),
)
-def enriched_task_adjust(ds: "GRLDatasets") -> EnrichedTaskAdjustMerge:
+def enriched_task_adjust(ds: GRLDatasets) -> EnrichedTaskAdjustMerge:
return EnrichedTaskAdjustMerge(
- start=datetime(year=2010, month=1, day=1, tzinfo=timezone.utc),
+ start=datetime(year=2010, month=1, day=1, tzinfo=UTC),
offset=None,
archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_TASK_ADJUST),
)
# --- Mergers: Others --- #
-def pop_ledger(ds: "GRLDatasets") -> PopLedgerMerge:
+def pop_ledger(ds: GRLDatasets) -> PopLedgerMerge:
return PopLedgerMerge(
# thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232
- start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc),
- offset="30d",
+ start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC),
+ offset="30D",
archive_path=ds.archive_path(enum_type=MergeType.POP_LEDGER),
)
-def ym_survey_wall(ds: "GRLDatasets") -> YMSurveyWallMerge:
+def ym_survey_wall(ds: GRLDatasets) -> YMSurveyWallMerge:
return YMSurveyWallMerge(
start=None,
offset="10D",
diff --git a/generalresearch/incite/exceptions.py b/generalresearch/incite/exceptions.py
new file mode 100644
index 0000000..112978f
--- /dev/null
+++ b/generalresearch/incite/exceptions.py
@@ -0,0 +1,19 @@
+class BuildItemsError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
+
+
+class BuildError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
+
+
+class FetchError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py
index b9c3789..e69de29 100644
--- a/generalresearch/incite/mergers/__init__.py
+++ b/generalresearch/incite/mergers/__init__.py
@@ -1,305 +0,0 @@
-import logging
-import os.path
-import subprocess
-from datetime import datetime, timezone
-from enum import Enum
-from sys import platform
-from typing import List, Optional, Type
-
-import dask.dataframe as dd
-import pandas as pd
-from dask.distributed import Client
-from pandera.pandas import DataFrameSchema
-from pydantic import Field, ValidationInfo, field_validator, model_validator
-from typing_extensions import Self
-
-from generalresearch.incite.base import CollectionBase, CollectionItemBase
-from generalresearch.incite.schemas import PARTITION_ON
-from generalresearch.incite.schemas.mergers.foundations.enriched_session import (
- EnrichedSessionSchema,
-)
-from generalresearch.incite.schemas.mergers.foundations.enriched_task_adjust import (
- EnrichedTaskAdjustSchema,
-)
-from generalresearch.incite.schemas.mergers.foundations.enriched_wall import (
- EnrichedWallSchema,
-)
-from generalresearch.incite.schemas.mergers.foundations.user_id_product import (
- UserIdProductSchema,
-)
-
-from generalresearch.incite.schemas.mergers.pop_ledger import (
- PopLedgerSchema,
-)
-from generalresearch.incite.schemas.mergers.ym_survey_wall import (
- YMSurveyWallSchema,
-)
-from generalresearch.incite.schemas.mergers.ym_wall_summary import (
- YMWallSummarySchema,
-)
-from generalresearch.models.custom_types import AwareDatetimeISO
-
-LOG = logging.getLogger("incite")
-
-
-class MergeType(str, Enum):
- TEST = "test"
- YM_SURVEY_WALL = "ym_survey_wall"
- YM_WALL_SUMMARY = "ym_wall_summary"
-
- POP_LEDGER = "pop_ledger"
-
- # --- Foundations ---
- USER_ID_PRODUCT = "user_id_product"
- ENRICHED_WALL = "enriched_wall"
- ENRICHED_SESSION = "enriched_session"
- ENRICHED_TASK_ADJUST = "enriched_task_adjust"
-
-
-MergeTypeSchemas = {
- MergeType.YM_SURVEY_WALL: YMSurveyWallSchema,
- MergeType.YM_WALL_SUMMARY: YMWallSummarySchema,
- MergeType.POP_LEDGER: PopLedgerSchema,
- # --- Foundations ---
- MergeType.USER_ID_PRODUCT: UserIdProductSchema,
- MergeType.ENRICHED_WALL: EnrichedWallSchema,
- MergeType.ENRICHED_SESSION: EnrichedSessionSchema,
- MergeType.ENRICHED_TASK_ADJUST: EnrichedTaskAdjustSchema,
-}
-
-
-class MergeCollectionItem(CollectionItemBase):
-
- # --- Properties ---
-
- @property
- def finish(self) -> datetime:
- # A MergeCollection can have offset = None
- if self._collection.offset:
- return (
- pd.Timestamp(self.start) + pd.Timedelta(self._collection.offset)
- ).to_pydatetime()
- else:
- return datetime.now(tz=timezone.utc).replace(microsecond=0)
-
- @property
- def filename(self) -> str:
- grouped_key = self._collection.grouped_key
- offset = self._collection.offset
- start = self.start.strftime("%Y-%m-%d-%H-%M-%S")
- f = [self._collection.merge_type.name.lower()]
- if offset:
- f.append(offset)
- if grouped_key:
- f.append(grouped_key)
- if self._collection.start is not None:
- # This is a collection that is "looking back" 'offset' time (1 item).
- f.append(start)
- s = "-".join(f)
- s += ".parquet"
- return s
-
- # --- ORM / Data handlers---
- def to_dict(self, *args, **kwargs) -> dict:
- res = self._to_dict()
- res["group_by"] = self._collection.group_by
- return res
-
- def to_archive(
- self,
- client: Client,
- ddf: dd.DataFrame,
- is_partial: bool = False,
- client_resources=None,
- ) -> bool:
- assert is_partial is False, "use to_archive_symlink"
- return self._to_archive(client=client, ddf=ddf, client_resources=None)
-
- def _to_archive(
- self, client: Client, ddf: dd.DataFrame, client_resources=None
- ) -> bool:
- """
- For archiving an item. Will write an empty file if ddf is empty.
- This is NOT for writing partials.
-
- :returns: bool (saved_successful)
- """
- if ddf is None:
- return False
-
- row_len = client.compute(collections=ddf.shape[0], sync=True)
- assert row_len > 0, "empty ddf"
-
- tmp_path = self.tmp_path()
- schema = self._collection._schema
- partition = schema.metadata.get(PARTITION_ON, None)
- f = ddf.to_parquet(
- compute=False,
- path=tmp_path,
- partition_on=partition,
- engine="pyarrow",
- overwrite=True,
- write_metadata_file=True,
- compression="brotli",
- )
- client.compute(f, sync=True, priority=2, resources=client_resources)
- assert not os.path.exists(
- self.path.as_posix()
- ), f"already exits!: {self.path.as_posix()}"
-
- if platform == "darwin":
- subprocess.call(["mv", tmp_path.as_posix(), self.path.as_posix()])
- else:
- # -T will (should) cause the mv to fail if `path` wasn't successfully deleted
- subprocess.call(["mv", "-T", tmp_path.as_posix(), self.path.as_posix()])
- return True
-
- def to_archive_symlink(
- self,
- client: Client,
- ddf: dd.DataFrame,
- is_partial: bool = False,
- client_resources=None,
- validate_after=True,
- ) -> bool:
- """
- This differs from to_archive():
- 1) to_parquet is run in this process. If the df is already
- computed, there is no point in sending it to another worker
- to write.
-
- 2) symlink to next_numbered_path is created whether or not
- is_partial (to_archive only does this on partials)
-
- 3) we do not validate the written file. seems not useful to do
- this, as the file will probably get overwritten on the next
- loop anyway
- """
- path = self.partial_path if is_partial else self.path
- next_numbered_path = self.next_numbered_path(path)
- collection = self._collection
- LOG.warning(f"{collection.merge_type.value}.to_archive_symlink()")
-
- if not isinstance(ddf, dd.DataFrame):
- raise ValueError("must pass a dask df")
-
- # We should validate before or after!!!
- # _validate_df(self.compute(ddf), coll._schema)
- target = (
- next_numbered_path.name
- ) # this is the symlink's target. it is a relative path (only the name)
-
- schema = self._collection._schema
- partition = schema.metadata.get(PARTITION_ON, None)
- f = ddf.to_parquet(
- compute=False,
- path=next_numbered_path.as_posix(),
- partition_on=partition,
- engine="pyarrow",
- overwrite=True,
- write_metadata_file=True,
- compression="brotli",
- )
- client.compute(f, sync=True, priority=2, resources=client_resources)
-
- if os.path.exists(path.as_posix()) and not os.path.islink(path.as_posix()):
- # This will fail when going from the old way to using symlinks,
- # if self.path already exists and is a directory.
- raise ValueError(
- f"first time we run this, make sure the path doesnt exist: {path.as_posix()}"
- )
-
- if platform == "darwin":
- subprocess.call(["ln", "-sfn", target, path.as_posix()])
- else:
- subprocess.call(["ln", "-sfnT", target, path.as_posix()])
-
- if validate_after:
- if not self.valid_archive(self.path):
- LOG.error(
- f"{collection.merge_type.value} failed validation: {self.path}"
- )
- self.delete_archive(self.path)
- return False
- return True
-
- # todo: unclear what the common interface should be here ... ?
- def fetch(self, *args, **kwargs) -> pd.DataFrame | dd.DataFrame:
- raise NotImplementedError("implement in subclass")
-
- def build(self, *args, **kwargs) -> pd.DataFrame | dd.DataFrame:
- raise NotImplementedError("implement in subclass")
-
-
-class MergeCollection(CollectionBase):
- """Mergers take instances of DFCollections, and/or other Mergers"""
-
- # In a merge, we can set offset = None which indicates that there is only 1
- # period/item where the range is 'start' until now.
- offset: Optional[str] = Field(default="72h")
- # In a merge, we can set start = None which indicates that there is only 1
- # period/item where the range is (now - offset) until now.
- start: Optional[AwareDatetimeISO] = Field(
- default=None,
- description="This is the starting point in which data will"
- " be retrieved in chunks from.",
- frozen=True,
- )
-
- merge_type: Optional[MergeType] = Field(default=None)
- group_by: Optional[str] = Field(default=None)
- grouped_key: Optional[str] = Field(default=None)
- collection_item_class: Type[MergeCollectionItem] = MergeCollectionItem
-
- @model_validator(mode="after")
- def check_start_and_offset_nullable(self) -> Self:
- if self.offset is None and self.start is None:
- raise AssertionError("cannot set both start and offset to None")
- return self
-
- @field_validator("merge_type")
- def check_merge_type(cls, merge_type, info: ValidationInfo):
- if merge_type is None:
- raise ValueError("Must explicitly provide a merge_type")
-
- if merge_type not in MergeTypeSchemas:
- raise ValueError("Must provide a supported merge_type")
-
- return merge_type
-
- # --- Properties ---
- @property
- def interval_start(self) -> Optional[datetime]:
- # if self.start is None and self.offset is set, the inferred start is (now - offset)
- if self.start is None:
- return datetime.now(tz=timezone.utc).replace(microsecond=0) - pd.Timedelta(
- self.offset
- )
- return self.start
-
- @property
- def items(self) -> List[MergeCollectionItem]:
- items = []
- for iv in self.interval_range:
- cm = self.collection_item_class(start=iv[0])
- cm._collection = self
- items.append(cm)
- return items
-
- @property
- def _schema(self) -> DataFrameSchema:
- return MergeTypeSchemas[self.merge_type]
-
- def signature(self) -> str:
- arr = [
- 1 if i.has_archive(include_empty=True) else 0
- for i in self.items
- if i.should_archive()
- ]
- repr_str = (
- f"path={self.archive_path.as_posix()}; "
- f"items={len(self.items)}; start={self.start} @ {self.offset}; {int(sum(arr) / len(arr) * 100)}% "
- f"archived"
- )
- res = f"{self.__repr_name__()}({repr_str})"
- return res
diff --git a/generalresearch/incite/mergers/base.py b/generalresearch/incite/mergers/base.py
new file mode 100644
index 0000000..a0874a7
--- /dev/null
+++ b/generalresearch/incite/mergers/base.py
@@ -0,0 +1,311 @@
+import os.path
+import subprocess
+from datetime import UTC, datetime
+from enum import StrEnum
+from sys import platform
+from typing import Self
+
+import dask.dataframe as dd
+import pandas as pd
+from dask.distributed import Client
+from pandera.pandas import DataFrameSchema
+from pydantic import Field, ValidationInfo, field_validator, model_validator
+
+from generalresearch.incite import LOG
+from generalresearch.incite.base import CollectionBase, CollectionItemBase
+from generalresearch.incite.schemas import PARTITION_ON
+from generalresearch.incite.schemas.mergers.foundations.enriched_session import (
+ EnrichedSessionSchema,
+)
+from generalresearch.incite.schemas.mergers.foundations.enriched_task_adjust import (
+ EnrichedTaskAdjustSchema,
+)
+from generalresearch.incite.schemas.mergers.foundations.enriched_wall import (
+ EnrichedWallSchema,
+)
+from generalresearch.incite.schemas.mergers.foundations.user_id_product import (
+ UserIdProductSchema,
+)
+from generalresearch.incite.schemas.mergers.pop_ledger import (
+ PopLedgerSchema,
+)
+from generalresearch.incite.schemas.mergers.ym_survey_wall import (
+ YMSurveyWallSchema,
+)
+from generalresearch.incite.schemas.mergers.ym_wall_summary import (
+ YMWallSummarySchema,
+)
+from generalresearch.models.custom_types import AwareDatetimeISO
+
+
+class MergeType(StrEnum):
+ TEST = "test"
+ YM_SURVEY_WALL = "ym_survey_wall"
+ YM_WALL_SUMMARY = "ym_wall_summary"
+
+ POP_LEDGER = "pop_ledger"
+
+ # --- Foundations ---
+ USER_ID_PRODUCT = "user_id_product"
+ ENRICHED_WALL = "enriched_wall"
+ ENRICHED_SESSION = "enriched_session"
+ ENRICHED_TASK_ADJUST = "enriched_task_adjust"
+
+
+MergeTypeSchemas = {
+ MergeType.YM_SURVEY_WALL: YMSurveyWallSchema,
+ MergeType.YM_WALL_SUMMARY: YMWallSummarySchema,
+ MergeType.POP_LEDGER: PopLedgerSchema,
+ # --- Foundations ---
+ MergeType.USER_ID_PRODUCT: UserIdProductSchema,
+ MergeType.ENRICHED_WALL: EnrichedWallSchema,
+ MergeType.ENRICHED_SESSION: EnrichedSessionSchema,
+ MergeType.ENRICHED_TASK_ADJUST: EnrichedTaskAdjustSchema,
+}
+
+
+class MergeCollectionItem(CollectionItemBase):
+ # --- Properties ---
+
+ @property
+ def finish(self) -> datetime:
+ # A MergeCollection can have offset = None
+ if self._collection.offset:
+ return (
+ pd.Timestamp(self.start) + pd.Timedelta(self._collection.offset)
+ ).to_pydatetime()
+ else:
+ return datetime.now(tz=UTC).replace(microsecond=0)
+
+ @property
+ def filename(self) -> str:
+ grouped_key = self._collection.grouped_key
+ offset = self._collection.offset
+ start = self.start.strftime("%Y-%m-%d-%H-%M-%S")
+ f = [self._collection.merge_type.name.lower()]
+ if offset:
+ f.append(offset)
+ if grouped_key:
+ f.append(grouped_key)
+ if self._collection.start is not None:
+ # This is a collection that is "looking back" 'offset' time (1 item).
+ f.append(start)
+ s = "-".join(f)
+ s += ".parquet"
+ return s
+
+ # --- ORM / Data handlers---
+ def to_dict(self, *args, **kwargs) -> dict:
+ res = self._to_dict()
+ res["group_by"] = self._collection.group_by
+ return res
+
+ def to_archive(
+ self,
+ client: Client,
+ ddf: dd.DataFrame,
+ is_partial: bool = False,
+ ) -> bool:
+ assert is_partial is False, "use to_archive_symlink"
+ return self._to_archive(client=client, ddf=ddf, client_resources=None)
+
+ def _to_archive(
+ self, client: Client, ddf: dd.DataFrame | None, client_resources=None
+ ) -> bool:
+ """
+ For archiving an item. Will write an empty file if ddf is empty.
+ This is NOT for writing partials.
+
+ :returns: bool (saved_successful)
+ """
+ if ddf is None:
+ return False
+
+ row_len: int = client.compute(collections=ddf.shape[0], sync=True)
+ assert row_len
+ assert row_len > 0, "empty ddf"
+
+ tmp_path = self.tmp_path()
+ schema = self._collection._schema
+ assert schema.metadata
+
+ partition = schema.metadata.get(PARTITION_ON)
+ f = ddf.to_parquet(
+ compute=False,
+ path=tmp_path,
+ partition_on=partition,
+ engine="pyarrow",
+ overwrite=True,
+ write_metadata_file=True,
+ compression="brotli",
+ )
+ client.compute(f, sync=True, priority=2, resources=client_resources)
+ assert not os.path.exists(self.path.as_posix()), (
+ f"already exits!: {self.path.as_posix()}"
+ )
+
+ if platform == "darwin":
+ subprocess.call(["mv", tmp_path.as_posix(), self.path.as_posix()])
+ else:
+ # -T will (should) cause the mv to fail if `path` wasn't successfully deleted
+ subprocess.call(["mv", "-T", tmp_path.as_posix(), self.path.as_posix()])
+ return True
+
+ def to_archive_symlink(
+ self,
+ client: Client,
+ ddf: dd.DataFrame,
+ is_partial: bool = False,
+ client_resources=None,
+ validate_after=True,
+ ) -> bool:
+ """
+ This differs from to_archive():
+ 1) to_parquet is run in this process. If the df is already
+ computed, there is no point in sending it to another worker
+ to write.
+
+ 2) symlink to next_numbered_path is created whether or not
+ is_partial (to_archive only does this on partials)
+
+ 3) we do not validate the written file. seems not useful to do
+ this, as the file will probably get overwritten on the next
+ loop anyway
+ """
+ path = self.partial_path if is_partial else self.path
+ next_numbered_path = self.next_numbered_path(path)
+ collection = self._collection
+ LOG.warning(f"{collection.merge_type.value}.to_archive_symlink()")
+
+ assert isinstance(ddf, dd.DataFrame), "must pass a dask df"
+
+ # We should validate before or after!!!
+ # _validate_df(self.compute(ddf), coll._schema)
+ target = (
+ next_numbered_path.name
+ ) # this is the symlink's target. it is a relative path (only the name)
+
+ schema = self._collection._schema
+ partition = schema.metadata.get(PARTITION_ON, None)
+ f = ddf.to_parquet(
+ compute=False,
+ path=next_numbered_path.as_posix(),
+ partition_on=partition,
+ engine="pyarrow",
+ overwrite=True,
+ write_metadata_file=True,
+ compression="brotli",
+ )
+ client.compute(f, sync=True, priority=2, resources=client_resources)
+
+ if os.path.exists(path.as_posix()) and not os.path.islink(path.as_posix()):
+ # This will fail when going from the old way to using symlinks,
+ # if self.path already exists and is a directory.
+ raise ValueError(
+ f"first time we run this, make sure the path doesnt exist: {path.as_posix()}"
+ )
+
+ if platform == "darwin":
+ subprocess.call(["ln", "-sfn", target, path.as_posix()])
+ else:
+ subprocess.call(["ln", "-sfnT", target, path.as_posix()])
+
+ if validate_after and not self.valid_archive(self.path):
+ LOG.error(f"{collection.merge_type.value} failed validation: {self.path}")
+ self.delete_archive(self.path)
+ return False
+ return True
+
+ # todo: unclear what the common interface should be here ... ?
+ def fetch(self, *args, **kwargs) -> pd.DataFrame | dd.DataFrame:
+ raise NotImplementedError("implement in subclass")
+
+ def build(self, *args, **kwargs) -> pd.DataFrame | dd.DataFrame:
+ raise NotImplementedError("implement in subclass")
+
+
+class MergeCollection(CollectionBase):
+ """Mergers take instances of DFCollections, and/or other Mergers"""
+
+ # In a merge, we can set offset = None which indicates that there is only 1
+ # period/item where the range is 'start' until now.
+ offset: str | None = Field(default="72h")
+ # In a merge, we can set start = None which indicates that there is only 1
+ # period/item where the range is (now - offset) until now.
+ start: AwareDatetimeISO | None = Field(
+ default=None,
+ description="This is the starting point in which data will"
+ " be retrieved in chunks from.",
+ frozen=True,
+ )
+
+ merge_type: MergeType | None = Field(default=None)
+ group_by: str | None = Field(default=None)
+ grouped_key: str | None = Field(default=None)
+ collection_item_class: type[MergeCollectionItem] = MergeCollectionItem
+
+ @model_validator(mode="after")
+ def check_model_after(self) -> Self:
+ if self.offset is None or self.start is None:
+ return self
+
+ offset_total_sec = pd.Timedelta(self.offset).total_seconds()
+ start_total_sec = (datetime.now(tz=UTC) - self.start).total_seconds()
+
+ if offset_total_sec > start_total_sec:
+ raise ValueError("Offset must be equal to, or smaller the start timestamp")
+
+ return self
+
+ @model_validator(mode="after")
+ def check_start_and_offset_nullable(self) -> Self:
+ if self.offset is None and self.start is None:
+ raise AssertionError("cannot set both start and offset to None")
+ return self
+
+ @field_validator("merge_type")
+ def check_merge_type(cls, merge_type: MergeType | None, info: ValidationInfo):
+ if merge_type is None:
+ raise ValueError("Must explicitly provide a merge_type")
+
+ if merge_type not in MergeTypeSchemas:
+ raise ValueError("Must provide a supported merge_type")
+
+ return merge_type
+
+ # --- Properties ---
+ @property
+ def interval_start(self) -> datetime | None:
+ # if self.start is None and self.offset is set, the inferred start is (now - offset)
+ if self.start is None:
+ return datetime.now(tz=UTC).replace(microsecond=0) - pd.Timedelta(
+ self.offset
+ )
+ return self.start
+
+ @property
+ def items(self) -> list[MergeCollectionItem]:
+ items = []
+ for iv in self.interval_range:
+ cm = self.collection_item_class(start=iv[0])
+ cm._collection = self
+ items.append(cm)
+ return items
+
+ @property
+ def _schema(self) -> DataFrameSchema:
+ return MergeTypeSchemas[self.merge_type]
+
+ def signature(self) -> str:
+ arr = [
+ 1 if i.has_archive(include_empty=True) else 0
+ for i in self.items
+ if i.should_archive()
+ ]
+ repr_str = (
+ f"path={self.archive_path.as_posix()}; "
+ f"items={len(self.items)}; start={self.start} @ {self.offset}; {int(sum(arr) / len(arr) * 100)}% "
+ f"archived"
+ )
+ res = f"{self.__repr_name__()}({repr_str})"
+ return res
diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py
index f7a45a8..561e759 100644
--- a/generalresearch/incite/mergers/foundations/__init__.py
+++ b/generalresearch/incite/mergers/foundations/__init__.py
@@ -68,11 +68,9 @@ def lookup_product_and_team_id(
assert len(user_ids) <= 1000, "you should chunk this bro"
res: list[dict[str, Any]] = []
- with pg_config.make_connection() as conn:
- try:
- with conn.cursor() as c:
- c.execute(
- query="""
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ query="""
SELECT u.id AS user_id,
u.product_id,
bp.team_id
@@ -81,13 +79,9 @@ def lookup_product_and_team_id(
ON bp.id = u.product_id
WHERE u.id = ANY(%s);
""",
- params=[list(user_ids)],
- )
- res.extend(c.fetchall())
-
- except Exception as e:
- LOG.exception(f"lookup_product_and_team_id: {e}")
- raise
+ params=[list(user_ids)],
+ )
+ res.extend(c.fetchall())
return res
@@ -117,7 +111,7 @@ def annotate_product_and_team_id(
try:
with conn.cursor() as c:
c.execute(
- query=f"""
+ query="""
SELECT u.id AS user_id, u.product_id,
bp.team_id
FROM thl_user u
diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py
index a368e6c..cea6501 100644
--- a/generalresearch/incite/mergers/foundations/enriched_session.py
+++ b/generalresearch/incite/mergers/foundations/enriched_session.py
@@ -1,20 +1,20 @@
from __future__ import annotations
-import logging
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Literal
import dask.dataframe as dd
import pandas as pd
+from dask.distributed import Client as DaskClient
from dask.distributed import as_completed
-from distributed import Client
from more_itertools import chunked, flatten
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import (
SessionDFCollection,
WallDFCollection,
)
-from generalresearch.incite.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -30,13 +30,11 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_session import
EnrichedSessionSchema,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
if TYPE_CHECKING:
from generalresearch.models.admin.request import ReportRequest
-
-LOG = logging.getLogger("incite")
+ from generalresearch.models.thl.user import User
class EnrichedSessionMergeItem(MergeCollectionItem):
@@ -46,7 +44,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
session_coll: SessionDFCollection,
wall_coll: WallDFCollection,
pg_config: PostgresConfig,
- client: Client | None = None,
+ client: DaskClient | None = None,
client_resources: dict[str, Any] | None = None,
) -> None:
@@ -60,17 +58,17 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
return
# --- Session ---
- LOG.warning(f"EnrichedSessionMergeItem: get session_collection")
+ LOG.warning("EnrichedSessionMergeItem: get session_collection")
session_items = [w for w in session_coll.items if w.interval.overlaps(ir)]
if len(session_items) == 0:
- LOG.warning(f"EnrichedSessionMergeItem: no session items. set_empty.")
+ LOG.warning("EnrichedSessionMergeItem: no session items. set_empty.")
if self.should_archive():
self.set_empty()
return
if not (
session_items[-1].has_partial_archive() or session_items[-1].has_archive()
):
- LOG.warning(f"EnrichedSessionMergeItem: session isn't updated!")
+ LOG.warning("EnrichedSessionMergeItem: session isn't updated!")
return
sddf = session_coll.ddf(
@@ -81,7 +79,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
)
# --- Walls ---
- LOG.warning(f"EnrichedSessionMergeItem: merge wall_collection")
+ LOG.warning("EnrichedSessionMergeItem: merge wall_collection")
wall_items = [
w
for w in wall_coll.items
@@ -95,7 +93,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
]
if len(wall_items) == 0:
- LOG.error(f"EnrichedSessionMergeItem: no wall items")
+ LOG.error("EnrichedSessionMergeItem: no wall items")
return
wddf = wall_coll.ddf(
@@ -140,9 +138,9 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
try:
results = client.gather(list(futures))
- except Exception as e:
+ except Exception:
client.cancel(futures, asynchronous=False, force=True)
- raise e
+ raise
dfp = pd.DataFrame(
list(flatten(results)), columns=["user_id", "product_id", "team_id"]
@@ -154,18 +152,13 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
df = df[df["started"].between(start, end)]
is_missing = df[["product_id"]].isna().sum().sum() > 0
- session_is_partial = any([w.should_archive() is False for w in session_items])
+ session_is_partial = any(w.should_archive() is False for w in session_items)
session_is_missing = any(
- [
- w.should_archive() is True and w.has_archive() is False
- for w in session_items
- ]
+ w.should_archive() is True and w.has_archive() is False
+ for w in session_items
)
wall_is_missing = any(
- [
- w.should_archive() is True and w.has_archive() is False
- for w in wall_items
- ]
+ w.should_archive() is True and w.has_archive() is False for w in wall_items
)
is_partial = (
is_missing or session_is_partial or session_is_missing or wall_is_missing
@@ -192,7 +185,6 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
client,
ddf=ddf,
is_partial=False,
- client_resources=client_resources,
)
@@ -203,7 +195,7 @@ class EnrichedSessionMerge(MergeCollection):
def build(
self,
- client: Client,
+ client: DaskClient,
session_coll: SessionDFCollection,
wall_coll: WallDFCollection,
pg_config: PostgresConfig,
@@ -232,7 +224,7 @@ class EnrichedSessionMerge(MergeCollection):
def to_admin_response(
self,
rr: ReportRequest,
- client: Client,
+ client: DaskClient,
product_ids: list[UUIDStr] | None = None,
user: User | None = None,
) -> pd.DataFrame:
@@ -243,6 +235,7 @@ class EnrichedSessionMerge(MergeCollection):
filters = []
if user:
+ assert product_ids
assert (
len(product_ids) <= 1
), "Can't search more than 1 Product ID for a specific User"
diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
index f3ab8d8..be491b7 100644
--- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
+++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
@@ -1,6 +1,5 @@
from __future__ import annotations
-import logging
from typing import Any, Literal
import dask.dataframe as dd
@@ -8,10 +7,12 @@ import pandas as pd
from distributed import Client
from sentry_sdk import capture_exception
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import (
TaskAdjustmentDFCollection,
)
-from generalresearch.incite.mergers import (
+from generalresearch.incite.exceptions import BuildError, BuildItemsError
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -27,8 +28,6 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_task_adjust imp
)
from generalresearch.pg_helper import PostgresConfig
-LOG = logging.getLogger("incite")
-
class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
"""Because a single wall event can have multiple "alerted" times,
@@ -55,13 +54,13 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
LOG.warning(f"EnrichedReconMergeItem.build({ir})")
# --- Task Adjustments ---
- LOG.warning(f"EnrichedReconMergeItem: get session_collection")
+ LOG.warning("EnrichedReconMergeItem: get session_collection")
task_adj_coll_items = [
w for w in task_adj_coll.items if w.interval.overlaps(ir)
]
if len(task_adj_coll_items) == 0:
- raise Exception("TaskAdjColl item collection failed")
+ raise BuildItemsError("TaskAdjColl item collection failed")
ddf: dd.DataFrame | None = task_adj_coll.ddf(
items=task_adj_coll_items,
@@ -83,6 +82,8 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
("started", "<", end),
],
)
+
+ assert isinstance(ddf, pd.DataFrame)
# Naked compute... don't log
# LOG.info(f"TaskAdjustmentDetailMergeCollectionItem.rows: {len(ddf.index)}")
@@ -91,7 +92,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
ew_items = [ew for ew in enriched_wall.items if ew.interval.overlaps(ir)]
if len(ew_items) == 0:
- raise Exception(
+ raise BuildItemsError(
"EnrichedWall item collection failed for EnrichedTaskAdjColl"
)
@@ -209,6 +210,5 @@ class EnrichedTaskAdjustMerge(MergeCollection):
enriched_wall=enriched_wall,
pg_config=pg_config,
)
- except (Exception,) as e:
+ except BuildError as e:
capture_exception(error=e)
- pass
diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py
index 5a7dd2b..70139c2 100644
--- a/generalresearch/incite/mergers/foundations/enriched_wall.py
+++ b/generalresearch/incite/mergers/foundations/enriched_wall.py
@@ -1,6 +1,5 @@
from __future__ import annotations
-import logging
from datetime import timedelta
from typing import TYPE_CHECKING, Any, Literal
@@ -8,11 +7,12 @@ import dask.dataframe as dd
import pandas as pd
from distributed import Client
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import (
SessionDFCollection,
WallDFCollection,
)
-from generalresearch.incite.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -25,13 +25,11 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_wall import (
EnrichedWallSchema,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
if TYPE_CHECKING:
from generalresearch.models.admin.request import ReportRequest
-
-LOG = logging.getLogger("incite")
+ from generalresearch.models.thl.user import User
class EnrichedWallMergeItem(MergeCollectionItem):
@@ -42,7 +40,6 @@ class EnrichedWallMergeItem(MergeCollectionItem):
session_coll: SessionDFCollection,
pg_config: PostgresConfig,
client: Client | None = None,
- client_resources: dict[str, Any] | None = None,
) -> None:
ir: pd.Interval = self.interval
@@ -55,10 +52,10 @@ class EnrichedWallMergeItem(MergeCollectionItem):
return
# --- Wall ---
- LOG.warning(f"EnrichedWallMergeItem: get wall_collection")
+ LOG.warning("EnrichedWallMergeItem: get wall_collection")
wall_items = [w for w in wall_coll.items if w.interval.overlaps(ir)]
if len(wall_items) == 0:
- LOG.warning(f"EnrichedWallMergeItem: no wall items. set_empty.")
+ LOG.warning("EnrichedWallMergeItem: no wall items. set_empty.")
if self.should_archive():
self.set_empty()
return
@@ -93,7 +90,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
wdf = wdf.reset_index(drop=False)
# --- Sessions ---
- LOG.warning(f"EnrichedWallMergeItem: merge session_collection")
+ LOG.warning("EnrichedWallMergeItem: merge session_collection")
session_items = [
s
for s in session_coll.items
@@ -107,7 +104,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
]
if len(session_items) == 0:
- LOG.error(f"EnrichedWallMergeItem: no session items. breaking early.")
+ LOG.error("EnrichedWallMergeItem: no session items. breaking early.")
return
sdf = session_coll.ddf(
@@ -148,7 +145,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
is_missing = False
df = df.dropna(subset=["product_id", "session_id"], how="any")
- wall_is_partial = any([w.should_archive() is False for w in wall_items])
+ wall_is_partial = any(w.should_archive() is False for w in wall_items)
is_partial = is_missing or wall_is_partial
# Lots of downstream issues with this...
@@ -162,7 +159,6 @@ class EnrichedWallMergeItem(MergeCollectionItem):
ddf=ddf,
is_partial=True,
validate_after=False,
- client_resources=client_resources,
)
else:
df = self.validate_df(df=df)
@@ -171,7 +167,6 @@ class EnrichedWallMergeItem(MergeCollectionItem):
client,
ddf=ddf,
is_partial=False,
- client_resources=client_resources,
)
diff --git a/generalresearch/incite/mergers/foundations/user_id_product.py b/generalresearch/incite/mergers/foundations/user_id_product.py
index 863a741..e682b7f 100644
--- a/generalresearch/incite/mergers/foundations/user_id_product.py
+++ b/generalresearch/incite/mergers/foundations/user_id_product.py
@@ -1,19 +1,17 @@
from __future__ import annotations
-import logging
from typing import Any, Literal
from distributed import Client
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import UserDFCollection
-from generalresearch.incite.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
)
-LOG = logging.getLogger("incite")
-
class UserIdProductMergeItem(MergeCollectionItem):
diff --git a/generalresearch/incite/mergers/pop_ledger.py b/generalresearch/incite/mergers/pop_ledger.py
index b32503c..54f1b7e 100644
--- a/generalresearch/incite/mergers/pop_ledger.py
+++ b/generalresearch/incite/mergers/pop_ledger.py
@@ -1,6 +1,5 @@
from __future__ import annotations
-import logging
from typing import Any, Literal
import dask.dataframe as dd
@@ -8,8 +7,9 @@ import pandas as pd
from distributed import Client
from more_itertools import flatten
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import LedgerDFCollection
-from generalresearch.incite.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -17,8 +17,6 @@ from generalresearch.incite.mergers import (
from generalresearch.incite.schemas.mergers.pop_ledger import PopLedgerSchema
from generalresearch.models.thl.ledger import Direction, TransactionType
-LOG = logging.getLogger("incite")
-
class PopLedgerMergeItem(MergeCollectionItem):
diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py
index c060aae..8b66eb5 100644
--- a/generalresearch/incite/mergers/ym_survey_wall.py
+++ b/generalresearch/incite/mergers/ym_survey_wall.py
@@ -1,6 +1,5 @@
from __future__ import annotations
-import logging
from datetime import timedelta
from typing import Any, Literal
@@ -9,8 +8,10 @@ import pandas as pd
from distributed import Client
from sentry_sdk import capture_exception
+from generalresearch.incite import LOG
from generalresearch.incite.collections.thl_web import WallDFCollection
-from generalresearch.incite.mergers import (
+from generalresearch.incite.exceptions import BuildError
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -23,8 +24,6 @@ from generalresearch.incite.schemas.mergers.ym_survey_wall import (
)
from generalresearch.models.custom_types import AwareDatetimeISO
-LOG = logging.getLogger("incite")
-
class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
@@ -38,7 +37,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
LOG.info(f"YMSurveyWallMerge.build({self.start=}, {self.finish=})")
ir: pd.Interval = self.interval
start, _ = self.start, self.finish
- ddf = wall_coll.ddf(
+ ddf: dd.DataFrame | None = wall_coll.ddf(
items=wall_coll.get_items(start),
force_rr_latest=False,
include_partial=True,
@@ -60,9 +59,10 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
],
filters=[("started", ">=", start)],
)
+ assert isinstance(ddf, dd.DataFrame)
ddf = ddf[ddf["started"] > start]
- LOG.warning(f"YMSurveyWallMerge: merge session_collection")
+ LOG.warning("YMSurveyWallMerge: merge session_collection")
session_items = [
s
for s in enriched_session.items
@@ -91,6 +91,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
)
df = client.compute(ddf, resources=client_resources, sync=True)
+ assert isinstance(df, pd.DataFrame)
df["elapsed"] = (df["finished"] - df["started"]).dt.total_seconds()
df["elapsed"] = df["elapsed"].round().astype("Int64")
df = df.drop(columns={"finished", "payout"}, errors="ignore")
@@ -98,18 +99,16 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
df.dropna(subset="product_id", how="any", inplace=True)
df.sort_values(by="started", inplace=True)
- LOG.debug(f"YMSurveyWallMerge.build() validation")
+ LOG.debug("YMSurveyWallMerge.build() validation")
df = self.validate_df(df=df)
if df is not None:
ddf = dd.from_pandas(df, npartitions=4)
- LOG.info(f"YMSurveyWallMerge.build() saving")
+ LOG.info("YMSurveyWallMerge.build() saving")
self.to_archive_symlink(client=client, ddf=ddf)
else:
LOG.warning("YMSurveyWallMerge failed validation")
- return None
-
class YMSurveyWallMerge(MergeCollection):
merge_type: Literal[MergeType.YM_SURVEY_WALL] = MergeType.YM_SURVEY_WALL
@@ -144,8 +143,7 @@ class YMSurveyWallMerge(MergeCollection):
wall_coll=wall_coll,
enriched_session=enriched_session,
)
- except (Exception,) as e:
+ except BuildError as e:
capture_exception(error=e)
- pass
item.delete_dangling_partials(keep_latest=2, target_path=item.path)
diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py
index 2f5995f..a01443b 100644
--- a/generalresearch/incite/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/mergers/ym_wall_summary.py
@@ -1,7 +1,7 @@
from __future__ import annotations
from datetime import datetime, time, timedelta
-from typing import Literal, Type
+from typing import Literal
import dask.dataframe as dd
import pandas as pd
@@ -12,7 +12,8 @@ from generalresearch.incite.collections.thl_web import (
SessionDFCollection,
WallDFCollection,
)
-from generalresearch.incite.mergers import (
+from generalresearch.incite.exceptions import FetchError
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeCollectionItem,
MergeType,
@@ -41,6 +42,8 @@ class YMWallSummaryMergeItem(MergeCollectionItem):
ddf = wall_collection.ddf(
items=wall_items, force_rr_latest=False, include_partial=True
)
+ assert isinstance(ddf, pd.DataFrame)
+
ddf = ddf[ddf["started"].between(start, end)]
# Then we need the sessions for these wall events. They'll have started
@@ -82,18 +85,20 @@ class YMWallSummaryMergeItem(MergeCollectionItem):
class YMWallSummaryMerge(MergeCollection):
merge_type: Literal[MergeType.YM_WALL_SUMMARY] = MergeType.YM_WALL_SUMMARY
_schema = YMWallSummarySchema
- collection_item_class: Type[YMWallSummaryMergeItem] = YMWallSummaryMergeItem
+ collection_item_class: type[YMWallSummaryMergeItem] = YMWallSummaryMergeItem
items: list[YMWallSummaryMergeItem] = Field(default_factory=list)
@field_validator("offset")
def check_offset_ym_wall_summary(cls, v: str | None):
# the offset MUST be on a whole day, no hourly
+ assert v
assert v.endswith("D"), "offset must be in days"
return v
@field_validator("start")
def check_start_ym_wall_summary(cls, v: datetime | None):
# the start MUST be start on midnight exactly
+ assert v
assert v.time() == time(0, 0, 0, 0), "start must no have a time component"
return v
@@ -117,9 +122,8 @@ class YMWallSummaryMerge(MergeCollection):
# item every time build is run even if it isn't closed
# if item.should_archive():
item.fetch(wall_collection, session_collection, user_id_product)
- except (Exception,) as e:
+ except FetchError as e:
capture_exception(e)
- pass
@staticmethod
def build_groupbys(df: pd.DataFrame) -> pd.DataFrame:
@@ -177,10 +181,10 @@ class YMWallSummaryMerge(MergeCollection):
# df.to_parquet(str(self.archive_path) + ".all.parquet")
pass
- def get_counts(self, product_id):
+ def get_counts(self, product_id: str):
# examples...
product_id = ""
- df = dd.read_parquet(
+ _ = dd.read_parquet(
str(self.archive_path) + ".all.parquet",
filters=[
("product_id", "=", product_id),
@@ -188,7 +192,7 @@ class YMWallSummaryMerge(MergeCollection):
],
).compute()
country_iso = "de"
- df = dd.read_parquet(
+ _ = dd.read_parquet(
str(self.archive_path) + ".all.parquet",
filters=[
("product_id", "=", product_id),
diff --git a/generalresearch/incite/schemas/__init__.py b/generalresearch/incite/schemas/__init__.py
index c0000d1..9d24f17 100644
--- a/generalresearch/incite/schemas/__init__.py
+++ b/generalresearch/incite/schemas/__init__.py
@@ -1,5 +1,3 @@
-from typing import List
-
import pandas as pd
import pandera.pandas as pa
@@ -11,8 +9,8 @@ ARCHIVE_AFTER = "archive_after"
PARTITION_ON = "partition_on"
-def empty_dataframe_from_schema(schema: pa.DataFrameSchema) -> "pd.DataFrame":
- index_names: List[str] = schema.index.names
+def empty_dataframe_from_schema(schema: pa.DataFrameSchema) -> pd.DataFrame:
+ index_names: list[str] = schema.index.names
columns = set(schema.dtypes.keys())
if len(index_names) > 1:
diff --git a/generalresearch/incite/schemas/admin_responses.py b/generalresearch/incite/schemas/admin_responses.py
index bb2852f..e65c6e2 100644
--- a/generalresearch/incite/schemas/admin_responses.py
+++ b/generalresearch/incite/schemas/admin_responses.py
@@ -1,5 +1,9 @@
-from datetime import datetime
+from __future__ import annotations
+from collections.abc import Callable
+from datetime import UTC, datetime
+
+import pandas as pd
from pandera.pandas import (
Check,
Column,
@@ -14,6 +18,14 @@ BIG_INT32 = 2_147_483_647
SIX_HOUR_SECONDS = 6 * 60 * 6
ROUNDING = 2
+_fillna: Callable[[pd.Series], pd.Series] = lambda s: s.fillna(value=0.00)
+_clip: Callable[[pd.Series], pd.Series] = lambda s: s.clip(
+ lower=0, upper=SIX_HOUR_SECONDS
+)
+_round: Callable[[pd.Series], pd.Series] = lambda s: s.round(decimals=ROUNDING)
+_tz_localize_none: Callable[[pd.Series], pd.Series] = lambda i: i.dt.tz_localize(None)
+
+
AdminPOPSchema = DataFrameSchema(
# Generic: used for Session or Wall
index=MultiIndex(
@@ -25,10 +37,15 @@ AdminPOPSchema = DataFrameSchema(
Index(
name="index0",
dtype=Timestamp,
- parsers=[Parser(lambda i: i.dt.tz_localize(None))],
+ parsers=[Parser(_tz_localize_none)],
checks=[
Check.less_than(
- max_value=datetime(year=datetime.now().year + 1, month=1, day=1)
+ max_value=datetime(
+ year=datetime.now(tz=UTC).year + 1,
+ month=1,
+ day=1,
+ tzinfo=UTC,
+ )
)
],
),
@@ -44,85 +61,85 @@ AdminPOPSchema = DataFrameSchema(
"elapsed_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.clip(lower=0, upper=SIX_HOUR_SECONDS)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_clip),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=SIX_HOUR_SECONDS),
),
"elapsed_total": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"payout_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=100),
),
"payout_total": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"entrances": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"completes": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"users": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"conversion": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0.00, max_value=1.00),
),
"epc": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=100),
),
"eph": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"cpc": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=250),
),
@@ -140,21 +157,21 @@ AdminPOPWallSchema = DataFrameSchema(
"buyers": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"surveys": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"sessions": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
@@ -169,15 +186,15 @@ AdminPOPSessionSchema = DataFrameSchema(
"attempts_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=25),
),
"attempts_total": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
index cc909f6..ac9a35a 100644
--- a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
+++ b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
@@ -1,19 +1,17 @@
-from typing import Set
-
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.incite.schemas import ARCHIVE_AFTER, ORDER_KEY
from generalresearch.incite.schemas.thl_web import THLTaskAdjustmentSchema
from generalresearch.locales import Localelator
-from generalresearch.models import DeviceType, Source
+from generalresearch.models.definitions import DeviceType, Source
from generalresearch.models.thl.definitions import (
WallAdjustedStatus,
)
thl_task_adj_columns = THLTaskAdjustmentSchema.columns.copy()
-COUNTRY_ISOS: Set[str] = Localelator().get_all_countries()
+COUNTRY_ISOS: set[str] = Localelator().get_all_countries()
kosovo = "xk"
COUNTRY_ISOS.add(kosovo)
BIGINT = 9223372036854775807
diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py b/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py
index 1443f28..71d0eab 100644
--- a/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py
+++ b/generalresearch/incite/schemas/mergers/foundations/enriched_wall.py
@@ -5,7 +5,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.incite.schemas import ARCHIVE_AFTER, PARTITION_ON
from generalresearch.locales import Localelator
-from generalresearch.models import DeviceType, Source
+from generalresearch.models.definitions import DeviceType, Source
from generalresearch.models.thl.definitions import (
ReportValue,
Status,
diff --git a/generalresearch/incite/schemas/mergers/foundations/user_id_product.py b/generalresearch/incite/schemas/mergers/foundations/user_id_product.py
index a07ca42..9e73aed 100644
--- a/generalresearch/incite/schemas/mergers/foundations/user_id_product.py
+++ b/generalresearch/incite/schemas/mergers/foundations/user_id_product.py
@@ -4,7 +4,7 @@ from pandera.pandas import Category, Check, Column, DataFrameSchema, Index
from generalresearch.incite.schemas import ARCHIVE_AFTER
-BIGINT = 9223372036854775807
+BIGINT = 9_223_372_036_854_775_807
UserIdIndex = Index(
name="id",
diff --git a/generalresearch/incite/schemas/mergers/pop_ledger.py b/generalresearch/incite/schemas/mergers/pop_ledger.py
index fdb1b05..8452eb9 100644
--- a/generalresearch/incite/schemas/mergers/pop_ledger.py
+++ b/generalresearch/incite/schemas/mergers/pop_ledger.py
@@ -1,3 +1,6 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import timedelta
import pandas as pd
@@ -22,6 +25,11 @@ from generalresearch.models.thl.ledger import Direction, TransactionType
# If an amount is "very" large, something is def wrong. Defining "very" somewhat arbitrarily here.
SUSPICIOUSLY_LARGE_NUMBER = (2**32 / 2) - 1 # 2147483647
+_tz_min_freq: Callable[[pd.Series], pd.Series] = lambda i: (i.dt.second == 0) & (
+ i.dt.microsecond == 0
+)
+
+
NonNegativeAmount = Column(
dtype="Int32",
nullable=True,
@@ -49,7 +57,7 @@ PopLedgerSchema = DataFrameSchema(
| {
"time_idx": Column(
dtype=pd.DatetimeTZDtype(tz="UTC"),
- checks=Check(lambda x: (x.dt.second == 0) & (x.dt.microsecond == 0)),
+ checks=Check(_tz_min_freq),
nullable=False,
),
"account_id": TxSchema.columns["account_id"],
diff --git a/generalresearch/incite/schemas/mergers/ym_wall_summary.py b/generalresearch/incite/schemas/mergers/ym_wall_summary.py
index 16cfc2f..737b925 100644
--- a/generalresearch/incite/schemas/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/schemas/mergers/ym_wall_summary.py
@@ -6,7 +6,7 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.incite.schemas import ARCHIVE_AFTER
from generalresearch.locales import Localelator
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
COUNTRY_ISOS: set[str] = Localelator().get_all_countries()
kosovo = "xk"
diff --git a/generalresearch/incite/schemas/thl_web.py b/generalresearch/incite/schemas/thl_web.py
index 5073a18..36ee8e9 100644
--- a/generalresearch/incite/schemas/thl_web.py
+++ b/generalresearch/incite/schemas/thl_web.py
@@ -1,11 +1,12 @@
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
import pandas as pd
+from grip_client.enums import AccessType
from pandera.pandas import Check, Column, DataFrameSchema, Index, MultiIndex
from generalresearch.incite.schemas import ARCHIVE_AFTER, ORDER_KEY
from generalresearch.locales import Localelator
-from generalresearch.models import DeviceType, Source
+from generalresearch.models.definitions import DeviceType, Source
from generalresearch.models.thl.definitions import (
ReportValue,
SessionAdjustedStatus,
@@ -16,7 +17,6 @@ from generalresearch.models.thl.definitions import (
WallStatusCode2,
)
from generalresearch.models.thl.ledger import TransactionMetadataColumns
-from generalresearch.models.thl.maxmind.definitions import UserType
IP_REGEX_PATTERN = (
r"^((([0-9]|[1-9][0-9]|1[0-9]{2}|2[0-4][0-9]|25[0-5])\.){3}([0-9]|[1-9][0-9]|1[0-9]{2}|2[0-4]["
@@ -105,7 +105,7 @@ THLWallSchema = DataFrameSchema(
),
"started": Column(
dtype=pd.DatetimeTZDtype(tz="UTC"),
- checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))],
+ checks=[Check(lambda x: x < datetime.now(tz=UTC))],
nullable=False,
),
"session_id": Column(
@@ -205,12 +205,12 @@ THLSessionSchema = DataFrameSchema(
),
"started": Column(
dtype=pd.DatetimeTZDtype(tz="UTC"),
- checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))],
+ checks=[Check(lambda x: x < datetime.now(tz=UTC))],
nullable=True,
),
"finished": Column(
dtype=pd.DatetimeTZDtype(tz="UTC"),
- checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))],
+ checks=[Check(lambda x: x < datetime.now(tz=UTC))],
nullable=True,
),
"loi_min": Column(dtype="Int64", nullable=True),
@@ -392,7 +392,7 @@ THLIPInfoSchema = DataFrameSchema(
dtype=str,
checks=[
Check.str_length(min_value=3, max_value=255),
- Check.isin([e.value for e in UserType]),
+ Check.isin([e.value for e in AccessType]),
],
nullable=True,
),
@@ -450,7 +450,7 @@ THLTaskAdjustmentSchema = DataFrameSchema(
),
"started": Column(
dtype=pd.DatetimeTZDtype(tz="UTC"),
- checks=[Check(lambda x: x < datetime.now(tz=timezone.utc))],
+ checks=[Check(lambda x: x < datetime.now(tz=UTC))],
),
"source": Column(
dtype=str,
diff --git a/generalresearch/locales/__init__.py b/generalresearch/locales/__init__.py
index 88b72e6..4bb10d0 100644
--- a/generalresearch/locales/__init__.py
+++ b/generalresearch/locales/__init__.py
@@ -12,7 +12,6 @@ https://en.wikipedia.org/wiki/List_of_ISO_639-1_codes
import json
import pkgutil
-from typing import Set
class Localelator:
@@ -20,10 +19,6 @@ class Localelator:
EVERYTHING IS LOWERCASE!!! (except this comment)
"""
- lang_alpha2_to_alpha3b = dict()
- lang_alpha3_to_alpha3b = dict()
- languages = set()
-
def __init__(self):
d = json.loads(pkgutil.get_data(__name__, "iso639-3.json"))
self.lang_alpha2_to_alpha3b = {x["alpha_2"]: x["alpha_3b"] for x in d}
@@ -43,11 +38,11 @@ class Localelator:
pkgutil.get_data(__name__, "country_default_lang.json")
)
- def get_all_languages(self) -> Set[str]:
+ def get_all_languages(self) -> set[str]:
# returns only the ISO 639-2/B (three-letter codes)
return set(self.lang_alpha2_to_alpha3b.values())
- def get_all_countries(self) -> Set[str]:
+ def get_all_countries(self) -> set[str]:
# returns only the ISO 3166-1 alpha-2 (two-letter codes)
return set(self.country_alpha3_to_alpha2.values())
diff --git a/generalresearch/locales/setup_json.py b/generalresearch/locales/setup_json.py
index 356084a..71beb09 100644
--- a/generalresearch/locales/setup_json.py
+++ b/generalresearch/locales/setup_json.py
@@ -1,61 +1,62 @@
-import json
-
-
-def country_default_lang():
- """
- Some marketplaces have no language specified. Surveys are in the "default
- language for that country", whatever that means. This helper is meant to
- provide a reasonable guess as to what language it is.
-
- Derived from: http://download.geonames.org/export/dump/countryInfo.txt
- """
- raise ValueError("no need to run this, I already ran it.")
- import pandas as pd
- from generalresearch.locales import Localelator
-
- l = Localelator()
-
- df = pd.read_csv(
- "http://download.geonames.org/export/dump/countryInfo.txt",
- sep="\t",
- skiprows=49,
- )
- df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0]
- df.default_lang = df.default_lang.fillna("en")
- df.default_lang = df.default_lang.map(
- lambda x: l.get_language_iso(x) if x in l.languages else "eng"
- )
- df["#ISO"] = df["#ISO"].str.lower()
- df["country_iso"] = df["#ISO"].map(
- lambda x: l.get_country_iso(x) if x in l.countries else None
- )
- df = df[df.country_iso.notnull()]
- d = df.set_index("country_iso").default_lang.to_dict()
- with open("country_default_lang.json", "w") as f:
- json.dump(d, f, indent=2)
- return d
-
-
-def setup_json():
- # pycountry is 30mb, which makes using this package on AWS lambda problematic.
- # These JSONs are stolen from pycountry and adapted.
-
- raise ValueError("no need to run this, I already ran it.")
-
- # languages
- d = json.load(open("iso639-3.json"))
- d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x]
- for x in d["639-3"]:
- x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"]
- del x["scope"]
- del x["type"]
- with open("iso639-3.json", "w") as f:
- json.dump(d["639-3"], f, indent=2)
-
- # countries
- d = json.load(open("iso3166-1.json"))["3166-1"]
- for x in d:
- x["alpha_2"] = x["alpha_2"].lower()
- x["alpha_3"] = x["alpha_3"].lower()
- with open("iso3166-1.json", "w") as f:
- json.dump(d, f, indent=2)
+# import json
+
+
+# def country_default_lang():
+# """
+# Some marketplaces have no language specified. Surveys are in the "default
+# language for that country", whatever that means. This helper is meant to
+# provide a reasonable guess as to what language it is.
+
+# Derived from: http://download.geonames.org/export/dump/countryInfo.txt
+# """
+# raise ValueError("no need to run this, I already ran it.")
+# import pandas as pd
+
+# from generalresearch.locales import Localelator
+
+# l = Localelator()
+
+# df = pd.read_csv(
+# "http://download.geonames.org/export/dump/countryInfo.txt",
+# sep="\t",
+# skiprows=49,
+# )
+# df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0]
+# df.default_lang = df.default_lang.fillna("en")
+# df.default_lang = df.default_lang.map(
+# lambda x: l.get_language_iso(x) if x in l.languages else "eng"
+# )
+# df["#ISO"] = df["#ISO"].str.lower()
+# df["country_iso"] = df["#ISO"].map(
+# lambda x: l.get_country_iso(x) if x in l.countries else None
+# )
+# df = df[df.country_iso.notnull()]
+# d = df.set_index("country_iso").default_lang.to_dict()
+# with open("country_default_lang.json", "w") as f:
+# json.dump(d, f, indent=2)
+# return d
+
+
+# def setup_json():
+# # pycountry is 30mb, which makes using this package on AWS lambda problematic.
+# # These JSONs are stolen from pycountry and adapted.
+
+# raise ValueError("no need to run this, I already ran it.")
+
+# # languages
+# d = json.load(open("iso639-3.json"))
+# d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x]
+# for x in d["639-3"]:
+# x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"]
+# del x["scope"]
+# del x["type"]
+# with open("iso639-3.json", "w") as f:
+# json.dump(d["639-3"], f, indent=2)
+
+# # countries
+# d = json.load(open("iso3166-1.json"))["3166-1"]
+# for x in d:
+# x["alpha_2"] = x["alpha_2"].lower()
+# x["alpha_3"] = x["alpha_3"].lower()
+# with open("iso3166-1.json", "w") as f:
+# json.dump(d, f, indent=2)
diff --git a/generalresearch/locales/timezone.py b/generalresearch/locales/timezone.py
index 50d539d..810dba3 100644
--- a/generalresearch/locales/timezone.py
+++ b/generalresearch/locales/timezone.py
@@ -1,9 +1,8 @@
-from typing import Optional
from pytz import country_timezones
-def get_default_timezone(country_iso: str) -> Optional[str]:
+def get_default_timezone(country_iso: str) -> str | None:
# to list all:
# from pytz import country_names, country_timezones
# [country_timezones.get(country) for country in country_names]
@@ -72,6 +71,6 @@ country_default_locale = {
}
-def get_default_locale(country_iso: str) -> Optional[str]:
+def get_default_locale(country_iso: str) -> str | None:
# todo: "https://cdn.simplelocalize.io/public/v1/locales" to fill in the rest?
return country_default_locale.get(country_iso, None)
diff --git a/generalresearch/logging.py b/generalresearch/logging.py
index 9b72e0b..40f2174 100644
--- a/generalresearch/logging.py
+++ b/generalresearch/logging.py
@@ -1,6 +1,7 @@
import decimal
import json
from datetime import date
+from typing import Any
class ThlJsonEncoder(json.JSONEncoder):
@@ -11,11 +12,11 @@ class ThlJsonEncoder(json.JSONEncoder):
datetime/date to isoformat
"""
- def default(self, o):
+ def default(self, o: Any) -> Any:
if isinstance(o, decimal.Decimal):
return str(o)
if isinstance(o, set):
- return sorted(list(o))
+ return sorted(o)
if isinstance(o, date):
return o.isoformat()
return super().default(o)
diff --git a/generalresearch/managers/__init__.py b/generalresearch/managers/__init__.py
index bc745fd..e69de29 100644
--- a/generalresearch/managers/__init__.py
+++ b/generalresearch/managers/__init__.py
@@ -1,16 +0,0 @@
-def parse_order_by(order_by_str: str) -> str:
- """
- Converts django-rest-framework ordering str to mysql clause
- :param order_by_str: e.g. 'created,-name'
- :return: mysql clause e.g. ORDER BY created ASC, name DESC
- """
- fields = order_by_str.split(",")
-
- order_clause = []
- for field in fields:
- if field.startswith("-"):
- order_clause.append(f"{field[1:]} DESC")
- else:
- order_clause.append(f"{field} ASC")
-
- return "ORDER BY " + ", ".join(order_clause)
diff --git a/generalresearch/managers/base.py b/generalresearch/managers/base.py
index b935413..2227d21 100644
--- a/generalresearch/managers/base.py
+++ b/generalresearch/managers/base.py
@@ -1,11 +1,14 @@
from __future__ import annotations
from collections.abc import Collection
+from contextlib import nullcontext
from enum import Enum
+from typing import TYPE_CHECKING
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
-from generalresearch.sql_helper import SqlHelper
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
+ from generalresearch.sql_helper import SqlHelper
class Permission(int, Enum):
@@ -45,6 +48,11 @@ class PostgresManager(Manager):
self.pg_config = pg_config
self.permissions = set(permissions) if permissions else set()
+ def connection(self, conn=None):
+ if conn is not None:
+ return nullcontext(conn)
+ return self.pg_config.make_connection()
+
class RedisManager(Manager):
CACHE_PREFIX = None
diff --git a/generalresearch/managers/cint/profiling.py b/generalresearch/managers/cint/profiling.py
index d549e94..c550632 100644
--- a/generalresearch/managers/cint/profiling.py
+++ b/generalresearch/managers/cint/profiling.py
@@ -1,10 +1,13 @@
from __future__ import annotations
import json
-from typing import Collection
+from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.cint.question import CintQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py
index f80542e..686a964 100644
--- a/generalresearch/managers/cint/survey.py
+++ b/generalresearch/managers/cint/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -107,7 +107,7 @@ class CintSurveyManager(SurveyManager):
return True
def update(self, surveys: list[CintSurvey]) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
for survey in surveys:
survey.last_updated = now
@@ -141,5 +141,5 @@ class CintSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/cint/user_pid.py b/generalresearch/managers/cint/user_pid.py
index 4f749a0..0265823 100644
--- a/generalresearch/managers/cint/user_pid.py
+++ b/generalresearch/managers/cint/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class CintUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py
index fe70732..760afd3 100644
--- a/generalresearch/managers/criteria.py
+++ b/generalresearch/managers/criteria.py
@@ -2,12 +2,24 @@ from __future__ import annotations
from abc import ABC
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
from more_itertools import chunked
from generalresearch.managers.base import SqlManager
-from generalresearch.models.thl.survey import MarketplaceCondition
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.survey import MarketplaceCondition
+
+DB_FIELDS = [
+ "hash",
+ "question_id",
+ "logical_operator",
+ "values",
+ "value_type",
+ "negate",
+]
class CriteriaManager(SqlManager, ABC):
@@ -15,14 +27,6 @@ class CriteriaManager(SqlManager, ABC):
Using the terms "criteria" & "condition" interchangeably!
"""
- DB_FIELDS = [
- "hash",
- "question_id",
- "logical_operator",
- "values",
- "value_type",
- "negate",
- ]
CONDITION_MODEL = None
TABLE_NAME = ""
@@ -30,7 +34,6 @@ class CriteriaManager(SqlManager, ABC):
"""
Create a single criterion
"""
- ...
def filter(self, hashes: Collection[str]) -> dict[str, MarketplaceCondition]:
"""
@@ -60,12 +63,12 @@ class CriteriaManager(SqlManager, ABC):
def update(self, conditions: Collection[MarketplaceCondition]) -> None:
# Add any new hashes into the DB
- this_hashes = set([condition.criterion_hash for condition in conditions])
+ this_hashes = {condition.criterion_hash for condition in conditions}
known_hashes = self.filter_exists(this_hashes)
new_hashes = this_hashes - known_hashes
if new_hashes:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
values = [
condition.to_mysql()
for condition in conditions
@@ -96,8 +99,6 @@ class CriteriaManager(SqlManager, ABC):
)
conn.commit()
- return None
-
@property
def mysql_fields(self) -> str:
return ", ".join([f"`{k}`" for k in self.DB_FIELDS])
diff --git a/generalresearch/managers/dynata/profiling.py b/generalresearch/managers/dynata/profiling.py
index 25afbe2..b76ecc7 100644
--- a/generalresearch/managers/dynata/profiling.py
+++ b/generalresearch/managers/dynata/profiling.py
@@ -1,10 +1,13 @@
from __future__ import annotations
import json
-from typing import Collection
+from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.dynata.question import DynataQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py
index 372a57d..21ec42e 100644
--- a/generalresearch/managers/dynata/survey.py
+++ b/generalresearch/managers/dynata/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -13,6 +13,34 @@ from generalresearch.models.dynata.survey import DynataCondition, DynataSurvey
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "status",
+ "is_live",
+ "client_id",
+ "bid_loi",
+ "bid_ir",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "expected_count",
+ "project_id",
+ "group_id",
+ "calculation_type",
+ "days_in_field",
+ "order_number",
+ "requirements",
+ "allowed_devices",
+ "category_exclusions",
+ "project_exclusions",
+ "live_link",
+ "category_ids",
+ "filters",
+ "quotas",
+ "used_question_ids",
+ "created",
+]
+
class DynataCriteriaManager(CriteriaManager):
CONDITION_MODEL = DynataCondition
@@ -20,33 +48,6 @@ class DynataCriteriaManager(CriteriaManager):
class DynataSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "status",
- "is_live",
- "client_id",
- "bid_loi",
- "bid_ir",
- "country_iso",
- "language_iso",
- "cpi",
- "expected_count",
- "project_id",
- "group_id",
- "calculation_type",
- "days_in_field",
- "order_number",
- "requirements",
- "allowed_devices",
- "category_exclusions",
- "project_exclusions",
- "live_link",
- "category_ids",
- "filters",
- "quotas",
- "used_question_ids",
- "created",
- ]
def get_survey_library(
self,
@@ -102,12 +103,12 @@ class DynataSurveyManager(SurveyManager):
return surveys
def create(self, survey: DynataSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = ["id"] + self.SURVEY_FIELDS + ["last_updated"]
+ create_fields = ["id"] + SURVEY_FIELDS + ["last_updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -123,11 +124,11 @@ class DynataSurveyManager(SurveyManager):
return True
def update(self, surveys: list[DynataSurvey]) -> bool:
- now = datetime.now(tz=timezone.utc)
- update_fields = self.SURVEY_FIELDS + ["last_updated"]
+ now = datetime.now(tz=UTC)
+ update_fields = SURVEY_FIELDS + ["last_updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update("dynata_survey", update_fields, survey_data)
return True
@@ -154,5 +155,6 @@ class DynataSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
+
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/dynata/user_pid.py b/generalresearch/managers/dynata/user_pid.py
index aefed34..67ff968 100644
--- a/generalresearch/managers/dynata/user_pid.py
+++ b/generalresearch/managers/dynata/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class DynataUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py
index 0be2bb9..4a2afb2 100644
--- a/generalresearch/managers/events.py
+++ b/generalresearch/managers/events.py
@@ -1,26 +1,25 @@
from __future__ import annotations
-import logging
import math
import socket
import threading
import time
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import TYPE_CHECKING
+from typing import TYPE_CHECKING, Any
from redis.client import PubSub, Redis
+from generalresearch.incite.base import LOG
from generalresearch.managers.base import RedisManager
-from generalresearch.models import Source
from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.definitions import Source
from generalresearch.models.events import (
AggregateBySource,
EventEnvelope,
EventMessage,
EventType,
MaxGaugeBySource,
- ServerToClientMessage,
ServerToClientMessageAdapter,
SessionEnterPayload,
SessionFinishPayload,
@@ -30,11 +29,14 @@ from generalresearch.models.events import (
TaskStatsSnapshot,
)
from generalresearch.models.thl.definitions import Status
-from generalresearch.models.thl.session import Session, Wall
-from generalresearch.models.thl.user import User
if TYPE_CHECKING:
from influxdb import InfluxDBClient
+
+ from generalresearch.models.events import ServerToClientMessage
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
+
else:
InfluxDBClient = object
@@ -141,7 +143,7 @@ class UserStatsManager(RedisManager):
pipe.execute()
def mark_user_active(self, user: User) -> None:
- now = datetime.now(tz=timezone.utc).isoformat()
+ now = datetime.now(tz=UTC).isoformat()
r = self.redis_client
pipe = r.pipeline(transaction=False)
@@ -175,7 +177,7 @@ class UserStatsManager(RedisManager):
# This call is idempotent; it can be called multiple times (for the
# same user) and won't falsely increase a counter; it will just
# reset the expiration for this user (times out after 60 min)
- now = datetime.now(tz=timezone.utc).isoformat()
+ now = datetime.now(tz=UTC).isoformat()
r = self.redis_client
pipe = r.pipeline(transaction=False)
@@ -216,16 +218,18 @@ class UserStatsManager(RedisManager):
class TaskStatsManager(RedisManager):
- task_stats = [
- "task_created_count_last_1h",
- "task_created_count_last_24h",
- "live_task_count",
- "live_tasks_max_payout",
- "TaskStatsManager:latest",
- ]
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
+
+ self.task_stats = [
+ "task_created_count_last_1h",
+ "task_created_count_last_24h",
+ "live_task_count",
+ "live_tasks_max_payout",
+ "TaskStatsManager:latest",
+ ]
+
self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT)
self.MAX_HASH_LUA = self.redis_client.register_script(MAX_HASH_LUA_SCRIPT)
@@ -343,8 +347,8 @@ class TaskStatsManager(RedisManager):
by_source=live_tasks_max_payout_by_source,
)
- task_created_count_last_1h = dict()
- task_created_count_last_24h = dict()
+ task_created_count_last_1h = {}
+ task_created_count_last_24h = {}
for source in sources:
task_created_count_last_1h[source] = pipe_res.pop(0)
task_created_count_last_24h[source] = pipe_res.pop(0)
@@ -381,25 +385,25 @@ class SessionStatsManager(RedisManager):
older than 1 hr (in the 1 hr bucket) will expire.
"""
- # Must be ordered. Don't change this
- global_keys = [
- "session_enters_last_1h",
- "session_enters_last_24h",
- "session_fails_last_1h",
- "session_fails_last_24h",
- "session_completes_last_1h",
- "session_completes_last_24h",
- "sum_payouts_last_1h",
- "sum_payouts_last_24h",
- "sum_user_payouts_last_1h",
- "sum_user_payouts_last_24h",
- # "session_fail_loi_sum_last_1h",
- "session_fail_loi_sum_last_24h",
- # "session_complete_loi_sum_last_1h",
- "session_complete_loi_sum_last_24h",
- ]
-
def __init__(self, *args, **kwargs):
+ # Must be ordered. Don't change this
+ self.global_keys = [
+ "session_enters_last_1h",
+ "session_enters_last_24h",
+ "session_fails_last_1h",
+ "session_fails_last_24h",
+ "session_completes_last_1h",
+ "session_completes_last_24h",
+ "sum_payouts_last_1h",
+ "sum_payouts_last_24h",
+ "sum_user_payouts_last_1h",
+ "sum_user_payouts_last_24h",
+ # "session_fail_loi_sum_last_1h",
+ "session_fail_loi_sum_last_24h",
+ # "session_complete_loi_sum_last_1h",
+ "session_complete_loi_sum_last_24h",
+ ]
+
super().__init__(*args, **kwargs)
self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT)
@@ -435,8 +439,8 @@ class SessionStatsManager(RedisManager):
pipe.hincrby(name, key, 1)
pipe.hexpire(name, ttl, key, nx=True)
# BP-specific tracker
- pipe.hincrby(name + ":" + user.product_id, key, 1)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, 1)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
# We're not returning this, but keep the sums, so we can
# calculate the avg
@@ -444,8 +448,8 @@ class SessionStatsManager(RedisManager):
value = round(session.elapsed.total_seconds())
pipe.hincrby(name, key, value)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, value)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, value)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
pipe.execute()
@@ -476,31 +480,31 @@ class SessionStatsManager(RedisManager):
pipe.hincrby(name, key, 1)
pipe.hexpire(name, ttl, key, nx=True)
# BP-specific tracker
- pipe.hincrby(name + ":" + user.product_id, key, 1)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, 1)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
name = "sum_payouts_" + name_postfix
amount = round(session.payout * 100)
pipe.hincrby(name, key, amount)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, amount)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, amount)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
if session.user_payout:
name = "sum_user_payouts_" + name_postfix
amount = round(session.user_payout * 100)
pipe.hincrby(name, key, amount)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, amount)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, amount)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
# We're not returning this, but keep the sums, so we can calculate the avg
name = "session_complete_loi_sum_" + name_postfix
value = round(session.elapsed.total_seconds())
pipe.hincrby(name, key, value)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, value)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, value)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
pipe.execute()
@@ -563,6 +567,7 @@ class SessionStatsManager(RedisManager):
res["session_avg_user_payout_last_24h"] = None
res["session_complete_avg_loi_last_24h"] = None
res["session_fail_avg_loi_last_24h"] = None
+
if res["session_completes_last_24h"]:
res["session_avg_payout_last_24h"] = math.ceil(
res["sum_payouts_last_24h"] / res["session_completes_last_24h"]
@@ -630,7 +635,7 @@ class EventManager(StatsManager):
def get_active_subscribers(self) -> set[UUIDStr]:
res = self.redis_client.pubsub_channels(f"{self.cache_prefix}:event-channel:*")
- product_ids = {x.rsplit(":", 1)[-1] for x in res}
+ product_ids = {str(x.rsplit(":", 1)[-1]) for x in res}
return product_ids
def stats_worker(self):
@@ -638,7 +643,7 @@ class EventManager(StatsManager):
try:
self.stats_worker_task()
except Exception as e:
- logging.exception(e)
+ LOG.exception(e)
finally:
time.sleep(60)
@@ -654,14 +659,14 @@ class EventManager(StatsManager):
lock_key = f"{self.cache_prefix}:event-channel-lock"
res = self.redis_client.set(lock_key, 1, ex=120, nx=True)
if not res:
- logging.debug("failed to acquire stats_worker_task lock")
+ LOG.debug("failed to acquire stats_worker_task lock")
return
- logging.info("Acquired stats_worker_task lock")
+ LOG.info("Acquired stats_worker_task lock")
for product_id in self.get_active_subscribers():
if time.monotonic() - now > 120:
- logging.exception("stats_worker_task is taking too long")
+ LOG.exception("stats_worker_task is taking too long")
break
channel = self.get_channel_name(product_id)
msg = self.get_stats_message(product_id=product_id)
@@ -680,7 +685,7 @@ class EventManager(StatsManager):
return
- def make_influx_point(self, channel: str, numsub: int):
+ def make_influx_point(self, channel: str, numsub: int) -> dict[str, Any]:
return {
"measurement": "redis_pubsub_subscribers",
"tags": {"hostname": socket.gethostname(), "channel": channel},
@@ -724,7 +729,6 @@ class EventManager(StatsManager):
)
)
self.publish_event(msg, product_id=user.product_id)
- return
def handle_task_finish(self, wall: Wall, session: Session, user: User):
self.mark_user_active(user=user)
@@ -818,7 +822,6 @@ class EventSubscriber(RedisManager):
p.subscribe(self.get_channel_name())
self.pubsub_client = r
self.pubsub = p
- return
def get_channel_name(self):
return f"{self.cache_prefix}:event-channel:{self.product_id}"
diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py
index a402693..f1ac2de 100644
--- a/generalresearch/managers/gr/authentication.py
+++ b/generalresearch/managers/gr/authentication.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import binascii
import logging
import os
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any
from psycopg import sql
@@ -11,13 +11,14 @@ from pydantic import AnyHttpUrl, PositiveInt
from generalresearch.managers.base import PostgresManager, PostgresManagerWithRedis
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
-
-LOG = logging.getLogger("gr")
if TYPE_CHECKING:
+
from generalresearch.models.gr.authentication import GRToken, GRUser
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
+
+LOG = logging.getLogger("gr")
class GRUserManager(PostgresManagerWithRedis):
@@ -29,7 +30,7 @@ class GRUserManager(PostgresManagerWithRedis):
) -> GRUser:
from generalresearch.models.gr.authentication import GRUser
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
instance = GRUser.model_validate(
{
@@ -146,8 +147,8 @@ class GRUserManager(PostgresManagerWithRedis):
for item in res:
for k, v in item.items():
- if isinstance(item[k], datetime):
- item[k] = item[k].replace(tzinfo=timezone.utc)
+ if isinstance(v, datetime):
+ item[k] = item[k].replace(tzinfo=UTC)
return [GRUser.model_validate(item) for item in res]
@@ -160,7 +161,7 @@ class GRUserManager(PostgresManagerWithRedis):
res = thl_pg_config.execute_sql_query(
query="""
- SELECT bp.id
+ SELECT bp.id::uuid as uuid
FROM userprofile_brokerageproduct AS bp
WHERE bp.business_id = ANY(%s)
""",
@@ -216,7 +217,7 @@ class GRTokenManager(PostgresManager):
"key": api_key,
"user_id": gr_user.id,
"user": gr_user,
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
}
)
@@ -234,10 +235,10 @@ class GRTokenManager(PostgresManager):
res = c.fetchall()
if len(res) == 0:
- raise Exception(f"No GRUser with token of '{api_key}'")
+ raise ValueError(f"No GRUser with token of '{api_key}'")
if len(res) > 1:
- raise Exception(f"Too many GRUsers found with token of '{api_key}'")
+ raise ValueError(f"Too many GRUsers found with token of '{api_key}'")
item = res[0]
@@ -251,7 +252,7 @@ class GRTokenManager(PostgresManager):
token = GRToken.model_validate(
{
"key": binascii.hexlify(os.urandom(20)).decode(),
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
"user_id": user_id,
}
)
@@ -270,8 +271,6 @@ class GRTokenManager(PostgresManager):
)
conn.commit()
- return
-
def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None:
# django authtoken_token table has (user_id) UNIQUE constraint
# therefore, this will only return 0 or 1 GRTokens
@@ -296,8 +295,8 @@ class GRTokenManager(PostgresManager):
res = result[0]
- for k, _ in res.items():
+ for k in res:
if isinstance(res[k], datetime):
- res[k] = res[k].replace(tzinfo=timezone.utc)
+ res[k] = res[k].replace(tzinfo=UTC)
return GRToken.model_validate(res)
diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py
index 4da0e7f..9338f03 100644
--- a/generalresearch/managers/gr/business.py
+++ b/generalresearch/managers/gr/business.py
@@ -12,14 +12,16 @@ from generalresearch.managers.base import (
PostgresManagerWithRedis,
)
from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessBankAccount,
+)
+from generalresearch.models.gr.definitions import BusinessType, TransferMethod
if TYPE_CHECKING:
+
from generalresearch.models.gr.business import (
- Business,
BusinessAddress,
- BusinessBankAccount,
- BusinessType,
- TransferMethod,
)
from generalresearch.models.gr.team import Team
@@ -36,8 +38,6 @@ class BusinessBankAccountManager(PostgresManager):
iban: str | None = None,
swift: str | None = None,
) -> BusinessBankAccount:
- from generalresearch.models.gr.business import BusinessBankAccount
-
ba = BusinessBankAccount.model_validate(
{
"business_id": business_id,
@@ -73,7 +73,6 @@ class BusinessBankAccountManager(PostgresManager):
return ba
def get_by_business_id(self, business_id: UUIDStr) -> list[BusinessBankAccount]:
- from generalresearch.models.gr.business import BusinessBankAccount
with self.pg_config.make_connection() as conn, conn.cursor() as c:
c.execute(
@@ -192,11 +191,7 @@ class BusinessManager(PostgresManagerWithRedis):
"""
Behavior: does this raise on duplicate?
"""
- from generalresearch.models.gr.business import (
- Business,
- BusinessType,
- )
-
+ # Business.model_rebuild()
business = Business.model_validate(
{
"uuid": uuid or uuid4().hex,
@@ -281,7 +276,6 @@ class BusinessManager(PostgresManagerWithRedis):
res = c.fetchall()
response = []
- from generalresearch.models.gr.business import Business
for i in res:
# i["contact"] = BusinessContact.model_validate(i)
@@ -370,8 +364,6 @@ class BusinessManager(PostgresManagerWithRedis):
self,
business_uuid: UUIDStr,
) -> Business | None:
- from generalresearch.models.gr.business import Business
-
assert UUID(hex=business_uuid).hex == business_uuid
with self.pg_config.make_connection() as conn, conn.cursor() as c:
@@ -397,8 +389,6 @@ class BusinessManager(PostgresManagerWithRedis):
return Business.model_validate(data)
def get_by_id(self, business_id: PositiveInt) -> Business | None:
- from generalresearch.models.gr.business import Business
-
assert isinstance(business_id, int)
with self.pg_config.make_connection() as conn, conn.cursor() as c:
diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py
index 6de82b0..41af709 100644
--- a/generalresearch/managers/gr/team.py
+++ b/generalresearch/managers/gr/team.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
from uuid import uuid4
@@ -11,15 +11,18 @@ from generalresearch.managers.base import (
PostgresManager,
PostgresManagerWithRedis,
)
+from generalresearch.managers.gr.authentication import GRUserManager
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.gr.team import Membership, MembershipPrivilege
+from generalresearch.models.gr.team import (
+ Membership,
+ MembershipPrivilege,
+)
if TYPE_CHECKING:
+
from generalresearch.models.gr.authentication import GRUser
from generalresearch.models.gr.business import Business
- from generalresearch.models.gr.team import (
- Team,
- )
+ from generalresearch.models.gr.team import Team
class MembershipManager(PostgresManager):
@@ -43,7 +46,7 @@ class MembershipManager(PostgresManager):
owner=False,
team_id=team.id,
user_id=gr_user.id,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
data = membership.model_dump(by_alias=True)
@@ -185,10 +188,12 @@ class TeamManager(PostgresManagerWithRedis):
return team
- def add_user(self, team: Team, gr_user: GRUser) -> Membership:
+ def add_user(
+ self, team: Team, gr_user: GRUser, gr_user_manager: GRUserManager
+ ) -> Membership:
"""Create a Membership between a GRUser and a Team"""
- team.prefetch_gr_users(pg_config=self.pg_config, redis_config=self.redis_config)
+ team.prefetch_gr_users(gr_user_manager=gr_user_manager)
assert gr_user not in team.gr_users, (
"Can't create multiple Memberships for " "the same User to the same Team"
diff --git a/generalresearch/managers/innovate/profiling.py b/generalresearch/managers/innovate/profiling.py
index bfa2685..0f32999 100644
--- a/generalresearch/managers/innovate/profiling.py
+++ b/generalresearch/managers/innovate/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.innovate.question import InnovateQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py
index 7db2f49..a4e36c1 100644
--- a/generalresearch/managers/innovate/survey.py
+++ b/generalresearch/managers/innovate/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -16,6 +16,44 @@ from generalresearch.models.innovate.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "buyer_id",
+ "job_id",
+ "survey_name",
+ "desired_count",
+ "remaining_count",
+ "supplier_completes_achieved",
+ "global_completes",
+ "global_starts",
+ "global_median_loi",
+ "global_conversion",
+ "bid_loi",
+ "bid_ir",
+ "allowed_devices",
+ "entry_link",
+ "category",
+ "requires_pii",
+ "excluded_surveys",
+ "duplicate_check_level",
+ "exclude_pids",
+ "include_pids",
+ "is_revenue_sharing",
+ "group_type",
+ "off_hour_traffic",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+ "modified_api",
+ "created_api",
+ "expected_end_date",
+]
+
class InnovateCriteriaManager(CriteriaManager):
CONDITION_MODEL = InnovateCondition
@@ -23,43 +61,6 @@ class InnovateCriteriaManager(CriteriaManager):
class InnovateSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "buyer_id",
- "job_id",
- "survey_name",
- "desired_count",
- "remaining_count",
- "supplier_completes_achieved",
- "global_completes",
- "global_starts",
- "global_median_loi",
- "global_conversion",
- "bid_loi",
- "bid_ir",
- "allowed_devices",
- "entry_link",
- "category",
- "requires_pii",
- "excluded_surveys",
- "duplicate_check_level",
- "exclude_pids",
- "include_pids",
- "is_revenue_sharing",
- "group_type",
- "off_hour_traffic",
- "qualifications",
- "quotas",
- "used_question_ids",
- "is_live",
- "modified_api",
- "created_api",
- "expected_end_date",
- ]
def get_survey_library(
self,
@@ -104,7 +105,7 @@ class InnovateSurveyManager(SurveyManager):
assert filters, "Must set at least 1 filter"
filter_str = " AND ".join(filters)
filter_str = "WHERE " + filter_str if filter_str else ""
- fields = set(self.SURVEY_FIELDS) | {"created", "updated"}
+ fields = set(SURVEY_FIELDS) | {"created", "updated"}
if exclude_fields:
fields -= exclude_fields
@@ -121,12 +122,12 @@ class InnovateSurveyManager(SurveyManager):
return surveys
def create(self, survey: InnovateSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -142,11 +143,11 @@ class InnovateSurveyManager(SurveyManager):
return True
def update(self, surveys: list[InnovateSurvey]) -> bool:
- now = datetime.now(tz=timezone.utc)
- update_fields = self.SURVEY_FIELDS + ["updated"]
+ now = datetime.now(tz=UTC)
+ update_fields = SURVEY_FIELDS + ["updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update(
table_name="innovate_survey",
field_names=update_fields,
@@ -179,5 +180,5 @@ class InnovateSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/innovate/user_pid.py b/generalresearch/managers/innovate/user_pid.py
index 100b0ca..7544c89 100644
--- a/generalresearch/managers/innovate/user_pid.py
+++ b/generalresearch/managers/innovate/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class InnovateUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/leaderboard/__init__.py b/generalresearch/managers/leaderboard/__init__.py
index d5138cd..048d2cb 100644
--- a/generalresearch/managers/leaderboard/__init__.py
+++ b/generalresearch/managers/leaderboard/__init__.py
@@ -1,8 +1,9 @@
from __future__ import annotations
+from zoneinfo import ZoneInfo
+
import pytz
from cachetools import LRUCache, cached
-from zoneinfo import ZoneInfo
@cached(cache=LRUCache(maxsize=1))
diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py
index 27d6a89..1673860 100644
--- a/generalresearch/managers/leaderboard/manager.py
+++ b/generalresearch/managers/leaderboard/manager.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from functools import cached_property
from typing import TYPE_CHECKING, cast
@@ -47,9 +47,7 @@ class LeaderboardManager:
self.product_id = product_id
self.country_iso = country_iso
if within_time is None:
- self.within_time_aware = datetime.now(tz=timezone.utc).astimezone(
- self.timezone
- )
+ self.within_time_aware = datetime.now(tz=UTC).astimezone(self.timezone)
elif within_time.tzinfo is not None:
self.within_time_aware = within_time.astimezone(self.timezone)
else:
@@ -59,7 +57,7 @@ class LeaderboardManager:
@cached_property
def period(self) -> Period:
local_ts = self.within_time_aware
- assert local_ts.tzinfo != timezone.utc and local_ts.tzinfo is not None
+ assert local_ts.tzinfo != UTC and local_ts.tzinfo is not None
t = pd.Timestamp(local_ts).tz_localize(tz=None)
freq_pd = {
LeaderboardFrequency.WEEKLY: "W-SUN",
diff --git a/generalresearch/managers/leaderboard/tasks.py b/generalresearch/managers/leaderboard/tasks.py
index 072a1aa..9e8dc9f 100644
--- a/generalresearch/managers/leaderboard/tasks.py
+++ b/generalresearch/managers/leaderboard/tasks.py
@@ -1,4 +1,5 @@
import logging
+from typing import TYPE_CHECKING
from redis import Redis
@@ -7,7 +8,9 @@ from generalresearch.models.thl.leaderboard import (
LeaderboardCode,
LeaderboardFrequency,
)
-from generalresearch.models.thl.session import Session
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.session import Session
logger = logging.getLogger()
diff --git a/generalresearch/managers/lucid/profiling.py b/generalresearch/managers/lucid/profiling.py
index fdd2d52..5937a59 100644
--- a/generalresearch/managers/lucid/profiling.py
+++ b/generalresearch/managers/lucid/profiling.py
@@ -2,12 +2,16 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from pydantic import ValidationError
-from generalresearch.decorators import LOG
from generalresearch.models.lucid.question import LucidQuestion, LucidQuestionType
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
+
+from generalresearch.decorators import LOG
def get_profiling_library(
diff --git a/generalresearch/managers/marketplace/__init__.py b/generalresearch/managers/marketplace/__init__.py
index 3349434..e69de29 100644
--- a/generalresearch/managers/marketplace/__init__.py
+++ b/generalresearch/managers/marketplace/__init__.py
@@ -1,23 +0,0 @@
-from generalresearch.managers.cint.user_pid import CintUserPidManager
-from generalresearch.managers.dynata.user_pid import DynataUserPidManager
-from generalresearch.managers.innovate.user_pid import InnovateUserPidManager
-from generalresearch.managers.morning.user_pid import MorningUserPidManager
-from generalresearch.managers.precision.user_pid import PrecisionUserPidManager
-from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager
-from generalresearch.managers.repdata.user_pid import RepdataUserPidManager
-from generalresearch.managers.sago.user_pid import SagoUserPidManager
-from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager
-
-_managers = [
- CintUserPidManager,
- DynataUserPidManager,
- InnovateUserPidManager,
- MorningUserPidManager,
- PrecisionUserPidManager,
- ProdegeUserPidManager,
- RepdataUserPidManager,
- SagoUserPidManager,
- SpectrumUserPidManager,
-]
-
-USER_PID_MANAGERS = {x.SOURCE: x for x in _managers}
diff --git a/generalresearch/managers/marketplace/managers.py b/generalresearch/managers/marketplace/managers.py
new file mode 100644
index 0000000..3349434
--- /dev/null
+++ b/generalresearch/managers/marketplace/managers.py
@@ -0,0 +1,23 @@
+from generalresearch.managers.cint.user_pid import CintUserPidManager
+from generalresearch.managers.dynata.user_pid import DynataUserPidManager
+from generalresearch.managers.innovate.user_pid import InnovateUserPidManager
+from generalresearch.managers.morning.user_pid import MorningUserPidManager
+from generalresearch.managers.precision.user_pid import PrecisionUserPidManager
+from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager
+from generalresearch.managers.repdata.user_pid import RepdataUserPidManager
+from generalresearch.managers.sago.user_pid import SagoUserPidManager
+from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager
+
+_managers = [
+ CintUserPidManager,
+ DynataUserPidManager,
+ InnovateUserPidManager,
+ MorningUserPidManager,
+ PrecisionUserPidManager,
+ ProdegeUserPidManager,
+ RepdataUserPidManager,
+ SagoUserPidManager,
+ SpectrumUserPidManager,
+]
+
+USER_PID_MANAGERS = {x.SOURCE: x for x in _managers}
diff --git a/generalresearch/managers/marketplace/user_pid.py b/generalresearch/managers/marketplace/user_pid.py
index 15d8a19..00dae8a 100644
--- a/generalresearch/managers/marketplace/user_pid.py
+++ b/generalresearch/managers/marketplace/user_pid.py
@@ -2,11 +2,14 @@ from __future__ import annotations
from abc import ABC
from collections.abc import Collection
+from typing import TYPE_CHECKING
from uuid import UUID
from generalresearch.managers.base import SqlManager
-from generalresearch.models import Source
-from generalresearch.sql_helper import SqlHelper
+from generalresearch.models.definitions import Source
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
class UserPidManager(SqlManager, ABC):
diff --git a/generalresearch/managers/morning/profiling.py b/generalresearch/managers/morning/profiling.py
index 01f99f3..7335e6b 100644
--- a/generalresearch/managers/morning/profiling.py
+++ b/generalresearch/managers/morning/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.morning.question import MorningQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py
index 2d86f0f..b43cd71 100644
--- a/generalresearch/managers/morning/survey.py
+++ b/generalresearch/managers/morning/survey.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import json
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -14,6 +14,50 @@ from generalresearch.models.morning.survey import MorningBid, MorningCondition
logger = logging.getLogger()
+STAT_FIELDS = [
+ "obs_median_loi",
+ "qualified_conversion",
+ "num_available",
+ "num_completes",
+ "num_failures",
+ "num_in_progress",
+ "num_over_quotas",
+ "num_qualified",
+ "num_quality_terminations",
+ "num_timeouts",
+]
+STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"]
+BID_FIELDS = (
+ [
+ "id",
+ "status",
+ "country_iso",
+ "language_isos",
+ "buyer_account_id",
+ "buyer_id",
+ "name",
+ "supplier_exclusive",
+ "survey_type",
+ "timeout",
+ "topic_id",
+ "bid_loi",
+ "exclusions",
+ "used_question_ids",
+ "expected_end",
+ "created_api",
+ "is_live",
+ ]
+ + STAT_FIELDS
+ + STAT_EXTENDED_FIELDS
+)
+QUOTA_FIELDS = [
+ "id",
+ "cpi",
+ "condition_hashes",
+] + STAT_FIELDS
+BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`"
+QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`"
+
class MorningCriteriaManager(CriteriaManager):
CONDITION_MODEL = MorningCondition
@@ -21,49 +65,6 @@ class MorningCriteriaManager(CriteriaManager):
class MorningSurveyManager(SurveyManager):
- STAT_FIELDS = [
- "obs_median_loi",
- "qualified_conversion",
- "num_available",
- "num_completes",
- "num_failures",
- "num_in_progress",
- "num_over_quotas",
- "num_qualified",
- "num_quality_terminations",
- "num_timeouts",
- ]
- STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"]
- BID_FIELDS = (
- [
- "id",
- "status",
- "country_iso",
- "language_isos",
- "buyer_account_id",
- "buyer_id",
- "name",
- "supplier_exclusive",
- "survey_type",
- "timeout",
- "topic_id",
- "bid_loi",
- "exclusions",
- "used_question_ids",
- "expected_end",
- "created_api",
- "is_live",
- ]
- + STAT_FIELDS
- + STAT_EXTENDED_FIELDS
- )
- QUOTA_FIELDS = [
- "id",
- "cpi",
- "condition_hashes",
- ] + STAT_FIELDS
- BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`"
- QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`"
def get_survey_library(
self,
@@ -138,7 +139,7 @@ class MorningSurveyManager(SurveyManager):
return bids
def create(self, bid: MorningBid) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = bid.to_mysql()
create_fields = self.BID_FIELDS + ["created", "updated"]
@@ -179,14 +180,14 @@ class MorningSurveyManager(SurveyManager):
return True
def update(self, surveys: list[MorningBid]) -> None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
for survey in surveys:
self.update_one(survey, now=now)
def update_one(self, bid: MorningBid, now: datetime | None = None) -> bool:
if now is None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = bid.to_mysql()
d["updated"] = now
@@ -258,5 +259,5 @@ class MorningSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/morning/user_pid.py b/generalresearch/managers/morning/user_pid.py
index 78de3bd..5896734 100644
--- a/generalresearch/managers/morning/user_pid.py
+++ b/generalresearch/managers/morning/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class MorningUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py
deleted file mode 100644
index 1efe875..0000000
--- a/generalresearch/managers/network/label.py
+++ /dev/null
@@ -1,147 +0,0 @@
-from __future__ import annotations
-
-from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
-
-from psycopg import sql
-from pydantic import IPvAnyNetwork, TypeAdapter
-
-from generalresearch.managers.base import PostgresManager
-from generalresearch.models.custom_types import (
- AwareDatetimeISO,
- IPvAnyAddressStr,
- IPvAnyNetworkStr,
-)
-from generalresearch.models.network.label import IPLabel, IPLabelKind, IPLabelSource
-
-
-class IPLabelManager(PostgresManager):
- def create(self, ip_label: IPLabel) -> IPLabel:
- query = sql.SQL("""
- INSERT INTO network_iplabel (
- ip, labeled_at, created_at,
- label_kind, source, confidence,
- provider, metadata
- ) VALUES (
- %(ip)s, %(labeled_at)s, %(created_at)s,
- %(label_kind)s, %(source)s, %(confidence)s,
- %(provider)s, %(metadata)s
- ) RETURNING id;""")
- params = ip_label.model_dump_postgres()
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
- pk = c.fetchone()["id"]
- return ip_label
-
- def make_filter_str(
- self,
- ips: Collection[IPvAnyNetworkStr] | None = None,
- ip_in_network: IPvAnyAddressStr | None = None,
- label_kind: IPLabelKind | None = None,
- source: IPLabelSource | None = None,
- labeled_at: AwareDatetimeISO | None = None,
- labeled_after: AwareDatetimeISO | None = None,
- labeled_before: AwareDatetimeISO | None = None,
- provider: str | None = None,
- ):
- filters = []
- params = {}
- if labeled_after or labeled_before:
- time_end = labeled_before or datetime.now(tz=timezone.utc)
- time_start = labeled_after or datetime(2017, 1, 1, tzinfo=timezone.utc)
- assert time_start.tzinfo.utcoffset(time_start) == timedelta(), "must be UTC"
- assert time_end.tzinfo.utcoffset(time_end) == timedelta(), "must be UTC"
- filters.append("labeled_at BETWEEN %(time_start)s AND %(time_end)s")
- params["time_start"] = time_start
- params["time_end"] = time_end
- if labeled_at:
- assert labeled_at.tzinfo.utcoffset(labeled_at) == timedelta(), "must be UTC"
- filters.append("labeled_at == %(labeled_at)s")
- params["labeled_at"] = labeled_at
- if label_kind:
- filters.append("label_kind = %(label_kind)s")
- params["label_kind"] = label_kind.value
- if source:
- filters.append("source = %(source)s")
- params["source"] = source.value
- if provider:
- filters.append("provider = %(provider)s")
- params["provider"] = provider
- if ips is not None:
- filters.append("ip = ANY(%(ips)s)")
- params["ips"] = list(ips)
- if ip_in_network:
- """
- Return matching networks.
- e.g. ip = '13f9:c462:e039:a38c::1', might return rows
- where ip = '13f9:c462:e039::/48' or '13f9:c462:e039:a38c::/64'
- """
- filters.append("ip >>= %(ip_in_network)s")
- params["ip_in_network"] = ip_in_network
-
- filter_str = "WHERE " + " AND ".join(filters) if filters else ""
- return filter_str, params
-
- def filter(
- self,
- ips: Collection[IPvAnyNetworkStr] | None = None,
- ip_in_network: IPvAnyAddressStr | None = None,
- label_kind: IPLabelKind | None = None,
- source: IPLabelSource | None = None,
- labeled_at: AwareDatetimeISO | None = None,
- labeled_after: AwareDatetimeISO | None = None,
- labeled_before: AwareDatetimeISO | None = None,
- provider: str | None = None,
- ) -> list[IPLabel]:
- filter_str, params = self.make_filter_str(
- ips=ips,
- ip_in_network=ip_in_network,
- label_kind=label_kind,
- source=source,
- labeled_at=labeled_at,
- labeled_after=labeled_after,
- labeled_before=labeled_before,
- provider=provider,
- )
- query = f"""
- SELECT
- ip, labeled_at, created_at,
- label_kind, source, confidence,
- provider, metadata
- FROM network_iplabel
- {filter_str}
- """
- res = self.pg_config.execute_sql_query(query, params)
- return [IPLabel.model_validate(rec) for rec in res]
-
- def get_most_specific_matching_network(self, ip: IPvAnyAddressStr) -> IPvAnyNetwork:
- """
- e.g. ip = 'b5f4:dc2:f136:70d5:5b6e:9a85:c7d4:3517', might return
- 'b5f4:dc2:f136:70d5::/64'
- """
- ip = TypeAdapter(IPvAnyAddressStr).validate_python(ip)
-
- query = """
- SELECT ip
- FROM network_iplabel
- WHERE ip >>= %(ip)s
- ORDER BY masklen(ip) DESC
- LIMIT 1;"""
- res = self.pg_config.execute_sql_query(query, {"ip": ip})
- if res:
- return IPvAnyNetwork(res[0]["ip"])
-
- def test_join(self, ip):
- query = """
- SELECT
- to_jsonb(i) AS ipinfo,
- to_jsonb(l) AS iplabel
- FROM thl_ipinformation i
- LEFT JOIN network_iplabel l
- ON l.ip >>= i.ip
- WHERE i.ip = %(ip)s
- ORDER BY masklen(l.ip) DESC;"""
- params = {"ip": ip}
- res = self.pg_config.execute_sql_query(query, params)
- return res
diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py
deleted file mode 100644
index 19c5caf..0000000
--- a/generalresearch/managers/network/mtr.py
+++ /dev/null
@@ -1,49 +0,0 @@
-from __future__ import annotations
-
-from psycopg import Cursor, sql
-
-from generalresearch.managers.base import PostgresManager
-from generalresearch.models.network.tool_run import MTRRun
-
-
-class MTRRunManager(PostgresManager):
-
- def _create(self, run: MTRRun, c: Cursor | None = None) -> None:
- """
- Do not use this directly. Must only be used in the context of a toolrun
- """
- query = sql.SQL("""
- INSERT INTO network_mtr (
- run_id, source_ip, facility_id,
- protocol, port, parsed,
- started_at, ip, scan_group_id
- )
- VALUES (
- %(run_id)s, %(source_ip)s, %(facility_id)s,
- %(protocol)s, %(port)s, %(parsed)s,
- %(started_at)s, %(ip)s, %(scan_group_id)s
- );
- """)
- params = run.model_dump_postgres()
-
- query_hops = sql.SQL("""
- INSERT INTO network_mtrhop (
- hop, ip, domain, asn, mtr_run_id
- ) VALUES (
- %(hop)s, %(ip)s, %(domain)s,
- %(asn)s, %(mtr_run_id)s
- )
- """)
- mtr_run = run.parsed
- params_hops = [h.model_dump_postgres(run_id=run.id) for h in mtr_run.hops]
-
- if c:
- c.execute(query, params)
- if params_hops:
- c.executemany(query_hops, params_hops)
- else:
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
- if params_hops:
- c.executemany(query_hops, params_hops)
diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py
deleted file mode 100644
index a8470c8..0000000
--- a/generalresearch/managers/network/nmap.py
+++ /dev/null
@@ -1,57 +0,0 @@
-from __future__ import annotations
-
-from psycopg import Cursor, sql
-
-from generalresearch.managers.base import PostgresManager
-from generalresearch.models.network.tool_run import NmapRun
-
-
-class NmapRunManager(PostgresManager):
-
- def _create(self, run: NmapRun, c: Cursor | None = None) -> None:
- """
- Insert a PortScan + PortScanPorts from a Pydantic NmapResult.
- Do not use this directly. Must only be used in the context of a toolrun
- """
- query = sql.SQL("""
- INSERT INTO network_portscan (
- run_id, xml_version, host_state,
- host_state_reason, latency_ms, distance,
- uptime_seconds, last_boot,
- parsed, scan_group_id, open_tcp_ports,
- started_at, ip, open_udp_ports
- )
- VALUES (
- %(run_id)s, %(xml_version)s, %(host_state)s,
- %(host_state_reason)s, %(latency_ms)s, %(distance)s,
- %(uptime_seconds)s, %(last_boot)s,
- %(parsed)s, %(scan_group_id)s, %(open_tcp_ports)s,
- %(started_at)s, %(ip)s, %(open_udp_ports)s
- );
- """)
- params = run.model_dump_postgres()
-
- query_ports = sql.SQL("""
- INSERT INTO network_portscanport (
- port_scan_id, protocol, port,
- state, reason, reason_ttl,
- service_name
- ) VALUES (
- %(port_scan_id)s, %(protocol)s, %(port)s,
- %(state)s, %(reason)s, %(reason_ttl)s,
- %(service_name)s
- )
- """)
- nmap_run = run.parsed
- params_ports = [p.model_dump_postgres(run_id=run.id) for p in nmap_run.ports]
-
- if c:
- c.execute(query, params)
- if nmap_run.ports:
- c.executemany(query_ports, params_ports)
- else:
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
- if nmap_run.ports:
- c.executemany(query_ports, params_ports)
diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py
deleted file mode 100644
index 0b41a9a..0000000
--- a/generalresearch/managers/network/rdns.py
+++ /dev/null
@@ -1,33 +0,0 @@
-from __future__ import annotations
-
-from psycopg import Cursor
-
-from generalresearch.managers.base import PostgresManager
-from generalresearch.models.network.tool_run import RDNSRun
-
-
-class RDNSRunManager(PostgresManager):
-
- def _create(self, run: RDNSRun, c: Cursor | None = None) -> None:
- """
- Do not use this directly. Must only be used in the context of a toolrun
- """
- query = """
- INSERT INTO network_rdnsresult (
- run_id, primary_hostname, primary_domain,
- hostname_count, hostnames,
- ip, started_at, scan_group_id
- )
- VALUES (
- %(run_id)s, %(primary_hostname)s, %(primary_domain)s,
- %(hostname_count)s, %(hostnames)s,
- %(ip)s, %(started_at)s, %(scan_group_id)s
- );
- """
- params = run.model_dump_postgres()
- if c:
- c.execute(query, params)
- else:
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, params)
diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py
deleted file mode 100644
index 026b3d3..0000000
--- a/generalresearch/managers/network/tool_run.py
+++ /dev/null
@@ -1,138 +0,0 @@
-from __future__ import annotations
-
-from collections.abc import Collection
-
-from psycopg import Cursor, sql
-
-from generalresearch.managers.base import Permission, PostgresManager
-from generalresearch.managers.network.mtr import MTRRunManager
-from generalresearch.managers.network.nmap import NmapRunManager
-from generalresearch.managers.network.rdns import RDNSRunManager
-from generalresearch.models.network.rdns.result import RDNSResult
-from generalresearch.models.network.tool_run import (
- MTRRun,
- NmapRun,
- RDNSRun,
- ToolName,
- ToolRun,
-)
-from generalresearch.pg_helper import PostgresConfig
-
-
-class ToolRunManager(PostgresManager):
- def __init__(
- self,
- pg_config: PostgresConfig,
- permissions: Collection[Permission] | None = None,
- ):
- super().__init__(pg_config=pg_config, permissions=permissions)
- self.nmap_manager = NmapRunManager(self.pg_config)
- self.rdns_manager = RDNSRunManager(self.pg_config)
- self.mtr_manager = MTRRunManager(self.pg_config)
-
- def _create_tool_run(self, run: NmapRun | RDNSRun | MTRRun, c: Cursor):
- query = sql.SQL("""
- INSERT INTO network_toolrun (
- ip, scan_group_id, tool_class,
- tool_name, tool_version, started_at,
- finished_at, status, raw_command,
- config
- )
- VALUES (
- %(ip)s, %(scan_group_id)s, %(tool_class)s,
- %(tool_name)s, %(tool_version)s, %(started_at)s,
- %(finished_at)s, %(status)s, %(raw_command)s,
- %(config)s
- ) RETURNING id;
- """)
- params = run.model_dump_postgres()
- c.execute(query, params)
- run_id = c.fetchone()["id"]
- run.id = run_id
- return None
-
- def create_tool_run(self, run: NmapRun | RDNSRun | MTRRun):
- if type(run) is NmapRun:
- return self.create_nmap_run(run)
- elif type(run) is RDNSRun:
- return self.create_rdns_run(run)
- elif type(run) is MTRRun:
- return self.create_mtr_run(run)
- else:
- raise ValueError("unrecognized run type")
-
- def get_latest_runs_by_tool(self, ip: str) -> dict[ToolName, ToolRun]:
- query = """
- SELECT DISTINCT ON (tool_name) *
- FROM network_toolrun
- WHERE ip = %(ip)s
- ORDER BY tool_name, started_at DESC;
- """
- params = {"ip": ip}
- res = self.pg_config.execute_sql_query(query, params=params)
- runs = [ToolRun.model_validate(x) for x in res]
- return {r.tool_name: r for r in runs}
-
- def create_nmap_run(self, run: NmapRun) -> NmapRun:
- """
- Insert a PortScan + PortScanPorts from a Pydantic NmapResult.
- """
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- self._create_tool_run(run, c)
- self.nmap_manager._create(run, c=c)
- return run
-
- def get_nmap_run(self, id: int) -> NmapRun:
- query = """
- SELECT tr.*, np.parsed
- FROM network_toolrun tr
- JOIN network_portscan np ON tr.id = np.run_id
- WHERE id = %(id)s
- """
- params = {"id": id}
- res = self.pg_config.execute_sql_query(query, params)[0]
- return NmapRun.model_validate(res)
-
- def create_rdns_run(self, run: RDNSRun) -> RDNSRun:
- """
- Insert a RDnsRun + RDNSResult
- """
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- self._create_tool_run(run, c)
- self.rdns_manager._create(run, c=c)
- return run
-
- def get_rdns_run(self, id: int) -> RDNSRun:
- query = """
- SELECT tr.*, hostnames
- FROM network_toolrun tr
- JOIN network_rdnsresult np ON tr.id = np.run_id
- WHERE id = %(id)s
- """
- params = {"id": id}
- res = self.pg_config.execute_sql_query(query, params)[0]
- parsed = RDNSResult.model_validate(
- {"ip": res["ip"], "hostnames": res["hostnames"]}
- )
- res["parsed"] = parsed
- return RDNSRun.model_validate(res)
-
- def create_mtr_run(self, run: MTRRun) -> MTRRun:
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- self._create_tool_run(run, c)
- self.mtr_manager._create(run, c=c)
- return run
-
- def get_mtr_run(self, id: int) -> MTRRun:
- query = """
- SELECT tr.*, mtr.parsed, mtr.source_ip, mtr.facility_id
- FROM network_toolrun tr
- JOIN network_mtr mtr ON tr.id = mtr.run_id
- WHERE id = %(id)s
- """
- params = {"id": id}
- res = self.pg_config.execute_sql_query(query, params)[0]
- return MTRRun.model_validate(res)
diff --git a/generalresearch/managers/pollfish/profiling.py b/generalresearch/managers/pollfish/profiling.py
index daf529b..43afd30 100644
--- a/generalresearch/managers/pollfish/profiling.py
+++ b/generalresearch/managers/pollfish/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.pollfish.question import PollfishQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/pollfish/user_pid.py b/generalresearch/managers/pollfish/user_pid.py
index 1068405..f3983cf 100644
--- a/generalresearch/managers/pollfish/user_pid.py
+++ b/generalresearch/managers/pollfish/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class PollfishUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/precision/profiling.py b/generalresearch/managers/precision/profiling.py
index 449fd25..813542c 100644
--- a/generalresearch/managers/precision/profiling.py
+++ b/generalresearch/managers/precision/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.precision.question import PrecisionQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py
index 6fb30f2..cc28287 100644
--- a/generalresearch/managers/precision/survey.py
+++ b/generalresearch/managers/precision/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -16,6 +16,30 @@ from generalresearch.models.precision.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ # 'country_iso', 'language_iso', # these come from join table
+ "survey_id",
+ "is_live",
+ "status",
+ "cpi",
+ "group_id",
+ "name",
+ "survey_guid",
+ "buyer_id",
+ "category_id",
+ "bid_loi",
+ "bid_ir",
+ "global_conversion",
+ "desired_count",
+ "achieved_count",
+ "allowed_devices",
+ "entry_link",
+ "excluded_surveys",
+ "quotas",
+ "used_question_ids",
+ "expected_end_date",
+]
+
class PrecisionCriteriaManager(CriteriaManager):
CONDITION_MODEL = PrecisionCondition
@@ -23,29 +47,6 @@ class PrecisionCriteriaManager(CriteriaManager):
class PrecisionSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- # 'country_iso', 'language_iso', # these come from join table
- "survey_id",
- "is_live",
- "status",
- "cpi",
- "group_id",
- "name",
- "survey_guid",
- "buyer_id",
- "category_id",
- "bid_loi",
- "bid_ir",
- "global_conversion",
- "desired_count",
- "achieved_count",
- "allowed_devices",
- "entry_link",
- "excluded_surveys",
- "quotas",
- "used_question_ids",
- "expected_end_date",
- ]
def get_survey_library(
self,
@@ -104,12 +105,12 @@ class PrecisionSurveyManager(SurveyManager):
return surveys
def create(self, survey: PrecisionSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(False)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -125,7 +126,7 @@ class PrecisionSurveyManager(SurveyManager):
country_data = [(survey.survey_id, c) for c in survey.country_isos]
c.executemany(
- f"""
+ """
INSERT INTO `thl-precision`.`precision_survey_country`
(survey_id, country_iso, is_active) VALUES
(%s, %s, TRUE)
@@ -134,7 +135,7 @@ class PrecisionSurveyManager(SurveyManager):
)
lang_data = [(survey.survey_id, c) for c in survey.language_isos]
c.executemany(
- f"""
+ """
INSERT INTO `thl-precision`.`precision_survey_language`
(survey_id, language_iso, is_active) VALUES
(%s, %s, TRUE)
@@ -151,7 +152,7 @@ class PrecisionSurveyManager(SurveyManager):
return True
def update_one(self, survey: PrecisionSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
d["updated"] = now
@@ -188,7 +189,7 @@ class PrecisionSurveyManager(SurveyManager):
country_data = [(survey.survey_id, c) for c in survey.country_isos]
# Turn ON countries in this survey's list of countries, insert row, if already exists, set active.
c.executemany(
- query=f"""
+ query="""
INSERT INTO `thl-precision`.`precision_survey_country`
(survey_id, country_iso, is_active) VALUES
(%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE;
@@ -207,7 +208,7 @@ class PrecisionSurveyManager(SurveyManager):
)
language_data = [(survey.survey_id, c) for c in survey.language_isos]
c.executemany(
- query=f"""
+ query="""
INSERT INTO `thl-precision`.`precision_survey_language`
(survey_id, language_iso, is_active) VALUES
(%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE;
@@ -241,5 +242,5 @@ class PrecisionSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/precision/user_pid.py b/generalresearch/managers/precision/user_pid.py
index 50e97e6..ed2d58d 100644
--- a/generalresearch/managers/precision/user_pid.py
+++ b/generalresearch/managers/precision/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class PrecisionUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/prodege/profiling.py b/generalresearch/managers/prodege/profiling.py
index 54a7b57..cf77fca 100644
--- a/generalresearch/managers/prodege/profiling.py
+++ b/generalresearch/managers/prodege/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.prodege.question import ProdegeQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py
index f555290..ef98d7b 100644
--- a/generalresearch/managers/prodege/survey.py
+++ b/generalresearch/managers/prodege/survey.py
@@ -1,7 +1,7 @@
from __future__ import annotations
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
@@ -9,6 +9,31 @@ from generalresearch.managers.criteria import CriteriaManager
from generalresearch.managers.survey import SurveyManager
from generalresearch.models.prodege.survey import ProdegeCondition, ProdegeSurvey
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_name",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "desired_count",
+ "remaining_count",
+ "achieved_completes",
+ "bid_loi",
+ "bid_ir",
+ "actual_loi",
+ "actual_ir",
+ "conversion_rate",
+ "entrance_url",
+ "max_clicks_settings",
+ "past_participation",
+ "include_psids",
+ "exclude_psids",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+]
+
class ProdegeCriteriaManager(CriteriaManager):
CONDITION_MODEL = ProdegeCondition
@@ -16,30 +41,6 @@ class ProdegeCriteriaManager(CriteriaManager):
class ProdegeSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_name",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "desired_count",
- "remaining_count",
- "achieved_completes",
- "bid_loi",
- "bid_ir",
- "actual_loi",
- "actual_ir",
- "conversion_rate",
- "entrance_url",
- "max_clicks_settings",
- "past_participation",
- "include_psids",
- "exclude_psids",
- "quotas",
- "used_question_ids",
- "is_live",
- ]
def get_survey_library(
self,
@@ -93,12 +94,12 @@ class ProdegeSurveyManager(SurveyManager):
return surveys
def create(self, survey: ProdegeSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -114,7 +115,7 @@ class ProdegeSurveyManager(SurveyManager):
return True
def update(self, surveys: list[ProdegeSurvey]) -> None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# Do to stupidity with bid/actual loi/ir values (see ProdegeSurvey.to_mysql), we now
# can't do a bulk update b/c the fields may be different in different rows. Just do
@@ -124,7 +125,7 @@ class ProdegeSurveyManager(SurveyManager):
def update_one(self, survey: ProdegeSurvey, now=None) -> bool:
if now is None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
# We have to have special logic for bid/actual loi/ir here. The api is
# stupid and only returns one set of them. If we just do the db
diff --git a/generalresearch/managers/prodege/user_pid.py b/generalresearch/managers/prodege/user_pid.py
index 7c92e28..c18c109 100644
--- a/generalresearch/managers/prodege/user_pid.py
+++ b/generalresearch/managers/prodege/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class ProdegeUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/repdata/profiling.py b/generalresearch/managers/repdata/profiling.py
index 4b97abd..ec30d78 100644
--- a/generalresearch/managers/repdata/profiling.py
+++ b/generalresearch/managers/repdata/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.repdata.question import RepDataQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py
index 2e2224f..2c3156c 100644
--- a/generalresearch/managers/repdata/survey.py
+++ b/generalresearch/managers/repdata/survey.py
@@ -2,7 +2,8 @@ from __future__ import annotations
import json
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
import pymysql
@@ -11,10 +12,46 @@ from generalresearch.managers.survey import SurveyManager
from generalresearch.models.repdata.survey import (
RepDataCondition,
RepDataStreamHashed,
- RepDataSurvey,
RepDataSurveyHashed,
)
+if TYPE_CHECKING:
+ from generalresearch.models.repdata.survey import RepDataSurvey
+
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_uuid",
+ "survey_name",
+ "project_uuid",
+ "survey_status",
+ "country_iso",
+ "language_iso",
+ "estimated_loi",
+ "estimated_ir",
+ "collects_pii",
+ "allowed_devices",
+]
+STREAM_FIELDS = [
+ "stream_id",
+ "stream_uuid",
+ "stream_name",
+ "stream_status",
+ "calculation_type",
+ "qualification_hashes",
+ "hashed_quotas",
+ "expected_count",
+ "cpi",
+ "days_in_field",
+ "actual_ir",
+ "actual_loi",
+ "actual_conversion",
+ "actual_complete_count",
+ "actual_count",
+ "used_question_ids",
+ "survey_id",
+ "remaining_count",
+]
+
class RepDataCriteriaManager(CriteriaManager):
CONDITION_MODEL = RepDataCondition
@@ -22,39 +59,6 @@ class RepDataCriteriaManager(CriteriaManager):
class RepDataSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_uuid",
- "survey_name",
- "project_uuid",
- "survey_status",
- "country_iso",
- "language_iso",
- "estimated_loi",
- "estimated_ir",
- "collects_pii",
- "allowed_devices",
- ]
- STREAM_FIELDS = [
- "stream_id",
- "stream_uuid",
- "stream_name",
- "stream_status",
- "calculation_type",
- "qualification_hashes",
- "hashed_quotas",
- "expected_count",
- "cpi",
- "days_in_field",
- "actual_ir",
- "actual_loi",
- "actual_conversion",
- "actual_complete_count",
- "actual_count",
- "used_question_ids",
- "survey_id",
- "remaining_count",
- ]
def get_survey_library(
self,
@@ -105,7 +109,7 @@ class RepDataSurveyManager(SurveyManager):
surveys = {s.survey_id: s for s in surveys}
if surveys:
res = self.sql_helper.execute_sql_query(
- query=f"""
+ query="""
SELECT *
FROM `thl-repdata`.`repdata_surveystream`
WHERE survey_id IN %s
@@ -122,12 +126,12 @@ class RepDataSurveyManager(SurveyManager):
return list(surveys.values())
def create(self, survey: RepDataSurvey | RepDataSurveyHashed) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "last_updated"]
+ create_fields = SURVEY_FIELDS + ["created", "last_updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -141,10 +145,10 @@ class RepDataSurveyManager(SurveyManager):
args=survey_data,
)
- fields_str = ", ".join([f"`{x}`" for x in self.STREAM_FIELDS])
- values_str = ", ".join([f"%({x})s" for x in self.STREAM_FIELDS])
+ fields_str = ", ".join([f"`{x}`" for x in STREAM_FIELDS])
+ values_str = ", ".join([f"%({x})s" for x in STREAM_FIELDS])
stream_data = [
- {k: v for k, v in stream.items() if k in self.STREAM_FIELDS}
+ {k: v for k, v in stream.items() if k in STREAM_FIELDS}
for stream in d["streams"]
]
for sd in stream_data:
@@ -160,11 +164,11 @@ class RepDataSurveyManager(SurveyManager):
return True
def update(self, surveys: list[RepDataSurveyHashed]) -> bool:
- now = datetime.now(tz=timezone.utc)
- update_fields = self.SURVEY_FIELDS + ["last_updated"]
+ now = datetime.now(tz=UTC)
+ update_fields = SURVEY_FIELDS + ["last_updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update(
table_name="repdata_survey",
field_names=update_fields,
@@ -175,7 +179,7 @@ class RepDataSurveyManager(SurveyManager):
for d in data:
for stream in d["streams"]:
stream["survey_id"] = d["survey_id"]
- stream_data.append([stream[k] for k in self.STREAM_FIELDS])
+ stream_data.append([stream[k] for k in STREAM_FIELDS])
self.sql_helper.bulk_update(
table_name="repdata_surveystream",
diff --git a/generalresearch/managers/repdata/user_pid.py b/generalresearch/managers/repdata/user_pid.py
index 9d53897..5fdeccf 100644
--- a/generalresearch/managers/repdata/user_pid.py
+++ b/generalresearch/managers/repdata/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class RepdataUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/sago/profiling.py b/generalresearch/managers/sago/profiling.py
index 4f5b2f5..3bbad3f 100644
--- a/generalresearch/managers/sago/profiling.py
+++ b/generalresearch/managers/sago/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.sago.question import SagoQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py
index 325639f..9228528 100644
--- a/generalresearch/managers/sago/survey.py
+++ b/generalresearch/managers/sago/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -13,6 +13,31 @@ from generalresearch.models.sago.survey import SagoCondition, SagoSurvey
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "is_live",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "buyer_id",
+ "account_id",
+ "study_type_id",
+ "industry_id",
+ "allowed_devices",
+ "collects_pii",
+ "bid_loi",
+ "bid_ir",
+ "live_link",
+ "survey_exclusions",
+ "ip_exclusions",
+ "remaining_count",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "modified_api",
+]
+
class SagoCriteriaManager(CriteriaManager):
CONDITION_MODEL = SagoCondition
@@ -20,30 +45,6 @@ class SagoCriteriaManager(CriteriaManager):
class SagoSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "is_live",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "buyer_id",
- "account_id",
- "study_type_id",
- "industry_id",
- "allowed_devices",
- "collects_pii",
- "bid_loi",
- "bid_ir",
- "live_link",
- "survey_exclusions",
- "ip_exclusions",
- "remaining_count",
- "qualifications",
- "quotas",
- "used_question_ids",
- "modified_api",
- ]
def get_survey_library(
self,
@@ -85,7 +86,7 @@ class SagoSurveyManager(SurveyManager):
assert filters, "Must set at least 1 filter"
filter_str = " AND ".join(filters)
filter_str = "WHERE " + filter_str if filter_str else ""
- fields = set(self.SURVEY_FIELDS) | {"created", "updated"}
+ fields = set(SURVEY_FIELDS) | {"created", "updated"}
if exclude_fields:
fields -= exclude_fields
fields_str = ", ".join([f"`{v}`" for v in fields])
@@ -101,12 +102,12 @@ class SagoSurveyManager(SurveyManager):
return surveys
def create(self, survey: SagoSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -122,16 +123,16 @@ class SagoSurveyManager(SurveyManager):
return True
def update(self, surveys: list[SagoSurvey]) -> bool:
- now = datetime.now(tz=timezone.utc)
- update_fields = self.SURVEY_FIELDS + ["updated"]
+ now = datetime.now(tz=UTC)
+ update_fields = SURVEY_FIELDS + ["updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update("sago_survey", update_fields, survey_data)
return True
def update_field(self, survey: SagoSurvey, field: str) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
conn: pymysql.Connection = self.sql_helper.make_connection()
value = survey.to_mysql()[field]
c = conn.cursor()
@@ -156,8 +157,8 @@ class SagoSurveyManager(SurveyManager):
return True
def create_or_update(self, surveys: list[SagoSurvey]) -> None:
- surveys = {s.survey_id: s for s in surveys}
- sns = set(surveys.keys())
+ _surveys = {s.survey_id: s for s in surveys}
+ sns = set(_surveys.keys())
existing_sns = {
x["survey_id"]
for x in self.sql_helper.execute_sql_query(
@@ -171,7 +172,7 @@ class SagoSurveyManager(SurveyManager):
}
create_sns = sns - existing_sns
for sn in create_sns:
- survey = surveys[sn]
+ survey = _surveys[sn]
try:
self.create(survey)
except IntegrityError as e:
@@ -179,6 +180,6 @@ class SagoSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
- self.update([surveys[sn] for sn in existing_sns])
+ self.update([_surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/sago/user_pid.py b/generalresearch/managers/sago/user_pid.py
index 311abb7..b7ce771 100644
--- a/generalresearch/managers/sago/user_pid.py
+++ b/generalresearch/managers/sago/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class SagoUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/spectrum/profiling.py b/generalresearch/managers/spectrum/profiling.py
index 8a0904a..5de218c 100644
--- a/generalresearch/managers/spectrum/profiling.py
+++ b/generalresearch/managers/spectrum/profiling.py
@@ -2,9 +2,12 @@ from __future__ import annotations
import json
from collections.abc import Collection
+from typing import TYPE_CHECKING
from generalresearch.models.spectrum.question import SpectrumQuestion
-from generalresearch.sql_helper import SqlHelper
+
+if TYPE_CHECKING:
+ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py
index 3ff2db8..5059716 100644
--- a/generalresearch/managers/spectrum/survey.py
+++ b/generalresearch/managers/spectrum/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pymysql
from pymysql import IntegrityError
@@ -16,6 +16,37 @@ from generalresearch.models.spectrum.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_name",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "field_end_date",
+ "category_code",
+ "calculation_type",
+ "requires_pii",
+ "buyer_id",
+ "survey_exclusions",
+ "exclusion_period",
+ "bid_loi",
+ "bid_ir",
+ "last_block_loi",
+ "last_block_ir",
+ "overall_ir",
+ "overall_loi",
+ "project_last_complete_date",
+ "include_psids",
+ "exclude_psids",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+ "modified_api",
+ "created_api",
+]
+
class SpectrumCriteriaManager(CriteriaManager):
CONDITION_MODEL = SpectrumCondition
@@ -23,36 +54,6 @@ class SpectrumCriteriaManager(CriteriaManager):
class SpectrumSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_name",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "field_end_date",
- "category_code",
- "calculation_type",
- "requires_pii",
- "buyer_id",
- "survey_exclusions",
- "exclusion_period",
- "bid_loi",
- "bid_ir",
- "last_block_loi",
- "last_block_ir",
- "overall_ir",
- "overall_loi",
- "project_last_complete_date",
- "include_psids",
- "exclude_psids",
- "qualifications",
- "quotas",
- "used_question_ids",
- "is_live",
- "modified_api",
- "created_api",
- ]
def get_survey_library(
self,
@@ -110,12 +111,12 @@ class SpectrumSurveyManager(SurveyManager):
return surveys
def create(self, survey: SpectrumSurvey) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["updated"]
+ create_fields = SURVEY_FIELDS + ["updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -134,7 +135,7 @@ class SpectrumSurveyManager(SurveyManager):
return True
def update(self, surveys: list[SpectrumSurvey]) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# Due to stupidity with bid/actual loi/ir values (last block nonsense),
# we can't do a bulk update b/c the fields may be different in
@@ -146,7 +147,7 @@ class SpectrumSurveyManager(SurveyManager):
def update_one(self, survey: SpectrumSurvey, now: datetime | None = None) -> bool:
if now is None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = survey.to_mysql()
# We have to have special logic for bid/actual loi/ir here. The api
@@ -212,6 +213,6 @@ class SpectrumSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/spectrum/user_pid.py b/generalresearch/managers/spectrum/user_pid.py
index 495e73c..980c28d 100644
--- a/generalresearch/managers/spectrum/user_pid.py
+++ b/generalresearch/managers/spectrum/user_pid.py
@@ -1,5 +1,5 @@
from generalresearch.managers.marketplace.user_pid import UserPidManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
class SpectrumUserPidManager(UserPidManager):
diff --git a/generalresearch/managers/survey.py b/generalresearch/managers/survey.py
index 964344f..65dbdd9 100644
--- a/generalresearch/managers/survey.py
+++ b/generalresearch/managers/survey.py
@@ -1,9 +1,12 @@
from __future__ import annotations
from abc import ABC
+from typing import TYPE_CHECKING
from generalresearch.managers.base import SqlManager
-from generalresearch.models.thl.survey import MarketplaceTask
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.survey import MarketplaceTask
class SurveyManager(SqlManager, ABC):
@@ -12,14 +15,12 @@ class SurveyManager(SqlManager, ABC):
"""
Create a single survey
"""
- ...
def update(self, surveys: list[MarketplaceTask]) -> bool:
"""
Update a list of surveys. Depending on the implementation, this may
operate one by one or as a bulk update.
"""
- ...
def update_field(self, survey: MarketplaceTask, field: str) -> bool:
"""
diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py
index 04452cd..38214c6 100644
--- a/generalresearch/managers/thl/buyer.py
+++ b/generalresearch/managers/thl/buyer.py
@@ -1,12 +1,15 @@
from __future__ import annotations
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
from generalresearch.managers.base import Permission, PostgresManager
-from generalresearch.models import Source
from generalresearch.models.thl.survey.buyer import Buyer
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.models.definitions import Source
+ from generalresearch.pg_helper import PostgresConfig
class BuyerManager(PostgresManager):
@@ -18,8 +21,8 @@ class BuyerManager(PostgresManager):
):
super().__init__(pg_config=pg_config, permissions=permissions)
# self.buyer_pk: Dict[Buyer, int] = dict()
- self.source_code_buyer: dict[str, Buyer] = dict()
- self.source_code_pk: dict[str, int] = dict()
+ self.source_code_buyer: dict[str, Buyer] = {}
+ self.source_code_pk: dict[str, int] = {}
self.populate_caches()
def populate_caches(self):
@@ -45,7 +48,7 @@ class BuyerManager(PostgresManager):
return None
def bulk_get_or_create(self, source: Source, codes: Collection[str]) -> list[Buyer]:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
buyers = []
params_seq = []
diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py
index 91d6a42..7d2366d 100644
--- a/generalresearch/managers/thl/cashout_method.py
+++ b/generalresearch/managers/thl/cashout_method.py
@@ -2,26 +2,28 @@ from __future__ import annotations
from collections.abc import Collection
from copy import copy
-from datetime import datetime, timezone
-from typing import Any
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any
from uuid import UUID, uuid4
from pydantic import NonNegativeInt
from generalresearch.managers.base import PostgresManager
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.wallet.cashout_method import (
- CashMailCashoutMethodData,
- CashoutMethod,
- PaypalCashoutMethodData,
-)
+from generalresearch.models.thl.user_ref import UserRef
+from generalresearch.models.thl.wallet.definitions import PayoutType
+if TYPE_CHECKING:
+ from generalresearch.models.thl.user import User
+ from generalresearch.models.thl.wallet.cashout_method import (
+ CashMailCashoutMethodData,
+ CashoutMethod,
+ PaypalCashoutMethodData,
+ )
-class CashoutMethodManager(PostgresManager):
+class CashoutMethodManager(PostgresManager):
def create(self, cm: CashoutMethod) -> None:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
query = """
INSERT INTO accounting_cashoutmethod (
id, last_updated, is_live, provider,
@@ -56,9 +58,9 @@ class CashoutMethodManager(PostgresManager):
res = next(iter(db_res), None)
assert res, f"cashout method id {cm_id} not found"
# Don't let anyone delete a non-user-scoped cashout method
- assert (
- res["user_id"] is not None
- ), "error trying to delete non user-scoped cashout method"
+ assert res["user_id"] is not None, (
+ "error trying to delete non user-scoped cashout method"
+ )
self.pg_config.execute_write(
query="""
@@ -78,6 +80,7 @@ class CashoutMethodManager(PostgresManager):
:return: the uuid of the created cashout method
"""
# todo: validate shipping address?
+ from generalresearch.models.thl.wallet.cashout_method import CashoutMethod
cm = CashoutMethod(
name="Cash in Mail",
@@ -90,7 +93,7 @@ class CashoutMethodManager(PostgresManager):
max_value=25000, # $250.00
data=data,
type=PayoutType.CASH_IN_MAIL,
- user=user,
+ user=user.to_user_ref(),
ext_id=data.delivery_address.md5sum(),
)
@@ -122,6 +125,8 @@ class CashoutMethodManager(PostgresManager):
:param user:
:return: the uuid of the created cashout method
"""
+ from generalresearch.models.thl.wallet.cashout_method import CashoutMethod
+
cm = CashoutMethod(
name="PayPal",
description="Cashout via PayPal",
@@ -132,7 +137,7 @@ class CashoutMethodManager(PostgresManager):
max_value=25_000, # $250.00
data=data,
type=PayoutType.PAYPAL,
- user=user,
+ user=user.to_user_ref(),
ext_id=data.email,
)
# Make sure this user doesn't already have one
@@ -154,13 +159,13 @@ class CashoutMethodManager(PostgresManager):
@staticmethod
def make_filter_str(
uuid: str | None = None,
- user: User | None = None,
+ user: UserRef | None = None,
ext_id: str | None = None,
payout_types: Collection[PayoutType] | None = None,
is_live: bool | None = True,
):
filters = []
- params = dict()
+ params = {}
if uuid is not None:
params["uuid"] = uuid
filters.append("id = %(uuid)s")
@@ -185,7 +190,7 @@ class CashoutMethodManager(PostgresManager):
def filter_count(
self,
uuid: str | None = None,
- user: User | None = None,
+ user: UserRef | None = None,
ext_id: str | None = None,
payout_types: Collection[PayoutType] | None = None,
is_live: bool | None = True,
@@ -210,7 +215,7 @@ class CashoutMethodManager(PostgresManager):
def filter(
self,
uuid: str | None = None,
- user: User | None = None,
+ user: UserRef | None = None,
ext_id: str | None = None,
payout_types: Collection[PayoutType] | None = None,
is_live: bool | None = True,
@@ -274,16 +279,18 @@ class CashoutMethodManager(PostgresManager):
# The data column here is inconsistent. Pulling keys from the mysql 'data' col
# and putting them into the base level. Renamed so that we don't overwrite
# a col called "data" within the "_data_" field.
+ from generalresearch.models.thl.wallet.cashout_method import CashoutMethod
+
for k in list(x["_data_"].keys()):
if k in CashoutMethod.model_fields:
x[k] = x["_data_"].pop(k)
x["type"] = PayoutType(x["provider"].upper())
if "data" not in x:
- x["data"] = dict()
+ x["data"] = {}
x["data"].update(x.pop("_data_"))
x["data"]["type"] = x["type"]
if user and x["type"] in {PayoutType.PAYPAL, PayoutType.CASH_IN_MAIL}:
- x["user"] = user
+ x["user"] = user.to_user_ref()
return CashoutMethod.model_validate(x)
diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py
index 05ceb8f..67c812a 100644
--- a/generalresearch/managers/thl/category.py
+++ b/generalresearch/managers/thl/category.py
@@ -1,16 +1,18 @@
from __future__ import annotations
from collections.abc import Collection
+from typing import TYPE_CHECKING
-from generalresearch.managers.base import Permission, PostgresManager
+from generalresearch.managers.base import PostgresManager
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.category import Category
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.base import Permission
+ from generalresearch.pg_helper import PostgresConfig
class CategoryManager(PostgresManager):
- categories = dict()
- category_label_map = dict()
def __init__(
self,
@@ -18,8 +20,9 @@ class CategoryManager(PostgresManager):
permissions: Collection[Permission] | None = None,
):
super().__init__(pg_config=pg_config, permissions=permissions)
- self.categories: dict[UUIDStr, Category] = dict()
- self.category_label_map: dict[str, Category] = dict()
+ self.categories: dict[UUIDStr, Category] = {}
+ self.category_label_map: dict[str, Category] = {}
+
self.populate_caches()
def populate_caches(self):
diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py
index 517f677..b0aa505 100644
--- a/generalresearch/managers/thl/contest_manager.py
+++ b/generalresearch/managers/thl/contest_manager.py
@@ -1,8 +1,8 @@
from __future__ import annotations
from collections.abc import Collection
-from datetime import datetime, timezone
-from typing import Any, Literal, cast
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any, Literal, cast
from uuid import UUID
import redis
@@ -10,22 +10,14 @@ from pydantic import NonNegativeInt, PositiveInt
from redis import Redis
from generalresearch.managers.base import PostgresManager
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
-)
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.contest import (
ContestPrize,
ContestWinner,
)
-from generalresearch.models.thl.contest.contest import (
- Contest,
- ContestUserView,
-)
+from generalresearch.models.thl.contest.contest_entry import ContestEntry
from generalresearch.models.thl.contest.definitions import (
+ ContestEntryType,
ContestStatus,
ContestType,
)
@@ -41,19 +33,29 @@ from generalresearch.models.thl.contest.leaderboard import (
LeaderboardContestUserView,
)
from generalresearch.models.thl.contest.milestone import (
- ContestEntryTrigger,
MilestoneContest,
MilestoneEntry,
MilestoneUserView,
)
from generalresearch.models.thl.contest.raffle import (
- ContestEntry,
- ContestEntryType,
RaffleContest,
RaffleUserView,
)
from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.user_manager.user_manager import (
+ UserManager,
+ )
+ from generalresearch.models.thl.contest.contest import (
+ Contest,
+ ContestUserView,
+ )
+ from generalresearch.models.thl.contest.milestone import ContestEntryTrigger
+
CONTEST_SELECT = """
c.id,
c.uuid::uuid,
@@ -173,7 +175,7 @@ class ContestBaseManager(PostgresManager):
except ValueError as e:
if e.args[0] == "Contest not found":
return None
- raise e
+ raise
@staticmethod
def make_filter_str(
@@ -187,7 +189,7 @@ class ContestBaseManager(PostgresManager):
has_participants: bool | None = None,
) -> tuple[str, dict[str, Any]]:
filters = []
- params = dict()
+ params = {}
if product_id:
params["product_id"] = product_id
@@ -199,10 +201,10 @@ class ContestBaseManager(PostgresManager):
params["contest_type"] = contest_type.value
filters.append("contest_type = %(contest_type)s")
if starts_at_before is True:
- params["starts_at"] = datetime.now(tz=timezone.utc)
+ params["starts_at"] = datetime.now(tz=UTC)
filters.append("starts_at < %(starts_at)s")
elif starts_at_before:
- assert starts_at_before.tzinfo == timezone.utc
+ assert starts_at_before.tzinfo == UTC
params["starts_at"] = starts_at_before
filters.append("starts_at < %(starts_at)s")
if name is not None:
@@ -681,7 +683,7 @@ class RaffleContestManager(ContestBaseManager):
raise ContestError(msg)
if contest.entry_type == ContestEntryType.CASH:
- tx = ledger_manager.create_tx_user_enter_contest(
+ ledger_manager.create_tx_user_enter_contest(
contest_uuid=contest.uuid, contest_entry=entry
)
@@ -822,13 +824,11 @@ class MilestoneContestManager(ContestBaseManager):
if decision:
contest.update(
status=ContestStatus.COMPLETED,
- ended_at=datetime.now(tz=timezone.utc),
+ ended_at=datetime.now(tz=UTC),
end_reason=reason,
)
self.end_milestone_contest(contest)
- return None
-
def enter_contest_db_work_milestone(
self, contest: MilestoneUserView, user: User, incr: PositiveInt
) -> MilestoneEntry:
@@ -1053,7 +1053,7 @@ class ContestManager(
) -> NonNegativeInt:
contests_closed = 0
for contest in contests:
- should_end, reason = contest.should_end()
+ should_end, _ = contest.should_end()
if should_end:
if hasattr(contest, "redis_client"):
contest.redis_client = redis_client
diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py
index e1143c2..93914c3 100644
--- a/generalresearch/managers/thl/ipinfo.py
+++ b/generalresearch/managers/thl/ipinfo.py
@@ -3,9 +3,10 @@ from __future__ import annotations
import ipaddress
from collections.abc import Collection
from decimal import Decimal
+from typing import TYPE_CHECKING
import faker
-import pymysql
+from grip_client.enums import AccessType
from more_itertools import chunked
from psycopg import Cursor
from pydantic import PositiveInt
@@ -14,18 +15,19 @@ from generalresearch.managers.base import (
PostgresManager,
PostgresManagerWithRedis,
)
-from generalresearch.models.custom_types import (
- CountryISOLike,
- IPvAnyAddressStr,
-)
from generalresearch.models.thl.ipinfo import (
GeoIPInformation,
IPGeoname,
IPInformation,
normalize_ip,
)
-from generalresearch.models.thl.maxmind.definitions import UserType
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.models.custom_types import (
+ CountryISOLike,
+ IPvAnyAddressStr,
+ )
+ from generalresearch.pg_helper import PostgresConfig
fake = faker.Faker()
@@ -156,17 +158,15 @@ class IPGeonameManager(PostgresManager):
if len(filter_ids) == 0:
return []
- with self.pg_config.make_connection() as sql_connection:
- sql_connection: pymysql.Connection
- with sql_connection.cursor() as c:
- res = []
- for chunk in chunked(filter_ids, 500):
- res.extend(
- self.fetch_geoname_ids_(
- c=c,
- filter_ids=chunk,
- )
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ res = []
+ for chunk in chunked(filter_ids, 500):
+ res.extend(
+ self.fetch_geoname_ids_(
+ c=c,
+ filter_ids=chunk,
)
+ )
return res
def fetch_geoname_ids_(
@@ -246,7 +246,7 @@ class IPInformationManager(PostgresManager):
network: str | None = None,
organization: str | None = None,
static_ip_score: float | None = None,
- user_type: UserType | None = None,
+ user_type: AccessType | None = None,
postal_code: str | None = None,
latitude: Decimal | None = None,
longitude: Decimal | None = None,
@@ -423,7 +423,9 @@ class IPInformationManager(PostgresManager):
FROM thl_ipinformation
WHERE updated >= NOW() - INTERVAL '12 hours'
"""
- denominator = list(pg_config.execute_sql_query(query=query))[0]["denominator"]
+ denominator = next(iter(pg_config.execute_sql_query(query=query)))[
+ "denominator"
+ ]
if denominator == 0:
pass
@@ -509,7 +511,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
res = [GeoIPInformation.model_validate_json(raw) for raw in res if raw]
gs = {x.ip: x for x in res}
- res2 = dict()
+ res2 = {}
for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items():
if normalized_ip not in gs:
# try the non-normalized (remove me also 28 days from 2025-11-15)
@@ -633,16 +635,14 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
if len(ips) == 0:
return {}
- with self.pg_config.make_connection() as sql_connection:
- sql_connection: pymysql.Connection
- with sql_connection.cursor() as c:
- res = {}
- for chunk in chunked(ips, 500):
- inner = self.get_mysql_multi_chunk(
- c=c,
- ips=chunk,
- )
- res.update(inner)
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ res = {}
+ for chunk in chunked(ips, 500):
+ inner = self.get_mysql_multi_chunk(
+ c=c,
+ ips=chunk,
+ )
+ res.update(inner)
return res
def get_mysql_multi_chunk(
@@ -719,7 +719,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
gs = [GeoIPInformation.from_mysql(i) for i in res]
gs = {g.ip: g for g in gs}
- res2 = dict()
+ res2 = {}
for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items():
if normalized_ip not in gs:
diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py
index a457f30..38398b1 100644
--- a/generalresearch/managers/thl/ledger_manager/conditions.py
+++ b/generalresearch/managers/thl/ledger_manager/conditions.py
@@ -1,27 +1,29 @@
from __future__ import annotations
import logging
-from datetime import datetime, timedelta, timezone
-from typing import TYPE_CHECKING, Callable
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING
from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.session import Session, Wall
-from generalresearch.models.thl.user import User
-
-logging.basicConfig()
-logger = logging.getLogger("LedgerManager")
-logger.setLevel(logging.INFO)
if TYPE_CHECKING:
+
from generalresearch.managers.thl.ledger_manager.ledger import (
LedgerManager,
)
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
)
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
+
+logging.basicConfig()
+logger = logging.getLogger("LedgerManager")
+logger.setLevel(logging.INFO)
def generate_condition_mp_payment(wall: Wall) -> Callable[..., bool]:
@@ -73,7 +75,7 @@ def generate_condition_bp_payout(
skip_one_per_day_check: bool = False,
skip_wallet_balance_check: bool = False,
) -> Callable[..., tuple[bool, str]]:
- created = datetime.now(tz=timezone.utc)
+ created = datetime.now(tz=UTC)
def _condition(
lm: ThlLedgerManager,
diff --git a/generalresearch/managers/thl/ledger_manager/exceptions.py b/generalresearch/managers/thl/ledger_manager/exceptions.py
index b79c153..48c102b 100644
--- a/generalresearch/managers/thl/ledger_manager/exceptions.py
+++ b/generalresearch/managers/thl/ledger_manager/exceptions.py
@@ -11,7 +11,6 @@ class LedgerTransactionCreateError(Exception):
Ledger transaction creation failed
"""
- pass
class LedgerTransactionCreateLockError(LedgerTransactionCreateError):
@@ -19,7 +18,6 @@ class LedgerTransactionCreateLockError(LedgerTransactionCreateError):
Ledger transaction creation failed because we could not acquire a lock
"""
- pass
class LedgerTransactionReleaseLockError(LedgerTransactionCreateError):
@@ -29,7 +27,6 @@ class LedgerTransactionReleaseLockError(LedgerTransactionCreateError):
back-populate as in sentry I see this very rarely.
"""
- pass
class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError):
@@ -38,7 +35,6 @@ class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError):
tx was already set
"""
- pass
class LedgerTransactionConditionFailedError(LedgerTransactionCreateError):
@@ -46,4 +42,3 @@ class LedgerTransactionConditionFailedError(LedgerTransactionCreateError):
We tried to create a transaction but the condition check failed.
"""
- pass
diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py
index a344430..3a02cdf 100644
--- a/generalresearch/managers/thl/ledger_manager/ledger.py
+++ b/generalresearch/managers/thl/ledger_manager/ledger.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import logging
from collections import defaultdict
from collections.abc import Callable, Collection
-from datetime import datetime, timedelta, timezone
-from typing import Any
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING, Any
from uuid import UUID
import redis
@@ -13,7 +13,6 @@ from pydantic import AwareDatetime, NonNegativeInt, PositiveInt
from redis.exceptions import LockError, LockNotOwnedError
from generalresearch.currency import LedgerCurrency
-from generalresearch.managers import parse_order_by
from generalresearch.managers.base import (
Permission,
PostgresManager,
@@ -28,16 +27,19 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerTransactionFlagAlreadyExistsError,
LedgerTransactionReleaseLockError,
)
+from generalresearch.managers.utils import parse_order_by
from generalresearch.models.custom_types import UUIDStr, check_valid_uuid
from generalresearch.models.thl.ledger import (
LedgerAccount,
LedgerEntry,
LedgerTransaction,
- UserLedgerTransactionType,
UserLedgerTransactionTypesSummary,
)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.ledger import UserLedgerTransactionType
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
logging.basicConfig()
logger = logging.getLogger("LedgerManager")
@@ -104,8 +106,8 @@ class LedgerManagerBasePostgres(PostgresManager, RedisManager):
filters = []
params = {}
if time_start or time_end:
- time_end = time_end or datetime.now(tz=timezone.utc)
- time_start = time_start or datetime(2017, 1, 1, tzinfo=timezone.utc)
+ time_end = time_end or datetime.now(tz=UTC)
+ time_start = time_start or datetime(2017, 1, 1, tzinfo=UTC)
assert time_start.tzinfo.utcoffset(time_start) == timedelta()
assert time_end.tzinfo.utcoffset(time_end) == timedelta()
filters.append("lt.created BETWEEN %(time_start)s AND %(time_end)s")
@@ -149,9 +151,9 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
)
if metadata is None:
- metadata = dict()
+ metadata = {}
if created is None:
- created = datetime.now(tz=timezone.utc)
+ created = datetime.now(tz=UTC)
t = LedgerTransaction(
created=created,
@@ -436,7 +438,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
)
}
else:
- metadata = dict()
+ metadata = {}
entries = [
LedgerEntry(
@@ -456,7 +458,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
id=row["transaction_id"],
entries=entries,
metadata=metadata,
- created=row["created"].replace(tzinfo=timezone.utc),
+ created=row["created"].replace(tzinfo=UTC),
ext_description=row["ext_description"],
tag=row["tag"],
)
@@ -757,7 +759,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
"""
- tx_ids = set([tx.id for tx in transactions])
+ tx_ids = {tx.id for tx in transactions}
res = self.pg_config.execute_sql_query(
query="""
SELECT
@@ -789,7 +791,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
from the database.
"""
- tx_ids = set([tx.id for tx in transactions])
+ tx_ids = {tx.id for tx in transactions}
res = self.pg_config.execute_sql_query(
query="""
SELECT tx_meta.id
@@ -799,7 +801,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
params=[list(tx_ids)],
)
- return set([i["id"] for i in res])
+ return {i["id"] for i in res}
class LedgerEntryManager(LedgerManagerBasePostgres):
@@ -809,7 +811,7 @@ class LedgerEntryManager(LedgerManagerBasePostgres):
def get_tx_entries_by_txs(
self, transactions: list[LedgerTransaction]
) -> list[LedgerEntry]:
- tx_ids = set([tx.id for tx in transactions])
+ tx_ids = {tx.id for tx in transactions}
tx_entries = self.pg_config.execute_sql_query(
query="""
SELECT
@@ -1147,4 +1149,5 @@ class LedgerManager(
}
for k, v in d.items():
v["total"] = (v["debit"] - v["credit"]) * k.normal_balance.value
+
return d
diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py
index 295385d..f1a76c4 100644
--- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py
+++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py
@@ -3,13 +3,14 @@ from __future__ import annotations
import logging
from collections import defaultdict
from collections.abc import Callable, Collection
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from typing import TYPE_CHECKING
from uuid import UUID
import numpy as np
import pandas as pd
+from generalresearch.models.thl.wallet.definitions import PayoutType
from pydantic import AwareDatetime, PositiveInt
from generalresearch.config import (
@@ -31,14 +32,12 @@ from generalresearch.managers.thl.ledger_manager.ledger import (
LedgerManager,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.contest.contest import Contest
from generalresearch.models.thl.contest.definitions import (
ContestPrizeKind,
ContestType,
)
from generalresearch.models.thl.contest.milestone import MilestoneContest
from generalresearch.models.thl.contest.raffle import (
- ContestEntry,
ContestEntryType,
RaffleContest,
)
@@ -59,7 +58,6 @@ from generalresearch.models.thl.payout_format import format_payout_format
from generalresearch.models.thl.product import Product
from generalresearch.models.thl.session import Session, Status, Wall
from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.user_wallet import (
UserDisplayedWalletBalance,
UserLedgerWallet,
@@ -276,10 +274,10 @@ class ThlLedgerManager(LedgerManager):
time_end: datetime | None = None,
):
if time_start is None:
- time_start = datetime(year=2017, month=1, day=1, tzinfo=timezone.utc)
+ time_start = datetime(year=2017, month=1, day=1, tzinfo=UTC)
if time_end is None:
- time_end = datetime.now(tz=timezone.utc)
+ time_end = datetime.now(tz=UTC)
assert all(isinstance(item, str) for item in account_uuids), (
"Must pass account_uuid as str"
@@ -850,9 +848,7 @@ class ThlLedgerManager(LedgerManager):
if skip_one_per_day_check or skip_wallet_balance_check:
skip_flag_check = True
- assert datetime.now(tz=timezone.utc) > created, (
- "created cannot be in the future"
- )
+ assert datetime.now(tz=UTC) > created, "created cannot be in the future"
f = lambda: self.create_tx_bp_payout_(
product=product,
amount=amount,
@@ -886,12 +882,13 @@ class ThlLedgerManager(LedgerManager):
created: datetime,
) -> LedgerTransaction:
+ tx_type = TransactionType.BP_PAYOUT
metadata = {
- tmc.TX_TYPE: TransactionType.BP_PAYOUT,
+ tmc.TX_TYPE: tx_type,
tmc.EVENT: payoutevent_uuid,
}
- # This tag might will uniquely identify this tx
- tag = f"{self.currency.value}:bp_payout:{payoutevent_uuid}"
+ # This tag will uniquely identify this tx
+ tag = f"{self.currency.value}:{tx_type.value}:{payoutevent_uuid}"
cash_account = self.get_account_cash()
bp_wallet_account = self.get_account_or_create_bp_wallet(product)
@@ -956,9 +953,7 @@ class ThlLedgerManager(LedgerManager):
:param skip_flag_check: If True, we skip the flag check to allow
for retry of a failed previous call.
"""
- assert datetime.now(tz=timezone.utc) > created, (
- "created cannot be in the future"
- )
+ assert datetime.now(tz=UTC) > created, "created cannot be in the future"
assert isinstance(amount, int)
assert isinstance(amount, USDCent)
@@ -1984,7 +1979,7 @@ class ThlLedgerManager(LedgerManager):
"Can't get wallet balance on non-managed account."
)
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
wallet = self.get_account_or_create_user_wallet(user)
if user.product_id == JAMES_BILLINGS_BPID:
assert since_days_ago is None
@@ -2014,7 +2009,7 @@ class ThlLedgerManager(LedgerManager):
After 3 days, about 25% of all "future" recons have happened,
7 days: 50%, 14 days: 75%, till end of next month: 100%.
"""
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# The redeemable balance can NOT ever be more than the actual user_wallet_balance
# Sum up the redeemable amount for each complete
diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py
index 1f3742e..59dc2e9 100644
--- a/generalresearch/managers/thl/payout.py
+++ b/generalresearch/managers/thl/payout.py
@@ -1,14 +1,13 @@
from __future__ import annotations
-from collections import defaultdict
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
-from time import sleep
-from typing import Any
-from uuid import UUID, uuid4
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any
+from uuid import uuid4
import numpy as np
import pandas as pd
+import psycopg
from psycopg import sql
from pydantic import AwareDatetime, NonNegativeInt, PositiveInt
@@ -17,30 +16,38 @@ from generalresearch.decorators import LOG
from generalresearch.managers.base import (
PostgresManagerWithRedis,
)
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
+from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerTransactionConditionFailedError,
+ LedgerTransactionReleaseLockError,
)
-from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-from generalresearch.models.gr.business import Business
from generalresearch.models.thl.definitions import PayoutStatus
from generalresearch.models.thl.ledger import (
Direction,
- LedgerAccount,
OrderBy,
+ TransactionType,
)
from generalresearch.models.thl.payout import (
BrokerageProductPayoutEvent,
BusinessPayoutEvent,
+ BusinessPayoutEventCreate,
PayoutEvent,
UserPayoutEvent,
)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailOrderData,
CashoutRequestInfo,
)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.thl.ledger import LedgerAccount
+ from generalresearch.models.thl.product import Product
class PayoutEventManager(PostgresManagerWithRedis):
@@ -51,25 +58,6 @@ class PayoutEventManager(PostgresManagerWithRedis):
"""
- def set_account_lookup_table(self, thl_lm: ThlLedgerManager) -> None:
- """This needs to run from grl-flow or from somewhere that has thl-redis
- access
- """
-
- res = self.pg_config.execute_sql_query(query=f"""
- SELECT uuid, reference_uuid
- FROM ledger_account
- WHERE qualified_name LIKE '{thl_lm.currency.value}:bp_wallet:%'
- """)
- account_to_product = {i["uuid"]: i["reference_uuid"] for i in res}
- product_to_account = {i["reference_uuid"]: i["uuid"] for i in res}
-
- rc = self.redis_client
- rc.hset(name="pem:account_to_product", mapping=account_to_product)
- rc.hset(name="pem:product_to_account", mapping=product_to_account)
-
- return None
-
def get_by_uuid(self, pe_uuid: UUIDStr) -> PayoutEvent:
res = self.pg_config.execute_sql_query(
query="""
@@ -100,7 +88,7 @@ class PayoutEventManager(PostgresManagerWithRedis):
order_data = order_data if order_data is not None else payout_event.order_data
payout_event.update(status=status, ext_ref_id=ext_ref_id, order_data=order_data)
- d = payout_event.model_dump_mysql()
+ d = payout_event.model_dump_postgres()
query = sql.SQL("""
UPDATE event_payout SET
status = %(status)s,
@@ -111,14 +99,13 @@ class PayoutEventManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=d)
- assert (
- c.rowcount == 1
- ), "Nothing was updated! Are you sure this payout_event exists?"
+ assert c.rowcount == 1, (
+ "Nothing was updated! Are you sure this payout_event exists?"
+ )
conn.commit()
class UserPayoutEventManager(PayoutEventManager):
-
def get_by_uuid(self, pe_uuid: UUIDStr) -> UserPayoutEvent:
res = self.pg_config.execute_sql_query(
@@ -157,7 +144,7 @@ class UserPayoutEventManager(PayoutEventManager):
# the purposes of returning to the user.
pe = self.get_by_uuid(pe_uuid=pe_uuid)
- transaction_info = dict()
+ transaction_info = {}
order: dict[str, Any] = pe.order_data
if pe.payout_type == PayoutType.TANGO and pe.status == PayoutStatus.COMPLETE:
reward = order["reward"]
@@ -304,7 +291,7 @@ class UserPayoutEventManager(PayoutEventManager):
account_reference_uuid=account_reference_uuid,
cashout_method_uuid=cashout_method_uuid,
description=description,
- created=created or datetime.now(tz=timezone.utc),
+ created=created or datetime.now(tz=UTC),
amount=amount,
status=status or PayoutStatus.PENDING,
ext_ref_id=ext_ref_id,
@@ -312,7 +299,7 @@ class UserPayoutEventManager(PayoutEventManager):
request_data=request_data or {},
order_data=order_data,
)
- d = payout_event.model_dump_mysql()
+ d = payout_event.model_dump_postgres()
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
@@ -345,43 +332,27 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
def get_by_uuid(
self,
pe_uuid: UUIDStr,
- # --- Support resources ---
- account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
) -> BrokerageProductPayoutEvent:
res = self.pg_config.execute_sql_query(
query="""
- SELECT ep.uuid,
- ep.debit_account_uuid,
- ep.cashout_method_uuid,
+ SELECT ep.uuid, ep.debit_account_uuid, ep.cashout_method_uuid,
ep.created, ep.amount, ep.status, ep.ext_ref_id, ep.payout_type,
ep.request_data::jsonb,
- ep.order_data::jsonb
+ ep.order_data::jsonb,
+ la.reference_uuid as product_id
FROM event_payout AS ep
+ JOIN ledger_account la on la.uuid = debit_account_uuid
WHERE ep.uuid = %s
""",
params=[pe_uuid],
)
assert len(res) == 1, f"{pe_uuid} expected 1 result, got {len(res)}"
-
- d = res[0]
-
- # This isn't really need for creation... but we're doing it so that
- # it can return back a full BrokerageProductPayoutEvent instance
- if account_product_mapping is None:
- rc = self.redis_client
- account_product_mapping: dict = rc.hgetall(name="pem:account_to_product")
- assert isinstance(account_product_mapping, dict)
-
- d["product_id"] = account_product_mapping[d["debit_account_uuid"]]
-
- return BrokerageProductPayoutEvent.model_validate(d)
+ return BrokerageProductPayoutEvent.model_validate(res[0])
@staticmethod
def check_for_ledger_tx(
thl_ledger_manager: ThlLedgerManager,
- product_id: UUIDStr,
- amount: USDCent,
payout_event: BrokerageProductPayoutEvent,
) -> bool:
"""
@@ -394,6 +365,9 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
are found, and raises a ValueError if something is inconsistent.
"""
tag = f"{thl_ledger_manager.currency.value}:bp_payout:{payout_event.uuid}"
+ amount = USDCent(payout_event.amount)
+ product_id = payout_event.product_id
+
txs = thl_ledger_manager.get_tx_by_tag(tag)
if not txs:
@@ -415,7 +389,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
product_uuid=product_id
)
- entry = [x for x in tx.entries if x.direction == Direction.DEBIT][0]
+ entry = next(x for x in tx.entries if x.direction == Direction.DEBIT)
if entry.account_uuid != bp_wallet_account.uuid:
raise ValueError(
f"Found existing tx with tag: {tag}, but for a different account!"
@@ -423,172 +397,80 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
return True
- def create(
- self,
- uuid: UUIDStr | None = None,
- debit_account_uuid: UUIDStr | None = None,
- created: AwareDatetimeISO = None,
- amount: PositiveInt = None,
- status: PayoutStatus | None = None,
- ext_ref_id: str | None = None,
- payout_type: PayoutType = None,
- request_data: dict[str, Any] | None = None,
- order_data: dict[str, Any] | CashMailOrderData | None = None,
- # --- Support resources ---
- account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
- ) -> BrokerageProductPayoutEvent:
-
- if request_data is None:
- request_data = dict()
-
- # This isn't really need for creation... but we're doing it so that
- # it can return back a full BrokerageProductPayoutEvent instance
- if account_product_mapping is None:
- rc = self.redis_client
- account_product_mapping: dict = rc.hgetall(name="pem:account_to_product")
- assert isinstance(account_product_mapping, dict)
- product_id = account_product_mapping[debit_account_uuid]
-
- bp_payout_event = BrokerageProductPayoutEvent(
- uuid=uuid or uuid4().hex,
- debit_account_uuid=debit_account_uuid,
- cashout_method_uuid=self.CASHOUT_METHOD_UUID,
- created=created or datetime.now(tz=timezone.utc),
- amount=amount,
- status=status,
- ext_ref_id=ext_ref_id,
- payout_type=payout_type,
- request_data=request_data,
- order_data=order_data,
- product_id=product_id,
- )
- d = bp_payout_event.model_dump_mysql()
-
- self.pg_config.execute_write(
- query="""
- INSERT INTO event_payout (
- uuid, debit_account_uuid, created, cashout_method_uuid, amount,
- status, ext_ref_id, payout_type, order_data, request_data
- ) VALUES (
- %(uuid)s, %(debit_account_uuid)s, %(created)s,
- %(cashout_method_uuid)s, %(amount)s, %(status)s,
- %(ext_ref_id)s, %(payout_type)s, %(order_data)s,
- %(request_data)s
- );
- """,
- params=d,
- )
-
- return bp_payout_event
-
def filter_by(
self,
- reference_uuid: str | None = None,
ext_ref_id: str | None = None,
debit_account_uuids: Collection[UUIDStr] | None = None,
amount: int | None = None,
created: datetime | None = None,
created_after: datetime | None = None,
- product_ids: str | None = None,
- bp_user_ids: Collection[str] | None = None,
+ product_ids: Collection[str] | None = None,
cashout_types: Collection[PayoutType] | None = None,
statuses: Collection[PayoutStatus] | None = None,
) -> list[BrokerageProductPayoutEvent]:
- """Try to retrieve payout events by the product_id/user_uuid, amount,
- and optionally timestamp.
+ """Try to retrieve BP payout events.
WARNING: This is only on the "payout events" table and nothing to
- do with the Ledger itself. Therefore, the product_ids query
- doesn't return Brokerage Product Payouts (the ACH or Wire events
- to Suppliers) as part of the query.
+ do with the Ledger itself
- *** IT IS ONLY FOR USER PAYOUTS ***
-
- Note: what used to be in thl-grpcs "ListCashoutRequests" calling
- "list_cashout_requests" was merged into this.
+ *** IT IS ONLY FOR Brokerage Product PAYOUTS ***
"""
- args = []
+ params = {}
filters = []
- if reference_uuid:
- # This could be a product_id or a user_uuid
- filters.append("la.reference_uuid = %s")
- args.append(reference_uuid)
if ext_ref_id:
- # This is transaction id for tracking ACH/Wires with a banking
- # institution
- filters.append("ep.ext_ref_id = %s")
- args.append(ext_ref_id)
+ # This is transaction id for tracking ACH/Wires with a banking institution
+ filters.append("ep.ext_ref_id = %(ext_ref_id)s")
+ params["ext_ref_id"] = ext_ref_id
if debit_account_uuids:
- # Or we could use the bp_wallet or user_wallet's account uuid
- # instead of looking up by the product/user
- filters.append("ep.debit_account_uuid = ANY(%s)")
- args.append(debit_account_uuids)
+ # Or we could use the bp_wallet's account uuid
+ # instead of looking up by the product
+ filters.append("ep.debit_account_uuid = ANY(%(debit_account_uuids)s)")
+ params["debit_account_uuids"] = debit_account_uuids
if amount:
- filters.append("ep.amount = %s")
- args.append(amount)
+ filters.append("ep.amount = %(amount)s")
+ params["amount"] = amount
if created:
- filters.append("ep.created = %s")
- args.append(created.replace(tzinfo=None))
+ filters.append("ep.created = %(created)s")
+ params["created"] = created
if created_after:
- filters.append("ep.created >= %s")
- args.append(created_after.replace(tzinfo=None))
- if product_ids:
- filters.append("product_id = ANY(%s)")
- args.append(product_ids)
- if bp_user_ids:
- filters.append("product_user_id = ANY(%s)")
- args.append(bp_user_ids)
- if cashout_types:
- filters.append("payout_type = ANY(%s)")
- args.append([x.value for x in cashout_types])
- if statuses:
- filters.append("status = ANY(%s)")
- args.append([x.value for x in statuses])
+ filters.append("ep.created >= %(created_after)s")
+ params["created_after"] = created_after
+ if product_ids is not None:
+ filters.append("la.reference_uuid = ANY(%(product_ids)s)")
+ params["product_ids"] = product_ids
+ if cashout_types is not None:
+ filters.append("payout_type = ANY(%(cashout_types)s)")
+ params["cashout_types"] = [x.value for x in cashout_types]
+ if statuses is not None:
+ filters.append("status = ANY(%(statuses)s)")
+ params["statuses"] = [x.value for x in statuses]
assert len(filters) > 0, "must pass at least 1 filter"
filter_str = " AND ".join(filters)
+ params["cashout_method_uuid"] = self.CASHOUT_METHOD_UUID
res = self.pg_config.execute_sql_query(
query=f"""
- SELECT ep.uuid,
- ep.debit_account_uuid,
- ep.cashout_method_uuid,
- ep.created,
- ep.amount, ep.status, ep.ext_ref_id, ep.payout_type,
+ SELECT ep.uuid, ep.debit_account_uuid, ep.cashout_method_uuid,
+ ep.created, ep.amount, ep.status, ep.ext_ref_id,
+ ep.payout_type, ep.supplier_payout_id,
ep.request_data::jsonb, ep.order_data::jsonb,
ac.name as description,
- la.reference_type as account_reference_type,
- la.reference_uuid as account_reference_uuid
+ la.reference_uuid as product_id
FROM event_payout AS ep
LEFT JOIN accounting_cashoutmethod AS ac
ON ep.cashout_method_uuid = ac.id
LEFT JOIN ledger_account AS la
ON la.uuid = ep.debit_account_uuid
- LEFT JOIN thl_user u
- ON la.reference_uuid = u.uuid
- WHERE cashout_method_uuid = '{self.CASHOUT_METHOD_UUID}'
+ WHERE cashout_method_uuid = %(cashout_method_uuid)s
+ AND la.reference_type = 'bp'
AND {filter_str}
""",
- params=args,
+ params=params,
)
-
- rc = self.redis_client
- account_product_mapping = rc.hgetall(name="pem:account_to_product")
-
pes = []
- for d in res:
- for k in [
- "uuid",
- "debit_account_uuid",
- "account_reference_uuid",
- "cashout_method_uuid",
- ]:
- if d[k] is not None:
- d[k] = UUID(d[k]).hex
-
- d["product_id"] = account_product_mapping[d["debit_account_uuid"]]
- pes.append(BrokerageProductPayoutEvent.model_validate(d))
-
+ for row in res:
+ pes.append(BrokerageProductPayoutEvent.model_validate(row))
return pes
def get_bp_payout_events_for_accounts(
@@ -601,159 +483,64 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
def get_bp_bp_payout_events_for_products(
self,
- thl_ledger_manager: ThlLedgerManager,
product_uuids: Collection[UUIDStr],
order_by: OrderBy | None = OrderBy.ASC,
) -> list[BrokerageProductPayoutEvent]:
"""This is a terrible name, but it returns the
BPPayoutEvent model type rather than a list of PayoutEvents.
- We do this for the Supplier centric APIs where they don't know,
+ We do this for the Supplier-centric APIs where they don't know
or care about the underlying ledger account structure.
"""
assert len(product_uuids) > 0, "Must provide product_uuids"
- accounts = thl_ledger_manager.get_accounts_bp_wallet_for_products(
- product_uuids=product_uuids
- )
-
- assert len(accounts) == len(product_uuids), "Unequal Product & Account lists"
-
- rc = self.redis_client
- account_product_mapping = rc.hgetall(name="pem:account_to_product")
+ order_by = order_by or OrderBy.ASC
- payout_events: list[BrokerageProductPayoutEvent] = (
- self.get_bp_payout_events_for_accounts(
- accounts=accounts,
- )
+ payout_events = self.filter_by(
+ product_ids=product_uuids,
+ cashout_types=[PayoutType.ACH],
)
-
- return BrokerageProductPayoutEvent.from_payout_events(
- payout_events=payout_events,
- account_product_mapping=account_product_mapping,
- order_by=order_by,
+ payout_events = sorted(
+ payout_events, key=lambda x: x.created, reverse=order_by == OrderBy.DESC
)
+ return payout_events
def retry_create_bp_payout_event_tx(
self,
thl_ledger_manager: ThlLedgerManager,
product: Product,
- payout_event_uuid: UUIDStr,
- skip_wallet_balance_check: bool = False,
- skip_one_per_day_check: bool = False,
+ bp_pe: BrokerageProductPayoutEvent,
) -> BrokerageProductPayoutEvent:
"""
If a create_bp_payout_event call fails, this can be called with
the associated payoutevent.
"""
- bp_pe: BrokerageProductPayoutEvent = self.get_by_uuid(payout_event_uuid)
- assert bp_pe.status == PayoutStatus.FAILED, "Only use this on failed payouts"
- created = bp_pe.created
-
- assert not self.check_for_ledger_tx(
- thl_ledger_manager=thl_ledger_manager,
- payout_event=bp_pe,
- product_id=bp_pe.product_id,
- amount=bp_pe.amount_usd,
- ), "Transaction exists! You should mark the payout event status as complete"
-
- return self._create_tx_bp_payout_from_payout_event(
- thl_ledger_manager=thl_ledger_manager,
- bp_pe=bp_pe,
- product=product,
- amount=bp_pe.amount_usd,
- created=created,
- skip_one_per_day_check=skip_one_per_day_check,
- skip_wallet_balance_check=skip_wallet_balance_check,
- )
-
- def create_bp_payout_event(
- self,
- thl_ledger_manager: ThlLedgerManager,
- product: Product,
- amount: USDCent,
- payout_type: PayoutType = PayoutType.ACH,
- ext_ref_id: str | None = None,
- created: AwareDatetime | None = None,
- skip_wallet_balance_check: bool = False,
- skip_one_per_day_check: bool = False,
- ) -> BrokerageProductPayoutEvent:
- """This should be called when a BP is paid out money from their
- wallet. Typically, this is an ACH payment. This function creates
- the PayoutEvent and the Ledger entries.
-
- :param thl_ledger_manager:
- :param product: The BP being paid. Assuming we're paying them out
- of the balance of their USD wallet account.
- :param amount: We're assuming everything is in USD, and we're
- paying out a USD currency account. We could theoretically also
- pay, for e.g. a Bitcoin account with a bitcoin transfer, but
- this is not supported for now.
- :param payout_type: PayoutType. default ACH
- :param cashout_method_uuid: The entry in the
- accounting_cashoutmethod table that records payment method
- details. By default, the generic ACH cashout method (that has
- no actual banking details).
-
- :param ext_ref_id: This is a unique ID for the Supplier Payment.
- Typically it'll be from JP Morgan Chase, but may also just be
- random if we can retrieve anything
-
- :param created:
-
- :param skip_wallet_balance_check: By default, this will fail unless
- the BP's wallet actually has the amount requested.
-
- :param skip_one_per_day_check: Safety mechanism, checks if there
- has already been a payout to this wallet in the past 24 hours.
-
- :return:
- """
-
- assert isinstance(amount, USDCent), "Must provide a USDCent"
+ assert bp_pe.status in {
+ PayoutStatus.FAILED,
+ PayoutStatus.PENDING,
+ }, "Only use this on pending or failed payouts"
- if created:
- # Try to do a quick dupe check first before we create the payout event
- pes = self.filter_by(
- reference_uuid=product.id, amount=amount, created=created
+ if self.check_for_ledger_tx(
+ thl_ledger_manager=thl_ledger_manager, payout_event=bp_pe
+ ):
+ LOG.warning(
+ f"Transaction for {bp_pe.uuid=} {bp_pe.product_id=} already exists! "
+ f"Marking the payout event status as complete."
)
- if len(pes) > 0:
- raise ValueError(f"Payout event already exists!: {pes}")
+ self.update(payout_event=bp_pe, status=PayoutStatus.COMPLETE)
+ return bp_pe
- if created is None:
- created = datetime.now(tz=timezone.utc)
-
- # TODO: Explain why we're doing this. Why is it important to have
- # Payout Events when the ledger has everything that should be
- # needed.
- bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
-
- bp_pe: BrokerageProductPayoutEvent = self.create(
- debit_account_uuid=bp_wallet.uuid,
- payout_type=payout_type,
- amount=amount,
- ext_ref_id=ext_ref_id,
- created=created,
- status=PayoutStatus.PENDING,
- )
- return self._create_tx_bp_payout_from_payout_event(
+ return self.create_tx_bp_payout_from_payout_event(
thl_ledger_manager=thl_ledger_manager,
bp_pe=bp_pe,
product=product,
- amount=amount,
- created=created,
- skip_one_per_day_check=skip_one_per_day_check,
- skip_wallet_balance_check=skip_wallet_balance_check,
)
- def _create_tx_bp_payout_from_payout_event(
+ def create_tx_bp_payout_from_payout_event(
self,
thl_ledger_manager: ThlLedgerManager,
bp_pe: BrokerageProductPayoutEvent,
product: Product,
- amount: USDCent,
created: AwareDatetime | None = None,
- skip_wallet_balance_check: bool = False,
- skip_one_per_day_check: bool = False,
) -> BrokerageProductPayoutEvent:
"""
This should not be called directly.
@@ -761,31 +548,37 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
Handles exceptions: Check if the ledger tx actually exists or not, and set the
payout event status accordingly.
"""
+ created = created if created else bp_pe.created
try:
thl_ledger_manager.create_tx_bp_payout(
product=product,
- amount=amount,
+ amount=USDCent(bp_pe.amount),
payoutevent_uuid=bp_pe.uuid,
created=created,
- skip_wallet_balance_check=skip_wallet_balance_check,
- skip_one_per_day_check=skip_one_per_day_check,
+ skip_wallet_balance_check=True,
+ skip_one_per_day_check=True,
)
-
- except Exception as e:
- e.pe_uuid = bp_pe.uuid
- if self.check_for_ledger_tx(
- thl_ledger_manager=thl_ledger_manager,
- product_id=product.uuid,
- amount=amount,
- payout_event=bp_pe,
- ):
- LOG.warning(f"Got exception {e} but ledger tx exists! Continuing ... ")
- self.update(payout_event=bp_pe, status=PayoutStatus.COMPLETE)
- return bp_pe
- else:
- LOG.warning(f"Got exception {e}. No ledger tx was created.")
+ except LedgerTransactionConditionFailedError as e:
+ if e.args[0] == "duplicate tag":
+ raise ValueError(f"""Payout event already exists! {e}
+ You are trying to create a tx that already exists. We can't know
+ if this is a new payout event with the same ref id, or you're
+ trying to run the same one twice ... So not setting the existing
+ payout event to FAILED, b/c the existing one is not failed!
+ Doing nothing ...
+ """) from e
+ self.update(payout_event=bp_pe, status=PayoutStatus.FAILED)
+ raise
+ except LedgerTransactionReleaseLockError as e:
+ # Redis error upon lock release. The tx was most likely created.
+ LOG.warning(e)
+ tag = f"{thl_ledger_manager.currency.value}:{TransactionType.BP_PAYOUT.value}:{bp_pe.uuid}"
+ # Check if it was created, and if so, swallow error.
+ if not thl_ledger_manager.get_tx_ids_by_tag(tag=tag):
self.update(payout_event=bp_pe, status=PayoutStatus.FAILED)
- raise e
+ except Exception:
+ self.update(payout_event=bp_pe, status=PayoutStatus.FAILED)
+ raise
self.update(payout_event=bp_pe, status=PayoutStatus.COMPLETE)
return bp_pe
@@ -814,156 +607,176 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
return self.get_bp_payout_events_for_accounts(accounts=accounts)
-class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
+class BusinessPayoutEventManager(PostgresManagerWithRedis):
+ def __init__(self, *arg, **kwargs):
+ super().__init__(*arg, **kwargs)
+ self.bp_pe_manager = BrokerageProductPayoutEventManager(*arg, **kwargs)
- def update_ext_reference_ids(
+ def get_by_ext_ref_id(self, ext_ref_id: str) -> BusinessPayoutEvent:
+ res = self.pg_config.execute_sql_query(
+ """
+ SELECT
+ sp.*,
+ ep.bp_payouts
+ FROM supplier_payout sp
+ JOIN (
+ SELECT
+ ep_inner.supplier_payout_id,
+ jsonb_agg(
+ to_jsonb(ep_inner)
+ || jsonb_build_object('product_id', la.reference_uuid)
+ ORDER BY ep_inner.created
+ ) AS bp_payouts
+ FROM event_payout ep_inner
+ JOIN ledger_account la
+ ON ep_inner.debit_account_uuid = la.uuid
+ GROUP BY ep_inner.supplier_payout_id
+ ) ep ON sp.id = ep.supplier_payout_id
+ WHERE sp.ext_ref_id = %(ext_ref_id)s
+ """,
+ {"ext_ref_id": ext_ref_id},
+ )
+ assert len(res) == 1, f"No Business Payout found with ext ref: {ext_ref_id}"
+ d = res[0]
+ for bp_payout in d["bp_payouts"]:
+ bp_payout["created"] = datetime.fromisoformat(bp_payout["created"])
+ bpe = BusinessPayoutEvent.model_validate(d)
+ assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, (
+ "No BP payouts found for this Business Payout Event. This shouldn't happen!"
+ )
+ return bpe
+
+ def filter_by(
self,
- new_value: str,
- current_value: str | None = None,
- ) -> None:
- """
- There are scenarios where an ACH/Wire payout event was saved with
- a generic or anonymized reference identifier. We may want to be
- able to go back and update all of those transaction IDs.
+ business_uuids: Collection[UUIDStr] | None = None,
+ ) -> list[BusinessPayoutEvent]:
- """
+ params = {}
+ filters = []
+ if business_uuids is not None:
+ filters.append("business_id = ANY(%(business_uuids)s)")
+ params["business_uuids"] = business_uuids
- if current_value is None:
- raise ValueError("Dangerous to do ambiguous updates")
+ assert len(filters) > 0, "must pass at least 1 filter"
+ filter_str = " AND ".join(filters)
- # SELECT first to check that records exist
- res = self.filter_by(ext_ref_id=current_value)
- if len(res) == 0:
- raise Warning("No event_payouts found to UPDATE")
+ res = self.pg_config.execute_sql_query(
+ f"""
+ SELECT
+ sp.*,
+ ep.bp_payouts
+ FROM supplier_payout sp
+ JOIN (
+ SELECT
+ ep_inner.supplier_payout_id,
+ jsonb_agg(
+ to_jsonb(ep_inner)
+ || jsonb_build_object('product_id', la.reference_uuid)
+ ORDER BY ep_inner.created
+ ) AS bp_payouts
+ FROM event_payout ep_inner
+ JOIN ledger_account la
+ ON ep_inner.debit_account_uuid = la.uuid
+ GROUP BY ep_inner.supplier_payout_id
+ ) ep ON sp.id = ep.supplier_payout_id
+ WHERE {filter_str}
+ """,
+ params,
+ )
+ bpes = []
+ for row in res:
+ for bp_payout in row["bp_payouts"]:
+ bp_payout["created"] = datetime.fromisoformat(bp_payout["created"])
+ bpe = BusinessPayoutEvent.model_validate(row)
+ assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, (
+ "No BP payouts found for this Business Payout Event. This shouldn't happen!"
+ )
+ bpes.append(bpe)
+ return bpes
- # As of 2025, no single Business has more than 10,000 Products,
- # leave the limit in as an additional safeguard.
- query = """
- UPDATE event_payout
- SET ext_ref_id = %s
- WHERE ext_ref_id = %s
+ def validate_business_payout_in_ledger(
+ self, ext_ref_id: str, thl_lm: ThlLedgerManager
+ ):
"""
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query=query, params=[new_value, current_value])
- assert c.rowcount < 10000
- conn.commit()
+ Check that there exist ledger TXs for the Brokerage Product payouts
+ for this Business Payout Event.
+ """
+ bpe = self.get_by_ext_ref_id(ext_ref_id=ext_ref_id)
+ tags = [
+ f"{thl_lm.currency.value}:bp_payout:{bp_pe.uuid}"
+ for bp_pe in bpe.bp_payouts
+ ]
+ txs = thl_lm.get_tx_ids_by_tags(tags=tags)
+ assert len(txs) == len(bpe.bp_payouts), (
+ f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!"
+ )
+ return True
- def delete_failed_business_payout(self, ext_ref_id: str, thl_lm: ThlLedgerManager):
+ def resume_failed_business_payout(
+ self, ext_ref_id: str, thl_lm: ThlLedgerManager, pm: ProductManager
+ ):
"""
- Sometimes ACH/Wire payouts fail due to multiple reasons (timeouts,
- Business Product having insufficient funds, etc). This is a utility
- method that finds all event_payouts, and deletes them with all the
- associated:
- (1) Transactions
- (2) Transaction Metadata
- (3) Transaction Entries
-
- and then proceeds to delete them all in reverse order (so there is
- no orphan / FK constraint issues).
+ Sometimes a business payout's BP payouts fail due to multiple reasons
+ (timeouts, BP having insufficient funds, etc). Grab the PENDING
+ BP payout events and retry them.
"""
+ bpe = self.get_by_ext_ref_id(ext_ref_id=ext_ref_id)
+ assert bpe.id
+ assert bpe.bp_payouts
- # (1) Find all by payout_event
- event_payouts = self.filter_by(ext_ref_id=ext_ref_id)
- if len(event_payouts) == 0:
- raise Warning("No event_payouts found to DELETE")
-
- # sum([i["amount"] for i in event_payouts])/100
- event_payout_uuids = [i.uuid for i in event_payouts]
-
- # (2) Find all ledger_transactions
- tags = [f"{thl_lm.currency.value}:bp_payout:{x}" for x in event_payout_uuids]
- transactions = thl_lm.get_txs_by_tags(tags=tags)
- transaction_ids = [tx.id for tx in transactions]
- print("XXX1", transaction_ids)
- # assert len(tags) == len(transactions)
-
- # (3) Find all ledger_transactionmetadata: assert two rows per tx
- tx_metadata_ids = thl_lm.get_tx_metadata_ids_by_txs(transactions=transactions)
- # assert len(tx_metadata) == len(transaction_ids)*2
-
- # (4) Find all ledger_entry: assert two rows per tx
- tx_entries = thl_lm.get_tx_entries_by_txs(transactions=transactions)
- tx_entry_ids = [tx_entry.id for tx_entry in tx_entries]
- # assert len(tx_entry) == len(transaction_ids)*2
-
- # (5) Delete records
-
- # DELETE: tx_entry
- self.pg_config.execute_write(
- query="""
- DELETE
- FROM ledger_entry
- WHERE transaction_id = ANY(%s)
- AND id = ANY(%s)
- """,
- params=[transaction_ids, tx_entry_ids],
- )
+ if all(bp_pe.status == PayoutStatus.COMPLETE for bp_pe in bpe.bp_payouts):
+ try:
+ self.validate_business_payout_in_ledger(
+ ext_ref_id=ext_ref_id, thl_lm=thl_lm
+ )
+ except AssertionError as e:
+ raise AssertionError(
+ f"Business Payout Event {ext_ref_id} is COMPLETE but BP payouts are not in the ledger! {e} "
+ f"This typically shouldn't happen, as if the ledger TX fails, the event_payout "
+ f"status won't be COMPLETE. If it does, set all the bp statuses to FAILED, and "
+ f"then try again. Any that do exist in the ledger will be found and marked COMPLETE."
+ ) from e
+ if bpe.status != PayoutStatus.COMPLETE:
+ self.update_business_payout_event(
+ pk=bpe.id, status=PayoutStatus.COMPLETE
+ )
+ LOG.warning(
+ "All BP payouts complete, setting Business Payout Event status to COMPLETE."
+ )
+ else:
+ LOG.warning(
+ "Nothing to do! Business Payout is COMPLETE and all Brokerage Product payouts are also COMPLETE!"
+ )
+ return
- # DELETE: tx_metadata
- self.pg_config.execute_write(
- query="""
- DELETE
- FROM ledger_transactionmetadata
- WHERE transaction_id = ANY(%s)
- AND id = ANY(%s)
- """,
- params=[transaction_ids, list(tx_metadata_ids)],
- )
+ for bp_pe in bpe.bp_payouts:
+ if bp_pe.status in {PayoutStatus.PENDING, PayoutStatus.FAILED}:
+ LOG.warning(
+ f"Found a {bp_pe.status} BP payout event: {bp_pe.uuid} - retrying ... "
+ )
+ product = pm.get_by_uuid(bp_pe.product_id)
+ self.retry_create_bp_payout_event_tx(
+ thl_ledger_manager=thl_lm, bp_pe=bp_pe, product=product
+ )
+ if bp_pe.status != PayoutStatus.COMPLETE:
+ raise ValueError(f"{bp_pe.uuid} has {bp_pe.status=}. Please check me.")
- # DELETE: transactions
- self.pg_config.execute_write(
- query="""
- DELETE
- FROM ledger_transaction
- WHERE id = ANY(%s)
- """,
- params=[transaction_ids],
- )
+ self.validate_business_payout_in_ledger(ext_ref_id=ext_ref_id, thl_lm=thl_lm)
+ self.update_business_payout_event(pk=bpe.id, status=PayoutStatus.COMPLETE)
- # DELETE: event_payouts
- self.pg_config.execute_write(
- query="""
- DELETE
- FROM event_payout
- WHERE ext_ref_id = %s
- AND uuid = ANY(%s)
- """,
- params=[ext_ref_id, event_payout_uuids],
- )
-
- return None
+ return
- def get_business_payout_events_for_products(
+ def get_business_payout_events_for_business(
self,
- thl_ledger_manager: ThlLedgerManager,
- product_uuids: Collection[UUIDStr],
+ business_uuid: UUIDStr,
order_by: OrderBy | None = OrderBy.ASC,
) -> list[BusinessPayoutEvent]:
- res = self.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_ledger_manager,
- product_uuids=product_uuids,
- order_by=order_by,
+ order_by = order_by or OrderBy.ASC
+ bpes = self.filter_by(
+ business_uuids=[business_uuid],
)
-
- return self.from_bp_payout_events(bp_payout_events=res)
-
- @staticmethod
- def from_bp_payout_events(
- bp_payout_events: Collection[BrokerageProductPayoutEvent],
- ) -> list[BusinessPayoutEvent]:
- if len(bp_payout_events) == 0:
- return []
-
- grouped = defaultdict(list)
- for bp_pe in bp_payout_events:
- grouped[bp_pe.ext_ref_id].append(bp_pe)
-
- res = []
- for _, members in grouped.items():
- res.append(BusinessPayoutEvent.model_validate({"bp_payouts": members}))
-
- return res
+ bpes = sorted(bpes, key=lambda x: x.created, reverse=order_by == OrderBy.DESC)
+ return bpes
@staticmethod
def recoup_proportional(
@@ -1022,9 +835,9 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
shortfall: int = int(target_amount) - w_df["deduction"].sum()
w_df["remaining_balance"] = w_df["available_balance"] - w_df["deduction"]
- assert w_df[
- w_df["deduction"] > w_df["available_balance"]
- ].empty, "Trying to deduct more from an Product than what is available"
+ assert w_df[w_df["deduction"] > w_df["available_balance"]].empty, (
+ "Trying to deduct more from an Product than what is available"
+ )
return w_df
@@ -1073,7 +886,6 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
shortage = int(amount) - allocation.sum()
if shortage > 0:
-
assert shortage < len(remainders), (
"The shortage cent amount must be less than or equal to the "
"length of the remainders if we intend of taking a penny "
@@ -1091,122 +903,297 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
return allocation
+ def create_business_payout_event(
+ self,
+ bpe: BusinessPayoutEventCreate,
+ ) -> BusinessPayoutEvent:
+ assert isinstance(bpe, BusinessPayoutEventCreate)
+ assert bpe.bp_payouts, "Must provide at least one BP Payout"
+ assert {bp_pe.status for bp_pe in bpe.bp_payouts} == {PayoutStatus.PENDING}, (
+ "All BP Payouts must be PENDING"
+ )
+ INSERT_SUPPLIER_PAYOUT = """
+ INSERT INTO supplier_payout (
+ business_id, created, amount,
+ status, ext_ref_id, payout_type,
+ request_data, order_data
+ ) VALUES (
+ %(business_id)s, %(created)s, %(amount)s,
+ %(status)s, %(ext_ref_id)s, %(payout_type)s,
+ %(request_data)s, %(order_data)s
+ ) RETURNING id;
+ """
+ INSERT_BP_PAYOUT = """
+ INSERT INTO event_payout (
+ uuid, debit_account_uuid, created, cashout_method_uuid,
+ amount, status, ext_ref_id, payout_type, order_data,
+ request_data, supplier_payout_id
+ ) VALUES (
+ %(uuid)s, %(debit_account_uuid)s, %(created)s, %(cashout_method_uuid)s,
+ %(amount)s, %(status)s, %(ext_ref_id)s, %(payout_type)s, %(order_data)s,
+ %(request_data)s, %(supplier_payout_id)s
+ );
+ """
+
+ with self.pg_config.make_connection() as conn:
+ with conn.cursor() as c:
+ # ext_ref_id (transaction_id) has a unique constraint
+ try:
+ c.execute(INSERT_SUPPLIER_PAYOUT, bpe.model_dump_postgres())
+ except psycopg.errors.UniqueViolation as e:
+ if e.diag.constraint_name == "supplier_payout_ext_ref_id_key":
+ raise ValueError(
+ f"Cannot create a BusinessPayoutEvent with an existing "
+ f"transaction_id. {e.diag.message_detail}"
+ )
+ raise
+ supplier_payout_pk = c.fetchone()["id"]
+ for bp_pe in bpe.bp_payouts:
+ c.execute(
+ INSERT_BP_PAYOUT,
+ bp_pe.model_dump_postgres()
+ | {"supplier_payout_id": supplier_payout_pk},
+ )
+ conn.commit()
+ return BusinessPayoutEvent(id=supplier_payout_pk, **bpe.model_dump())
+
def create_from_ach_or_wire(
self,
business: Business,
amount: USDCent,
+ transaction_id: str,
pm: ProductManager,
thl_lm: ThlLedgerManager,
created: datetime | None = None,
- transaction_id: str | None = None,
) -> BusinessPayoutEvent | None:
"""This records a single banking transfer to a supplier. Takes a
specific Business that was paid out and how much. It then determines
how to distribute the amount to each Brokerage Product in the
Business.
-
- :param business
- :param amount
- :param pm
- :param thl_lm: this must have rw permissions to add transactions to
- the ledger
- :param created
- :param transaction_id
-
- :return:
"""
assert business.balance is not None, (
"Must provide a full version of a Business in order to calculate"
"the required Brokerage Product amounts."
)
- assert amount > 100_00, "Must issue Supplier Payouts at least $100 minimum."
- LOG.warning("Paying out ")
+ assert amount >= 100_00, "Must issue Supplier Payouts at least $100 minimum."
+ LOG.warning(f"Paying out {business.name} {amount.to_usd_str()}")
if created:
LOG.warning("Payouts in the past, require the parquet files to be rebuilt.")
- assert created < datetime.now(tz=timezone.utc)
-
+ assert created.tzinfo == UTC, "created must be UTC"
+ assert created < datetime.now(tz=UTC), "created must be in the past"
else:
- created = datetime.now(tz=timezone.utc)
+ created = datetime.now(tz=UTC)
# Gather the total amount available balance from each and put into
# a simple DF. We're using the available balance because we need it
- # to always be positive.. and we never want to get into a negative
+ # to always be positive. We never want to get into a negative
# situation again, so it's best to be extra conservative.
- res = {
+ balances = {
pb.product_id: pb.available_balance
for pb in business.balance.product_balances
}
- df = pd.DataFrame.from_dict(res, orient="index").reset_index()
+ df = pd.DataFrame.from_dict(balances, orient="index").reset_index()
df.columns = ["product_id", "available_balance"]
- res = BusinessPayoutEventManager.recoup_proportional(
+ df = BusinessPayoutEventManager.recoup_proportional(
df=df, target_amount=business.balance.recoup
)
# Can't pay any Products that don't have a remaining balance
- res = res[res["remaining_balance"] > 0]
+ df = df[df["remaining_balance"] > 0].copy()
- assert (
- res.deduction.sum() == business.balance.recoup
- ), "recoup_proportional failure"
+ assert df.deduction.sum() == business.balance.recoup, (
+ "recoup_proportional failure"
+ )
- res["issue_amount"] = BusinessPayoutEventManager.distribute_amount(
- df=res, amount=amount
+ df["issue_amount"] = BusinessPayoutEventManager.distribute_amount(
+ df=df, amount=amount
)
- assert res.issue_amount.sum() == amount, "issue_amount failure"
+ assert df.issue_amount.sum() == amount, "issue_amount failure"
# Can't pay any Products that don't have an issue amount
- res = res[res["issue_amount"] > 0]
+ df = df[df["issue_amount"] > 0].copy()
- recouped_amounts: list[dict[str, int]] = res[
- ["product_id", "remaining_balance", "issue_amount"]
- ].to_dict(orient="records")
+ amounts: dict[str, dict[str, int]] = df.set_index("product_id")[
+ ["remaining_balance", "issue_amount"]
+ ].to_dict(orient="index")
- # Get all of the products at once so we're not doing it for every interation
- products = pm.get_by_uuids(
- product_uuids=[i["product_id"] for i in recouped_amounts]
- )
+ products = pm.get_by_uuids(product_uuids=list(amounts.keys()))
+ product_lookup = {p.uuid: p for p in products}
- bp_payouts: list[BrokerageProductPayoutEvent] = []
- for idx, item in enumerate(recouped_amounts):
- product = next((p for p in products if p.uuid == item["product_id"]), None)
- assert product is not None
+ # Bulk version of this ---v
+ # bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=product)
+ qualified_names = [
+ f"{thl_lm.currency.value}:bp_wallet:{bp.id}" for bp in products
+ ]
+ bp_wallets = thl_lm.get_accounts(qualified_names)
+ wallet_lookup = {bpw.reference_uuid: bpw.uuid for bpw in bp_wallets}
- try:
- bp_pe: BrokerageProductPayoutEvent = self.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
+ bp_payouts: list[BrokerageProductPayoutEvent] = []
+ for product_id, item in amounts.items():
+ product = product_lookup[product_id]
+ bp_payouts.append(
+ BrokerageProductPayoutEvent(
+ created=created,
+ payout_type=PayoutType.ACH,
+ status=PayoutStatus.PENDING,
+ uuid=uuid4().hex,
amount=USDCent(item["issue_amount"]),
- created=created + timedelta(milliseconds=idx + 1),
ext_ref_id=transaction_id,
- skip_wallet_balance_check=True,
+ product_id=product.uuid,
+ cashout_method_uuid=self.bp_pe_manager.CASHOUT_METHOD_UUID,
+ debit_account_uuid=wallet_lookup[product_id],
)
+ )
- assert bp_pe.status == PayoutStatus.COMPLETE
- bp_payouts.append(bp_pe)
+ bpe_create = BusinessPayoutEventCreate(
+ business_id=business.uuid,
+ payout_type=PayoutType.ACH,
+ amount=amount,
+ created=created,
+ ext_ref_id=transaction_id,
+ # The ACH payment was sent! We haven't yet recorded it
+ # in the ledger, but it was sent by the bank. This
+ # is kind of ambiguous the meaning, we'll say it
+ # is not yet COMPLETE b/c the bp payouts
+ # haven't all been created yet.
+ status=PayoutStatus.APPROVED,
+ bp_payouts=bp_payouts,
+ )
+ # The supplier_payout db row and all event_payout (BP rows) are all
+ # created in the same DB transaction.
+ bpe = self.create_business_payout_event(bpe=bpe_create)
+ assert bpe.id is not None, "Something failed creating BusinessPayoutEvent"
+
+ # Now, go through each and create ledger txs. This is resumable
+ # from the BrokerageProductPayoutEvents
+ for bp_pe in bpe.bp_payouts:
+ product = product_lookup[bp_pe.product_id]
+ self.bp_pe_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_lm,
+ bp_pe=bp_pe,
+ product=product,
+ )
- except (Exception,) as e:
- # Cleanup bp_payouts
- print("Exception", e)
- return None
+ self.update_business_payout_event(pk=bpe.id, status=PayoutStatus.COMPLETE)
- if bp_pe.status == PayoutStatus.FAILED:
- sleep(1)
+ return bpe
- try:
- bp_pe = self.retry_create_bp_payout_event_tx(
- thl_ledger_manager=thl_lm,
- product=product,
- payout_event_uuid=bp_pe.uuid,
- )
- assert bp_pe.status == PayoutStatus.COMPLETE
- bp_payouts.append(bp_pe)
+ def update_business_payout_event(self, pk: int, status: PayoutStatus):
+ with self.connection() as conn:
+ with conn.cursor() as c:
+ c.execute(
+ """
+ UPDATE supplier_payout
+ SET status = %(status)s
+ WHERE id = %(pk)s""",
+ {"pk": pk, "status": status},
+ )
+ assert c.rowcount == 1, f"{id=} not found"
+ conn.commit()
+
+ def create_bp_payout_event(
+ self,
+ thl_ledger_manager: ThlLedgerManager,
+ product: Product,
+ amount: USDCent,
+ ext_ref_id: str,
+ created: datetime | None = None,
+ ):
+ """
+ This should NOT be called directly normally. It is just a shortcut
+ for tests. However, instead of just making a naked BP payout,
+ it created the business payout also, but with just one BP Payout
+ """
+ created = created or datetime.now(tz=UTC)
+ account = thl_ledger_manager.get_account(
+ f"{thl_ledger_manager.currency.value}:bp_wallet:{product.uuid}"
+ )
+ bpe = BusinessPayoutEventCreate(
+ business_id=product.business_uuid,
+ payout_type=PayoutType.ACH,
+ amount=amount,
+ created=created,
+ ext_ref_id=ext_ref_id,
+ status=PayoutStatus.APPROVED,
+ bp_payouts=[
+ BrokerageProductPayoutEvent(
+ created=created,
+ payout_type=PayoutType.ACH,
+ status=PayoutStatus.PENDING,
+ uuid=uuid4().hex,
+ amount=amount,
+ ext_ref_id=ext_ref_id,
+ product_id=product.uuid,
+ cashout_method_uuid=self.bp_pe_manager.CASHOUT_METHOD_UUID,
+ debit_account_uuid=account.uuid,
+ )
+ ],
+ )
+ bpe = self.create_business_payout_event(bpe=bpe)
+ self.bp_pe_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ bp_pe=bpe.bp_payouts[0],
+ product=product,
+ )
+ self.update_business_payout_event(pk=bpe.id, status=PayoutStatus.COMPLETE)
+ return bpe
+
+ def update_ext_reference_ids(
+ self,
+ new_value: str,
+ current_value: str,
+ ) -> None:
+ """
+ There are scenarios where an ACH/Wire payout event was saved with
+ a generic or anonymized reference identifier. We may want to be
+ able to go back and update all of those transaction IDs.
+
+ """
+ assert new_value and current_value
+
+ # Will raise if doesn't exist
+ self.get_by_ext_ref_id(ext_ref_id=current_value)
+
+ query1 = """
+ UPDATE supplier_payout
+ SET ext_ref_id = %(new_value)s
+ WHERE ext_ref_id = %(old_value)s
+ """
+ query2 = """
+ UPDATE event_payout
+ SET ext_ref_id = %(new_value)s
+ WHERE ext_ref_id = %(old_value)s
+ """
+ params = {"new_value": new_value, "old_value": current_value}
+ with self.pg_config.make_connection() as conn:
+ with conn.cursor() as c:
+ c.execute(query1, params)
+ assert c.rowcount == 1
+ c.execute(query2, params)
+ # As of 2025, no single Business has more than 10,000 Products,
+ # leave the limit in as an additional safeguard.
+ assert c.rowcount < 10000
+ conn.commit()
- except (Exception,) as e:
- # Cleanup bp_payouts
- return None
- return BusinessPayoutEvent.model_validate({"bp_payouts": bp_payouts})
+# import duckdb
+# conn = duckdb.connect()
+# conn.execute("""
+# select * from read_parquet('/mnt/thl-incite/raw/df-collections/ledger/*/*.parquet')
+# where event_payout is not null
+# and direction =1
+# and reference_uuid in ?
+# """, [b.product_uuids])
+# df = conn.fetch_df()
+# df['ext_description'].value_counts()
+#
+# tx_ids = [35554404, 37210650]
+# conn.execute("""
+# select * from read_parquet('/mnt/thl-incite/raw/df-collections/ledger/*/*.parquet')
+# where tx_id in ?
+# """, [tx_ids])
+# df = conn.fetch_df()
diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py
index 46280b8..66dd131 100644
--- a/generalresearch/managers/thl/product.py
+++ b/generalresearch/managers/thl/product.py
@@ -4,7 +4,7 @@ import json
import logging
import operator
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from threading import Lock
from typing import TYPE_CHECKING
@@ -18,26 +18,28 @@ from sentry_sdk import capture_exception
from generalresearch.decorators import LOG
from generalresearch.managers.base import (
- Permission,
PostgresManager,
)
from generalresearch.models.custom_types import UUIDStr, is_valid_uuid
-from generalresearch.pg_helper import PostgresConfig
-
-logger = logging.getLogger()
if TYPE_CHECKING:
+ from generalresearch.managers.base import (
+ Permission,
+ )
from generalresearch.models.thl.product import (
PayoutConfig,
Product,
ProfilingConfig,
SessionConfig,
SourcesConfig,
- SupplyConfigs,
+ SupplyConfig,
UserCreateConfig,
UserHealthConfig,
UserWalletConfig,
)
+ from generalresearch.pg_helper import PostgresConfig
+
+logger = logging.getLogger()
class ProductManager(PostgresManager):
@@ -94,9 +96,9 @@ class ProductManager(PostgresManager):
return self.fetch_uuids(
product_uuids=[product_uuid],
)[0]
- except (AssertionError,):
+ except AssertionError:
return None
- except (IndexError,):
+ except IndexError:
return None
def get_by_uuids_if_exists(
@@ -168,15 +170,14 @@ class ProductManager(PostgresManager):
if filter_uuids is None or len(filter_uuids) == 0:
return []
- with self.pg_config.make_connection() as sql_connection:
- with sql_connection.cursor() as c:
- res = []
- for chunk in chunked(filter_uuids, 500):
- res.extend(
- self.fetch_uuids_(
- c=c, filter_uuids=chunk, filter_column=filter_column
- )
+ with self.pg_config.make_connection() as sql_connection, sql_connection.cursor() as c:
+ res = []
+ for chunk in chunked(filter_uuids, 500):
+ res.extend(
+ self.fetch_uuids_(
+ c=c, filter_uuids=chunk, filter_column=filter_column
)
+ )
return res
def fetch_uuids_(
@@ -259,9 +260,10 @@ class ProductManager(PostgresManager):
for k, v in res1.items():
try:
r.append(Product.model_validate(v))
- except ValidationError as e:
+ except ValidationError:
logger.info(f"failed to parse product: {k}")
- raise e
+ raise
+
return r
def create(
@@ -273,7 +275,7 @@ class ProductManager(PostgresManager):
business_id: UUIDStr | None = None,
harmonizer_domain: str | None = None,
commission_pct: Decimal = Decimal("0.05"),
- sources_config: SourcesConfig | SupplyConfigs | None = None,
+ sources_config: SourcesConfig | SupplyConfig | None = None,
payout_config: PayoutConfig | None = None,
session_config: SessionConfig | None = None,
profiling_config: ProfilingConfig | None = None,
@@ -293,7 +295,7 @@ class ProductManager(PostgresManager):
UserWalletConfig,
)
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# TODO: Add product_id, and possibly name uniqueness validation to the
# pydantic model definition itself. The create manager doesn't need
@@ -361,10 +363,10 @@ class ProductManager(PostgresManager):
insert_data["payments_enabled"] = instance.payments_enabled
try:
- insert_data["id_int"] = list(self.pg_config.execute_sql_query(query="""
+ insert_data["id_int"] = next(iter(self.pg_config.execute_sql_query(query="""
SELECT COALESCE(MAX(id_int), 0) + 1 as id_int
FROM userprofile_brokerageproduct
- """))[0]["id_int"]
+ """)))["id_int"]
instance.id_int = insert_data["id_int"]
query = """
@@ -397,18 +399,18 @@ class ProductManager(PostgresManager):
#
# from pymysql import IntegrityError
# except IntegrityError as e:
- except (Exception,) as e:
+ except Exception as e:
try:
return self.get_by_uuid(product_uuid=instance.id)
- except (Exception,) as e2:
+ except AssertionError:
pass
finally:
self.cache_clear(instance.id)
# If we couldn't find the Product, then go ahead and raise.
capture_exception(e)
- raise e
+ raise
bpconfig = instance.model_dump(
include={"sources_config", "user_wallet"}, mode="json"
@@ -478,7 +480,7 @@ class ProductManager(PostgresManager):
data["grs_domain"] = data.pop("harmonizer_domain")
data = {k: v for k, v in data.items() if k in in_bp_keys}
data["id"] = product_uuid
- update_str = ", ".join(f"{k}=%({k})s" for k in data.keys())
+ update_str = ", ".join(f"{k}=%({k})s" for k in data)
self.pg_config.execute_write(
f"""
UPDATE userprofile_brokerageproduct
diff --git a/generalresearch/managers/thl/profiling/question.py b/generalresearch/managers/thl/profiling/question.py
index 10a9e32..078894a 100644
--- a/generalresearch/managers/thl/profiling/question.py
+++ b/generalresearch/managers/thl/profiling/question.py
@@ -85,7 +85,7 @@ class QuestionManager(PostgresManager):
def lookup_by_property(
self, property_code: str, country_iso: str, language_iso: str
) -> UpkQuestion:
- query = f"""
+ query = """
SELECT data, property_code, explanation_template, explanation_fragment_template
FROM marketplace_question
WHERE property_code = %(property_code)s
diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py
index cbe39e7..d240eea 100644
--- a/generalresearch/managers/thl/profiling/uqa.py
+++ b/generalresearch/managers/thl/profiling/uqa.py
@@ -2,14 +2,17 @@ from __future__ import annotations
import logging
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING
from generalresearch.managers.base import PostgresManagerWithRedis
from generalresearch.models.thl.profiling.user_question_answer import (
DUMMY_UQA,
UserQuestionAnswer,
)
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.user import User
logger = logging.getLogger()
@@ -128,7 +131,7 @@ class UQAManager(PostgresManagerWithRedis):
def get_from_db(self, user: User) -> list[UserQuestionAnswer]:
logger.info(f"get_uqa_from_db: {user.user_id}")
# Only store the latest row per question_id. We don't need it multiple times.
- since = datetime.now(tz=timezone.utc) - timedelta(days=30)
+ since = datetime.now(tz=UTC) - timedelta(days=30)
# We CAN use the RR, b/c either
# 1) the cache expired and the user hasn't sent an answer recently
diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py
index 1820103..53475b5 100644
--- a/generalresearch/managers/thl/profiling/user_upk.py
+++ b/generalresearch/managers/thl/profiling/user_upk.py
@@ -3,28 +3,30 @@ from __future__ import annotations
import json
from collections import defaultdict
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
-from typing import Any
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING, Any
from uuid import UUID
from psycopg import Cursor
from pydantic import PositiveInt
from generalresearch.managers.base import (
- Permission,
PostgresManagerWithRedis,
)
from generalresearch.managers.thl.profiling.schema import UpkSchemaManager
from generalresearch.models.thl.profiling.upk_property import (
Cardinality,
PropertyType,
- UpkProperty,
)
from generalresearch.models.thl.profiling.upk_question_answer import (
UpkQuestionAnswer,
)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.base import Permission
+ from generalresearch.models.thl.profiling.upk_property import UpkProperty
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class UserUpkManager(PostgresManagerWithRedis):
@@ -56,7 +58,7 @@ class UserUpkManager(PostgresManagerWithRedis):
return res
def get_user_upk_mysql(self, user_id: int) -> list[UpkQuestionAnswer]:
- since = datetime.now(tz=timezone.utc) - timedelta(days=89)
+ since = datetime.now(tz=UTC) - timedelta(days=89)
query = """
SELECT
@@ -158,7 +160,7 @@ class UserUpkManager(PostgresManagerWithRedis):
country_isos = {x["country_iso"] for x in upk_ans_dict}
assert len(country_isos) == 1
- country_iso = list(country_isos)[0]
+ country_iso = next(iter(country_isos))
for x in upk_ans_dict:
x["pred"] = x["pred"].replace("gr:", "")
x["obj"] = x["obj"].replace("gr:", "")
@@ -304,15 +306,15 @@ class UserUpkManager(PostgresManagerWithRedis):
def set_user_upk(self, upk_ans: list[UpkQuestionAnswer]) -> None:
user_id = {x.user_id for x in upk_ans}
assert len(user_id) == 1, "only run for 1 user at a time"
- user_id = list(user_id)[0]
+ user_id = next(iter(user_id))
curr_upk = self.get_user_upk(user_id=user_id)
curr_upk_simple = self.get_user_upk_simple(user_id=user_id)
new_upk_simple = defaultdict(set)
delete_items = set()
- upk_multi = list()
- delete_upk_multi = list()
+ upk_multi = []
+ delete_upk_multi = []
for x in upk_ans:
# For zero or more (multiple values) We want all values to equal these.
# Might involve deleting values if they exist and are not in upk_ans
diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py
index 9d3661c..3f235fd 100644
--- a/generalresearch/managers/thl/session.py
+++ b/generalresearch/managers/thl/session.py
@@ -1,29 +1,23 @@
from __future__ import annotations
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Any
+from typing import TYPE_CHECKING, Any
from uuid import UUID, uuid4
from faker import Faker
from psycopg import sql
from pydantic import NonNegativeInt, PositiveInt
-from generalresearch.managers import parse_order_by
from generalresearch.managers.base import (
Permission,
PostgresManager,
)
from generalresearch.managers.thl.product import ProductManager
-from generalresearch.models import DeviceType
+from generalresearch.managers.utils import parse_order_by
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.legacy.bucket import Bucket
-from generalresearch.models.thl.definitions import (
- SessionStatusCode2,
- Status,
- StatusCode1,
-)
from generalresearch.models.thl.session import (
Session,
Wall,
@@ -34,6 +28,15 @@ from generalresearch.models.thl.task_status import (
)
from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+
+ from generalresearch.models.definitions import DeviceType
+ from generalresearch.models.thl.definitions import (
+ SessionStatusCode2,
+ Status,
+ StatusCode1,
+ )
+
fake = Faker()
@@ -194,7 +197,7 @@ class SessionManager(PostgresManager):
# validation errors. There doesn't seem to be a clean way of doing this.
# model_copy with update doesn't trigger the validators, so we
# re-run model_validate after
- finished = finished if finished else datetime.now(tz=timezone.utc)
+ finished = finished if finished else datetime.now(tz=UTC)
session.update(
status=status,
status_code_1=status_code_1,
@@ -455,33 +458,31 @@ class SessionManager(PostgresManager):
params = {}
if started_before or started_after:
- started_after = started_after or datetime(2017, 1, 1, tzinfo=timezone.utc)
- started_before = started_before or datetime.now(tz=timezone.utc)
- assert started_after.tzinfo == timezone.utc, (
- "started_after must be tz-aware as UTC"
- )
- assert started_before.tzinfo == timezone.utc, (
- "started_before must be tz-aware as UTC"
- )
- assert started_after < started_before, (
- "started_after must be before started_before"
- )
+ started_after = started_after or datetime(2017, 1, 1, tzinfo=UTC)
+ started_before = started_before or datetime.now(tz=UTC)
+ assert started_after.tzinfo == UTC, "started_after must be tz-aware as UTC"
+ assert (
+ started_before.tzinfo == UTC
+ ), "started_before must be tz-aware as UTC"
+ assert (
+ started_after < started_before
+ ), "started_after must be before started_before"
filters.append("started BETWEEN %(started_after)s AND %(started_before)s")
params["started_after"] = started_after
params["started_before"] = started_before
if adjusted_before or adjusted_after:
- adjusted_after = adjusted_after or datetime(2017, 1, 1, tzinfo=timezone.utc)
- adjusted_before = adjusted_before or datetime.now(tz=timezone.utc)
- assert adjusted_after.tzinfo == timezone.utc, (
- "adjusted_after must be tz-aware as UTC"
- )
- assert adjusted_before.tzinfo == timezone.utc, (
- "adjusted_before must be tz-aware as UTC"
- )
- assert adjusted_after < adjusted_before, (
- "adjusted_after must be before adjusted_before"
- )
+ adjusted_after = adjusted_after or datetime(2017, 1, 1, tzinfo=UTC)
+ adjusted_before = adjusted_before or datetime.now(tz=UTC)
+ assert (
+ adjusted_after.tzinfo == UTC
+ ), "adjusted_after must be tz-aware as UTC"
+ assert (
+ adjusted_before.tzinfo == UTC
+ ), "adjusted_before must be tz-aware as UTC"
+ assert (
+ adjusted_after < adjusted_before
+ ), "adjusted_after must be before adjusted_before"
filters.append(
"adjusted_timestamp BETWEEN %(adjusted_after)s AND %(adjusted_before)s"
)
diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py
index c7671ce..92777e5 100644
--- a/generalresearch/managers/thl/survey.py
+++ b/generalresearch/managers/thl/survey.py
@@ -2,8 +2,8 @@ from __future__ import annotations
from collections import defaultdict
from collections.abc import Collection
-from datetime import datetime, timezone
-from typing import Any
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any
import pandas as pd
from more_itertools import chunked
@@ -13,13 +13,15 @@ from pydantic import NonNegativeInt
from generalresearch.managers.base import Permission, PostgresManager
from generalresearch.managers.thl.buyer import BuyerManager
from generalresearch.managers.thl.category import CategoryManager
-from generalresearch.models import Source
-from generalresearch.models.custom_types import SurveyKey
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.survey.model import (
Survey,
SurveyStat,
)
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.models.custom_types import SurveyKey
+ from generalresearch.pg_helper import PostgresConfig
class SurveyManager(PostgresManager):
@@ -134,7 +136,7 @@ class SurveyManager(PostgresManager):
if len(survey_keys) == 0:
return []
- params = dict()
+ params = {}
survey_source_ids = defaultdict(set)
for sk in survey_keys:
@@ -280,7 +282,6 @@ class SurveyManager(PostgresManager):
query=query,
params={"survey_pks": survey_pks},
)
- return None
def update_surveys_categories(self, surveys: list[Survey] | None = None) -> None:
for chunk in chunked(surveys, 500):
@@ -329,12 +330,11 @@ class SurveyManager(PostgresManager):
]
with self.pg_config.make_connection() as conn:
# noinspection PyArgumentList
- with conn.transaction():
- with conn.cursor() as c:
- c.execute(temp_table_sql)
- c.executemany(insert_values_sql, rows)
- c.execute(delete_sql)
- c.execute(upsert_sql)
+ with conn.transaction(), conn.cursor() as c:
+ c.execute(temp_table_sql)
+ c.executemany(insert_values_sql, rows)
+ c.execute(delete_sql)
+ c.execute(upsert_sql)
conn.commit()
def get_survey_categories(self):
@@ -356,59 +356,6 @@ class SurveyManager(PostgresManager):
class SurveyStatManager(PostgresManager):
- KEYS = [
- "survey_id",
- "quota_id",
- "country_iso",
- "version",
- "cpi",
- "complete_too_fast_cutoff",
- "prescreen_conv_alpha",
- "prescreen_conv_beta",
- "conv_alpha",
- "conv_beta",
- "dropoff_alpha",
- "dropoff_beta",
- "completion_time_mu",
- "completion_time_sigma",
- "mobile_eligible_alpha",
- "mobile_eligible_beta",
- "desktop_eligible_alpha",
- "desktop_eligible_beta",
- "tablet_eligible_alpha",
- "tablet_eligible_beta",
- "long_fail_rate",
- "user_report_coeff",
- "recon_likelihood",
- "score_x0",
- "score_x1",
- "score",
- "updated_at",
- "survey_is_live",
- "survey_survey_id",
- "survey_source",
- ]
-
- SURVEY_STATS_COL_MAP = {
- "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha",
- "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta",
- "CONVERSION.alpha": "conv_alpha",
- "CONVERSION.beta": "conv_beta",
- "COMPLETION_TIME.mu": "completion_time_mu",
- "COMPLETION_TIME.sigma": "completion_time_sigma",
- "LONG_FAIL.value": "long_fail_rate",
- "USER_REPORT_COEFF.value": "user_report_coeff",
- "RECON_LIKELIHOOD.value": "recon_likelihood",
- "DROPOFF_RATE.alpha": "dropoff_alpha",
- "DROPOFF_RATE.beta": "dropoff_beta",
- "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha",
- "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta",
- "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha",
- "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta",
- "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha",
- "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta",
- "cpi": "cpi",
- }
def __init__(
self,
@@ -421,6 +368,60 @@ class SurveyStatManager(PostgresManager):
)
# self.ensure_surveystat_key_type()
+ self.KEYS = [
+ "survey_id",
+ "quota_id",
+ "country_iso",
+ "version",
+ "cpi",
+ "complete_too_fast_cutoff",
+ "prescreen_conv_alpha",
+ "prescreen_conv_beta",
+ "conv_alpha",
+ "conv_beta",
+ "dropoff_alpha",
+ "dropoff_beta",
+ "completion_time_mu",
+ "completion_time_sigma",
+ "mobile_eligible_alpha",
+ "mobile_eligible_beta",
+ "desktop_eligible_alpha",
+ "desktop_eligible_beta",
+ "tablet_eligible_alpha",
+ "tablet_eligible_beta",
+ "long_fail_rate",
+ "user_report_coeff",
+ "recon_likelihood",
+ "score_x0",
+ "score_x1",
+ "score",
+ "updated_at",
+ "survey_is_live",
+ "survey_survey_id",
+ "survey_source",
+ ]
+
+ self.SURVEY_STATS_COL_MAP = {
+ "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha",
+ "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta",
+ "CONVERSION.alpha": "conv_alpha",
+ "CONVERSION.beta": "conv_beta",
+ "COMPLETION_TIME.mu": "completion_time_mu",
+ "COMPLETION_TIME.sigma": "completion_time_sigma",
+ "LONG_FAIL.value": "long_fail_rate",
+ "USER_REPORT_COEFF.value": "user_report_coeff",
+ "RECON_LIKELIHOOD.value": "recon_likelihood",
+ "DROPOFF_RATE.alpha": "dropoff_alpha",
+ "DROPOFF_RATE.beta": "dropoff_beta",
+ "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha",
+ "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta",
+ "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha",
+ "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta",
+ "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha",
+ "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta",
+ "cpi": "cpi",
+ }
+
#
# def ensure_surveystat_key_type(self):
# SQL = """
@@ -544,7 +545,7 @@ class SurveyStatManager(PostgresManager):
VALUES ({values_str})
ON CONFLICT ({unique_cols_str})
DO UPDATE SET {update_str};"""
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
params = [ss.model_dump_sql() | {"updated_at": now} for ss in survey_stats]
with self.pg_config.make_connection() as conn:
@@ -572,12 +573,10 @@ class SurveyStatManager(PostgresManager):
= (v.survey_id, v.quota_id, v.country_iso, v.version);
"""
params = [item for row in keys for item in row]
- with self.pg_config.make_connection() as conn:
- # self.register_surveystat_key(conn)
- with conn.cursor() as c:
- c.execute(query, params=params)
- res = c.fetchall()
- # print('\n'.join([x['QUERY PLAN'] for x in res]))
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, params=params)
+ res = c.fetchall()
+ # print('\n'.join([x['QUERY PLAN'] for x in res]))
return [SurveyStat.model_validate(x) for x in res]
def update_surveystats_for_source(
@@ -635,7 +634,7 @@ class SurveyStatManager(PostgresManager):
country_iso: str | None = None,
) -> tuple[str, dict[str, Any]]:
filters = []
- params = dict()
+ params = {}
if updated_after is not None:
params["updated_after"] = updated_after
filters.append("ss.updated_at >= %(updated_after)s")
@@ -760,12 +759,11 @@ class SurveyStatManager(PostgresManager):
print(query)
print(params)
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("SET work_mem = '256MB';")
- c.execute("SET statement_timeout = '10s';")
- c.execute(query, params=params)
- res = c.fetchall()
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute("SET work_mem = '256MB';")
+ c.execute("SET statement_timeout = '10s';")
+ c.execute(query, params=params)
+ res = c.fetchall()
return [SurveyStat.model_validate(x) for x in res]
diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py
index efaa930..f176e8d 100644
--- a/generalresearch/managers/thl/survey_penalty.py
+++ b/generalresearch/managers/thl/survey_penalty.py
@@ -4,21 +4,23 @@ import json
import threading
from collections import defaultdict
from datetime import timedelta
+from typing import TYPE_CHECKING
from cachetools import TTLCache, cachedmethod
from generalresearch.decorators import LOG
from generalresearch.managers.base import RedisManager
-from generalresearch.models.custom_types import (
- UUIDStr,
-)
-from generalresearch.models.thl.survey.penalty import (
- BPSurveyPenalty,
- Penalty,
- PenaltyListAdapter,
- TeamSurveyPenalty,
-)
-from generalresearch.redis_helper import RedisConfig
+from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.thl.survey.penalty import PenaltyListAdapter
+
+if TYPE_CHECKING:
+
+ from generalresearch.models.thl.survey.penalty import (
+ BPSurveyPenalty,
+ Penalty,
+ TeamSurveyPenalty,
+ )
+ from generalresearch.redis_helper import RedisConfig
class SurveyPenaltyManager(RedisManager):
@@ -61,7 +63,6 @@ class SurveyPenaltyManager(RedisManager):
return f"{self.redis_prefix}:{uuid_id}"
def set_penalties(self, penalties: list[Penalty]):
- """ """
if len(penalties) > 1000:
LOG.warning("SurveyPenaltyManager.set_penalties batch me!")
assert len(penalties) < 10_000, "something is surely wrong"
diff --git a/generalresearch/managers/thl/tango_api.py b/generalresearch/managers/thl/tango_api.py
index 657224e..dab560e 100644
--- a/generalresearch/managers/thl/tango_api.py
+++ b/generalresearch/managers/thl/tango_api.py
@@ -122,7 +122,7 @@ class TangoClient:
return self.get_order(reference_order_id)
except TangoError as e:
if "The order you requested cannot be found" not in e.args[0]:
- raise e
+ raise
return None
def create_order(self, order: TangoOrderRequest) -> dict[str, Any]:
diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py
index e4736d4..91d914f 100644
--- a/generalresearch/managers/thl/task_adjustment.py
+++ b/generalresearch/managers/thl/task_adjustment.py
@@ -1,19 +1,17 @@
from __future__ import annotations
import logging
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from functools import cached_property
+from typing import TYPE_CHECKING
-from generalresearch.managers import parse_order_by
from generalresearch.managers.base import (
PostgresManager,
)
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
from generalresearch.managers.thl.session import SessionManager
from generalresearch.managers.thl.wall import WallManager
+from generalresearch.managers.utils import parse_order_by
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.definitions import (
Status,
@@ -24,6 +22,15 @@ from generalresearch.models.thl.session import (
)
from generalresearch.models.thl.task_adjustment import TaskAdjustmentEvent
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+
+
+logging.basicConfig()
+logger = logging.getLogger(__name__)
+
class TaskAdjustmentManager(PostgresManager):
@@ -120,8 +127,8 @@ class TaskAdjustmentManager(PostgresManager):
CHANGES/DELTAS as just communicated by the marketplace, not
what the Wall's final adjusted_* will be.
"""
- alert_time = alert_time or datetime.now(tz=timezone.utc)
- assert alert_time.tzinfo == timezone.utc
+ alert_time = alert_time or datetime.now(tz=UTC)
+ assert alert_time.tzinfo == UTC
wall = self.wall_manager.get_from_uuid(wall_uuid)
session = self.session_manager.get_from_id(wall.session_id)
@@ -129,8 +136,9 @@ class TaskAdjustmentManager(PostgresManager):
user.prefetch_product(self.pg_config)
if adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL:
+ assert wall.cpi
amount_usd = wall.cpi * -1
- adjusted_cpi = 0
+ adjusted_cpi = Decimal(0)
elif adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE:
amount_usd = wall.cpi
adjusted_cpi = wall.cpi
@@ -148,10 +156,7 @@ class TaskAdjustmentManager(PostgresManager):
if (
wall.status == Status.COMPLETE
and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE
- ):
- new_adjusted_status = None
- new_adjusted_cpi = None
- elif (
+ ) or (
wall.status != Status.COMPLETE
and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL
):
@@ -172,7 +177,7 @@ class TaskAdjustmentManager(PostgresManager):
new_adjusted_cpi=new_adjusted_cpi,
)
except AssertionError as e:
- logging.warning(e)
+ logger.warning(e)
return
event = TaskAdjustmentEvent(
diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py
index 543de87..a3cad90 100644
--- a/generalresearch/managers/thl/user_compensate.py
+++ b/generalresearch/managers/thl/user_compensate.py
@@ -1,16 +1,19 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
from pydantic import NonNegativeInt
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.models.thl.user import User
def user_compensate(
@@ -28,7 +31,7 @@ def user_compensate(
pg_config = ledger_manager.pg_config
redis_client = ledger_manager.redis_client
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
assert type(amount_int) is int
user.prefetch_product(pg_config=pg_config)
assert (
diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py
index 0392edb..6414ef2 100644
--- a/generalresearch/managers/thl/user_manager/__init__.py
+++ b/generalresearch/managers/thl/user_manager/__init__.py
@@ -2,30 +2,18 @@ from __future__ import annotations
import csv
import logging
-import os
-import threading
-import time
from pathlib import Path
from threading import RLock
-from typing import Any
+from typing import TYPE_CHECKING, Any
from cachetools import TTLCache, cached
-from generalresearch.models.thl.product import Product
+if TYPE_CHECKING:
+ from generalresearch.models.thl.product import Product
logger = logging.getLogger()
-
-class UserDoesntExistError(Exception):
- pass
-
-
-class UserCreateNotAllowedError(Exception):
- pass
-
-
-def download_bp_trust():
- raise DeprecationWarning("No more S3")
+convert_int = lambda x: int(float(x))
@cached(TTLCache(maxsize=1, ttl=5 * 60), lock=RLock())
@@ -39,22 +27,11 @@ def get_bp_trust_df():
# 'product_name', 'bp_trust', 'team_trust', 'entrance_limit_expire_sec',
# 'entrance_limit_value']
- if not os.path.exists(fp):
- Path(fp).touch()
- threading.Thread(target=download_bp_trust).start()
- # raise exception so its not cached
- raise FileNotFoundError()
- if time.time() - os.path.getmtime(fp) > 3600:
- Path(fp).touch()
- threading.Thread(target=download_bp_trust).start()
bptrust = parse_bp_trust_df(fp)
return bptrust
-convert_int = lambda x: int(float(x))
-
-
def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]:
dtype = {
"bp_trust": float,
@@ -63,7 +40,7 @@ def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]:
"entrance_limit_value": convert_int,
"median_daily_completes_7d": convert_int,
}
- bptrust = dict()
+ bptrust = {}
with open(fp, newline="") as csvfile:
reader = csv.reader(csvfile)
diff --git a/generalresearch/managers/thl/user_manager/exceptions.py b/generalresearch/managers/thl/user_manager/exceptions.py
new file mode 100644
index 0000000..6d8c87c
--- /dev/null
+++ b/generalresearch/managers/thl/user_manager/exceptions.py
@@ -0,0 +1,6 @@
+class UserDoesntExistError(Exception):
+ pass
+
+
+class UserCreateNotAllowedError(Exception):
+ pass
diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
index 109b601..d2732ed 100644
--- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py
+++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
@@ -3,17 +3,20 @@ from __future__ import annotations
import logging
import operator
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from threading import Lock
+from typing import TYPE_CHECKING
from uuid import uuid4
-from cachetools import LRUCache, cachedmethod
import psycopg
+from cachetools import LRUCache, cachedmethod
from psycopg import sql
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.user import User
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
logging.basicConfig()
logger = logging.getLogger()
@@ -30,7 +33,7 @@ class MysqlUserManager:
def _set_last_seen(self, user: User) -> None:
# Don't call this directly. Use UserManager.set_last_seen()
assert not self.is_read_replica
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
self.pg_config.execute_write(
"""
UPDATE thl_user
@@ -40,19 +43,16 @@ class MysqlUserManager:
params=[now, user.user_id],
)
- def _change_product_user_id(
- self, *, user: User, new_product_user_id: str
- ) -> User:
+ def _change_product_user_id(self, *, user: User, new_product_user_id: str) -> User:
"""Change a user's supplier-provided ID in the primary database."""
assert not self.is_read_replica
assert user.user_id is not None
assert user.product_id is not None
assert user.product_user_id is not None
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(
- query="""
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ query="""
UPDATE thl_user
SET product_user_id = %(new_product_user_id)s
WHERE id = %(user_id)s
@@ -61,14 +61,14 @@ class MysqlUserManager:
RETURNING id AS user_id, product_id, product_user_id,
uuid, blocked, created, last_seen
""",
- params={
- "new_product_user_id": new_product_user_id,
- "user_id": user.user_id,
- "product_id": user.product_id,
- "old_product_user_id": user.product_user_id,
- },
- )
- row = c.fetchone()
+ params={
+ "new_product_user_id": new_product_user_id,
+ "user_id": user.user_id,
+ "product_id": user.product_id,
+ "old_product_user_id": user.product_user_id,
+ },
+ )
+ row = c.fetchone()
if row is None:
raise RuntimeError(
@@ -89,16 +89,16 @@ class MysqlUserManager:
logger.info(
f"get_user_from_mysql: {product_id}, {product_user_id}, {user_id}, {user_uuid}"
)
- assert (
- (product_id and product_user_id) or user_id or user_uuid
- ), "Must pass either (product_id, product_user_id), or user_id, or uuid"
+ assert (product_id and product_user_id) or user_id or user_uuid, (
+ "Must pass either (product_id, product_user_id), or user_id, or uuid"
+ )
if product_id or product_user_id:
- assert (
- product_id and product_user_id
- ), "Must pass both product_id and product_user_id"
- assert (
- sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1
- ), "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid"
+ assert product_id and product_user_id, (
+ "Must pass both product_id and product_user_id"
+ )
+ assert sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1, (
+ "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid"
+ )
# Using RR: Assume we check redis first for newly created users
if can_use_read_replica is False:
@@ -158,7 +158,7 @@ class MysqlUserManager:
if not self.product_id_exists(product_id=product_id):
raise ValueError(f"userprofile_brokerageproduct not found: {product_id}")
- now = created or datetime.now(tz=timezone.utc)
+ now = created or datetime.now(tz=UTC)
user_uuid = uuid4().hex
params = {
"user_uuid": user_uuid,
@@ -179,10 +179,9 @@ class MysqlUserManager:
""")
try:
- with self.pg_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query=query, params=params)
- user_id = c.fetchone()["id"]
+ with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query=query, params=params)
+ user_id = c.fetchone()["id"]
except psycopg.IntegrityError:
# Two machines/processes are trying to create this same (product_id, product_user_id)
# at the same time. There's a unique index, so mysql will not let two be created.
@@ -268,9 +267,9 @@ class MysqlUserManager:
assert product_id, "must pass product_id"
assert len(product_user_ids) > 0, "must pass 1 or more product_user_ids"
assert len(product_user_ids) <= 500, "limit 500 product_user_ids"
- assert isinstance(
- product_user_ids, (list, set)
- ), "must pass a collection of product_user_ids"
+ assert isinstance(product_user_ids, (list, set)), (
+ "must pass a collection of product_user_ids"
+ )
res = self.pg_config.execute_sql_query(
query="""
SELECT id AS user_id, product_id, product_user_id,
@@ -293,13 +292,13 @@ class MysqlUserManager:
user_ids: Collection[int] | None = None,
user_uuids: Collection[str] | None = None,
) -> list[User]:
- assert (user_ids or user_uuids) and not (
- user_ids and user_uuids
- ), "Must pass ONE of user_ids, user_uuids"
+ assert (user_ids or user_uuids) and not (user_ids and user_uuids), (
+ "Must pass ONE of user_ids, user_uuids"
+ )
if user_ids:
- assert isinstance(
- user_ids, (list, set)
- ), "must pass a collection of user_ids"
+ assert isinstance(user_ids, (list, set)), (
+ "must pass a collection of user_ids"
+ )
assert len(user_ids) <= 500, "limit 500 user_ids"
res = self.pg_config.execute_sql_query(
@@ -313,9 +312,9 @@ class MysqlUserManager:
params={"user_ids": user_ids},
)
else:
- assert isinstance(
- user_uuids, (list, set)
- ), "must pass a collection of user_uuids"
+ assert isinstance(user_uuids, (list, set)), (
+ "must pass a collection of user_uuids"
+ )
assert len(user_uuids) <= 500, "limit 500 user_uuids"
res = self.pg_config.execute_sql_query(
query="""
diff --git a/generalresearch/managers/thl/user_manager/rate_limit.py b/generalresearch/managers/thl/user_manager/rate_limit.py
index f938664..2aaa134 100644
--- a/generalresearch/managers/thl/user_manager/rate_limit.py
+++ b/generalresearch/managers/thl/user_manager/rate_limit.py
@@ -1,14 +1,19 @@
import logging
+from typing import TYPE_CHECKING
from limits import RateLimitItem, RateLimitItemPerHour, storage, strategies
from limits.limits import TIME_TYPES, safe_string
from pydantic import RedisDsn
from generalresearch.managers.thl.user_manager import (
- UserCreateNotAllowedError,
get_bp_user_create_limit_hourly,
)
-from generalresearch.models.thl.product import Product
+from generalresearch.managers.thl.user_manager.exceptions import (
+ UserCreateNotAllowedError,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.product import Product
logger = logging.getLogger()
diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py
index 7f24029..df3fb6b 100644
--- a/generalresearch/managers/thl/user_manager/user_manager.py
+++ b/generalresearch/managers/thl/user_manager/user_manager.py
@@ -5,13 +5,16 @@ import operator
from collections.abc import Collection
from datetime import datetime
from threading import Lock
+from typing import TYPE_CHECKING
from cachetools import TTLCache, cachedmethod
from pydantic import RedisDsn
from generalresearch.managers.base import Permission
from generalresearch.managers.thl.product import ProductManager
-from generalresearch.managers.thl.user_manager import UserDoesntExistError
+from generalresearch.managers.thl.user_manager.exceptions import (
+ UserDoesntExistError,
+)
from generalresearch.managers.thl.user_manager.mysql_user_manager import (
MysqlUserManager,
)
@@ -22,11 +25,15 @@ from generalresearch.managers.thl.user_manager.redis_user_manager import (
RedisUserManager,
)
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
-from generalresearch.pg_helper import PostgresConfig
from generalresearch.utils.copying_cache import deepcopy_return
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.userhealth import AuditLogManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.models.thl.userhealth import AuditLog
+ from generalresearch.pg_helper import PostgresConfig
+
logging.basicConfig()
logger = logging.getLogger()
auditlog = logging.getLogger("auditlog")
@@ -116,18 +123,17 @@ class UserManager:
def audit_log(
self,
+ alm: AuditLogManager,
user: User,
level: int,
event_type: str,
event_msg: str | None = None,
event_value: float | None = None,
- ) -> None:
- from generalresearch.managers.thl.userhealth import AuditLogManager
+ ) -> AuditLog:
from generalresearch.models.thl.userhealth import AuditLogLevel
- alm = AuditLogManager(pg_config=self.mysql_user_manager.pg_config)
assert user.user_id is not None
- alm.create(
+ return alm.create(
user_id=user.user_id,
level=AuditLogLevel(level),
event_type=event_type,
@@ -183,6 +189,7 @@ class UserManager:
# We can use the read-replica here b/c when we create a user we'll
# put it in the in-memory cache
mysql_user_manager = self.mysql_user_manager_rr or self.mysql_user_manager
+ assert mysql_user_manager
user = mysql_user_manager.get_user_from_mysql(
product_id=product_id,
product_user_id=product_user_id,
@@ -269,6 +276,8 @@ class UserManager:
"""
Given a bp_user_id and a product_id, get or create a User
"""
+ from generalresearch.models.thl.user import User
+
assert Permission.CREATE in self.sql_permissions
assert self.mysql_user_manager is not None
assert self.redis_user_manager is not None, (
@@ -279,6 +288,7 @@ class UserManager:
"Need user_manager_limiter to get_or_create_user"
)
# Attempt to create common_struct solely for validation purposes
+
if not User.is_valid_ubp(
product_id=product_id, product_user_id=product_user_id
):
@@ -316,6 +326,7 @@ class UserManager:
# if product.id not in {}:
# self.user_manager_limiter.raise_allow_user_create(product=product)
+ assert self.mysql_user_manager
user = self.mysql_user_manager.create_user(
product_user_id=product_user_id,
product_id=product.id,
@@ -328,6 +339,7 @@ class UserManager:
def product_id_exists(self, product_id: str) -> bool:
mysql_user_manager = self.mysql_user_manager_rr or self.mysql_user_manager
+ assert mysql_user_manager
return mysql_user_manager.product_id_exists(product_id)
def block_user(self, user: User) -> bool:
@@ -360,6 +372,7 @@ class UserManager:
Currently, this sets a key in the userprofile_userstat table.
TODO: this should be a property of the user?
"""
+ assert self.mysql_user_manager
return self.mysql_user_manager.is_whitelisted(user=user)
def fetch_by_bpuids(
@@ -370,6 +383,7 @@ class UserManager:
) -> list[User]:
assert product_id, "must pass product_id"
assert len(product_user_ids) > 0, "must pass 1 or more product_user_ids"
+ assert self.mysql_user_manager_rr
return self.mysql_user_manager_rr.fetch_by_bpuids(
product_id=product_id, product_user_ids=product_user_ids
)
@@ -383,6 +397,7 @@ class UserManager:
assert (user_ids or user_uuids) and not (user_ids and user_uuids), (
"Must pass ONE of user_ids, user_uuids"
)
+ assert self.mysql_user_manager_rr
return self.mysql_user_manager_rr.fetch(
user_ids=user_ids, user_uuids=user_uuids
)
diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py
index fe2163f..b986256 100644
--- a/generalresearch/managers/thl/userhealth.py
+++ b/generalresearch/managers/thl/userhealth.py
@@ -2,31 +2,34 @@ from __future__ import annotations
import ipaddress
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from itertools import zip_longest
-from typing import Any
+from typing import TYPE_CHECKING, Any
import faker
from pydantic import NonNegativeInt, PositiveInt
from generalresearch.decorators import LOG
from generalresearch.managers.base import (
- Permission,
PostgresManager,
PostgresManagerWithRedis,
)
from generalresearch.managers.thl.ipinfo import GeoIpInfoManager
-from generalresearch.models.custom_types import IPvAnyAddressStr
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
from generalresearch.models.thl.user_iphistory import (
IPRecord,
UserIPHistory,
UserIPRecord,
)
-from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+from generalresearch.models.thl.userhealth import AuditLog
+
+if TYPE_CHECKING:
+ from generalresearch.managers.base import Permission
+ from generalresearch.models.custom_types import IPvAnyAddressStr
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.models.thl.userhealth import AuditLogLevel
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
fake = faker.Faker()
@@ -210,7 +213,7 @@ class IPRecordManager(PostgresManagerWithRedis):
data = {
"user_id": user_id,
"ip": ipaddress.ip_address(ip).exploded,
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
}
fips_cols = [
@@ -221,7 +224,7 @@ class IPRecordManager(PostgresManagerWithRedis):
"forwarded_ip5",
"forwarded_ip6",
]
- for col, ip in zip_longest(
+ for col, fwd_ip in zip_longest(
fips_cols,
[
forwarded_ip1,
@@ -233,7 +236,7 @@ class IPRecordManager(PostgresManagerWithRedis):
],
fillvalue=None,
):
- data[col] = ipaddress.ip_address(ip).exploded if ip else ip
+ data[col] = ipaddress.ip_address(fwd_ip).exploded if fwd_ip else fwd_ip
self.pg_config.execute_write(
query="""
@@ -335,7 +338,7 @@ class AuditLogManager(PostgresManager):
al = AuditLog.model_validate(
{
"user_id": user_id,
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
"level": level,
"event_type": event_type,
"event_msg": event_msg,
@@ -374,10 +377,10 @@ class AuditLogManager(PostgresManager):
)
if len(res) == 0:
- raise Exception(f"No AuditLog with id of '{auditlog_id}'")
+ raise ValueError(f"No AuditLog with id of '{auditlog_id}'")
if len(res) > 1:
- raise Exception(f"Too many AuditLog found with id of '{auditlog_id}'")
+ raise ValueError(f"Too many AuditLog found with id of '{auditlog_id}'")
return AuditLog.from_mysql(res[0])
@@ -490,12 +493,10 @@ class AuditLogManager(PostgresManager):
created_after: datetime | None = None,
) -> tuple[str, dict[str, Any]]:
assert user_ids, "must pass at least 1 user_id"
- assert all(
- [isinstance(uid, int) for uid in user_ids]
- ), "must pass user_id as int"
+ assert all(isinstance(uid, int) for uid in user_ids), "must pass user_id as int"
if created_after is None:
- created_after = datetime.now(tz=timezone.utc) - timedelta(days=7)
+ created_after = datetime.now(tz=UTC) - timedelta(days=7)
filters = [
"user_id = ANY(%(user_ids)s)",
diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py
index c2eb821..bffd4c8 100644
--- a/generalresearch/managers/thl/wall.py
+++ b/generalresearch/managers/thl/wall.py
@@ -3,9 +3,10 @@ from __future__ import annotations
import logging
from collections import defaultdict
from collections.abc import Collection
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from functools import cached_property
+from typing import TYPE_CHECKING
from uuid import uuid4
from faker import Faker
@@ -13,20 +14,15 @@ from psycopg import sql
from psycopg.rows import dict_row
from pydantic import AwareDatetime, PositiveInt
-from generalresearch.managers import parse_order_by
from generalresearch.managers.base import (
- Permission,
PostgresManager,
PostgresManagerWithRedis,
)
-from generalresearch.models import Source
+from generalresearch.managers.utils import parse_order_by
from generalresearch.models.custom_types import SurveyKey, UUIDStr
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
- ReportValue,
- Status,
- StatusCode1,
WallAdjustedStatus,
- WallStatusCode2,
)
from generalresearch.models.thl.ledger import OrderBy
from generalresearch.models.thl.session import (
@@ -35,7 +31,18 @@ from generalresearch.models.thl.session import (
check_adjusted_status_wall_consistent,
)
from generalresearch.models.thl.survey.model import TaskActivity
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.base import (
+ Permission,
+ )
+ from generalresearch.models.thl.definitions import (
+ ReportValue,
+ Status,
+ StatusCode1,
+ WallStatusCode2,
+ )
+ from generalresearch.pg_helper import PostgresConfig
logger = logging.getLogger("WallManager")
fake = Faker()
@@ -360,12 +367,10 @@ class WallManager(PostgresManager):
params = {}
filters.append("user_id = %(user_id)s")
params["user_id"] = user_id
- default_started = datetime.now(tz=timezone.utc) - timedelta(days=90)
+ default_started = datetime.now(tz=UTC) - timedelta(days=90)
started_after = started_after or default_started
- started_before = started_before or datetime.now(tz=timezone.utc)
- assert (
- started_before.tzinfo == timezone.utc
- ), "started_before must be tz-aware as UTC"
+ started_before = started_before or datetime.now(tz=UTC)
+ assert started_before.tzinfo == UTC, "started_before must be tz-aware as UTC"
assert (
started_after < started_before
), "started_after must be before started_before"
@@ -412,7 +417,7 @@ class WallManager(PostgresManager):
started_before: datetime | None = None,
order_by: str | None = "-started",
) -> list[WallAttempt]:
- started_before = started_before or datetime.now(tz=timezone.utc)
+ started_before = started_before or datetime.now(tz=UTC)
res = []
page = 1
while True:
@@ -486,7 +491,7 @@ class WallManager(PostgresManager):
ORDER BY rs.source, rs.survey_id;
"""
- params = dict()
+ params = {}
filters = []
# Instead of doing a big IN with a big set of tuples, since we know
@@ -581,7 +586,7 @@ class WallCacheManager(PostgresManagerWithRedis):
# b as second element and a as third element"
attempts = sorted(attempts, key=lambda x: x.started)
json_res = [attempt.model_dump_json() for attempt in attempts]
- res = self.redis_client.lpush(redis_key, *json_res)
+ _ = self.redis_client.lpush(redis_key, *json_res)
self.redis_client.expire(redis_key, time=60 * 60 * 24)
# So this doesn't grow forever, keep only the most recent 5k
diff --git a/generalresearch/managers/thl/wallet/__init__.py b/generalresearch/managers/thl/wallet/__init__.py
index 70fcacf..7a8b4c9 100644
--- a/generalresearch/managers/thl/wallet/__init__.py
+++ b/generalresearch/managers/thl/wallet/__init__.py
@@ -1,27 +1,29 @@
from decimal import Decimal
-from typing import Any, Dict, Optional, Union
+from typing import TYPE_CHECKING, Any
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.payout import (
- PayoutEventManager,
- UserPayoutEventManager,
-)
-from generalresearch.managers.thl.tango_api import TangoClient
-from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
-)
-from generalresearch.managers.thl.userhealth import UserIpHistoryManager
from generalresearch.managers.thl.wallet.approve import (
approve_paypal_order,
)
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.wallet.cashout_method import (
- CashMailOrderData,
-)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.payout import (
+ PayoutEventManager,
+ UserPayoutEventManager,
+ )
+ from generalresearch.managers.thl.tango_api import TangoClient
+ from generalresearch.managers.thl.user_manager.user_manager import (
+ UserManager,
+ )
+ from generalresearch.managers.thl.userhealth import UserIpHistoryManager
+ from generalresearch.models.thl.payout import UserPayoutEvent
+ from generalresearch.models.thl.wallet.cashout_method import (
+ CashMailOrderData,
+ )
def manage_pending_cashout(
diff --git a/generalresearch/managers/thl/wallet/approve.py b/generalresearch/managers/thl/wallet/approve.py
index 012b406..7cae025 100644
--- a/generalresearch/managers/thl/wallet/approve.py
+++ b/generalresearch/managers/thl/wallet/approve.py
@@ -1,10 +1,16 @@
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.payout import PayoutEventManager
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.payout import PayoutEventManager
+ from generalresearch.models.thl.payout import UserPayoutEvent
+ from generalresearch.models.thl.user import User
def approve_paypal_order(
diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py
index 2f2dc52..038f67f 100644
--- a/generalresearch/managers/thl/wallet/tango.py
+++ b/generalresearch/managers/thl/wallet/tango.py
@@ -1,16 +1,21 @@
-from typing import Any, Dict
+from __future__ import annotations
+
+from typing import TYPE_CHECKING, Any
from generalresearch.config import (
is_debug,
)
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.payout import PayoutEventManager
-from generalresearch.managers.thl.tango_api import TangoClient, TangoOrderRequest
+from generalresearch.managers.thl.tango_api import TangoOrderRequest
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.payout import PayoutEventManager
+ from generalresearch.managers.thl.tango_api import TangoClient
+ from generalresearch.models.thl.payout import UserPayoutEvent
+ from generalresearch.models.thl.user import User
def complete_tango_order(
@@ -44,7 +49,7 @@ def complete_tango_order(
tango_client=tango_client,
)
- except Exception as e:
+ except AssertionError:
# todo: its possible the order went through, but something else was wrong
# we should try to retrieve the order by its ref_id and confirm it really
# failed...
@@ -65,13 +70,13 @@ def complete_tango_order(
def create_tango_order(
- request_data: Dict[str, Any], ref_id: str, tango_client: TangoClient
-) -> Dict[str, Any]:
+ request_data: dict[str, Any], ref_id: str, tango_client: TangoClient
+) -> dict[str, Any]:
"""
Create a tango gift card order.
Throws exception if anything is not right.
- # https://integration-www.tangocard.com/raas_api_console/v2/
- # https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python
+ - https://integration-www.tangocard.com/raas_api_console/v2/
+ - https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python
:param utid: Card identifier
:param amount: requested card value in USD
diff --git a/generalresearch/managers/utils.py b/generalresearch/managers/utils.py
new file mode 100644
index 0000000..bc745fd
--- /dev/null
+++ b/generalresearch/managers/utils.py
@@ -0,0 +1,16 @@
+def parse_order_by(order_by_str: str) -> str:
+ """
+ Converts django-rest-framework ordering str to mysql clause
+ :param order_by_str: e.g. 'created,-name'
+ :return: mysql clause e.g. ORDER BY created ASC, name DESC
+ """
+ fields = order_by_str.split(",")
+
+ order_clause = []
+ for field in fields:
+ if field.startswith("-"):
+ order_clause.append(f"{field[1:]} DESC")
+ else:
+ order_clause.append(f"{field} ASC")
+
+ return "ORDER BY " + ", ".join(order_clause)
diff --git a/generalresearch/mariadb.py b/generalresearch/mariadb.py
index 8bcd8ee..5d43f97 100644
--- a/generalresearch/mariadb.py
+++ b/generalresearch/mariadb.py
@@ -32,11 +32,3 @@ def example():
for m in zip(c.metadata["field"], c.metadata["ext_type_or_format"]):
# here we can just check if the field's ext_field_flag == 'UUID' (2)
print(m[0], ext_field_flags_rev[m[1]])
-
-
-def get_column_types():
- # How does django do this?
- res = """
- SELECT column_name, data_type
- FROM information_schema.columns
- WHERE table_name = 'morning_userpid' AND table_schema = DATABASE()"""
diff --git a/generalresearch/models/__init__.py b/generalresearch/models/__init__.py
index 560e5d9..e69de29 100644
--- a/generalresearch/models/__init__.py
+++ b/generalresearch/models/__init__.py
@@ -1,115 +0,0 @@
-from __future__ import annotations
-
-from enum import Enum
-
-from generalresearch.utils.enum import ReprEnumMeta
-
-
-class Source(str, Enum, metaclass=ReprEnumMeta):
- # The external marketplace, or the source of the survey / work.
- # Max length of the value is 2.
- GRS = "g"
- CINT = "c"
- DALIA = "a" # deprecated
- DYNATA = "d"
- ETX = "et"
- FULL_CIRCLE = "f"
- INNOVATE = "i"
- LUCID = "l"
- MORNING_CONSULT = "m"
- OPEN_LABS = "n"
- POLLFISH = "o"
- PRECISION = "e"
- PRODEGE_USER = "r" # deprecated
- PRODEGE = "pr" # using 'r' for vendor_wall
- PULLEY = "p" # deprecated
- REPDATA = "rd" # using 'q' for vendor_wall
- SAGO = "h"
- SPECTRUM = "s"
- TESTING = "t" # Used internally for testing
- TESTING2 = "u" # Used internally for testing
- WXET = "w"
-
-
-class DebitKey(int, Enum, metaclass=ReprEnumMeta):
- # The debit key for marketplaces
- CINT = 8
- DALIA = 9
- DYNATA = 6
- # ETX = None
- FULL_CIRCLE = 15
- INNOVATE = 7
- LUCID = 0
- MORNING_CONSULT = 12
- # OPEN_LABS = None
- POLLFISH = 13
- PRECISION = 14
- PRODEGE = 11
- SAGO = 10
- SPECTRUM = 5
- # WXET = None
-
-
-class DeviceType(int, Enum, metaclass=ReprEnumMeta):
- UNKNOWN = 0
- MOBILE = 1
- DESKTOP = 2
- TABLET = 3
-
-
-class LogicalOperator(str, Enum, metaclass=ReprEnumMeta):
- OR = "OR"
- AND = "AND"
- # There is currently no use case for NOT. See MarketplaceCondition.explain_not
- NOT = "NOT"
-
-
-class TaskStatus(str, Enum, metaclass=ReprEnumMeta):
- # A survey is live if it is open and, given all conditions are met, is
- # possible to send in traffic. All other statuses are just variants of
- # NOT Live (not accepting traffic)
- LIVE = "LIVE"
-
- # This is a generic NOT Live status. A marketplace may use other more
- # specific statuses but in practice they don't matter because all we care
- # about is if the task is LIVE.
- NOT_LIVE = "NOT_LIVE"
-
- # We need a status to mark if a survey we thought was live does not come
- # back from the API, we'll mark it as NOT_FOUND.
- NOT_FOUND = "NOT_FOUND"
-
-
-class TaskCalculationType(str, Enum):
- COMPLETES = "COMPLETES"
- STARTS = "STARTS"
-
- @classmethod
- def from_api(cls, v: str) -> TaskCalculationType:
- return {
- "complete": cls.COMPLETES,
- "completes": cls.COMPLETES,
- "survey start": cls.STARTS,
- "survey starts": cls.STARTS,
- "start": cls.STARTS,
- "starts": cls.STARTS,
- "prescreens": cls.STARTS,
- "prescreen": cls.STARTS,
- }[v.lower()]
-
- @classmethod
- def prodege_from_api(cls, v: int) -> TaskCalculationType:
- return {1: cls.COMPLETES, 2: cls.STARTS}[v]
-
- @classmethod
- def innovate_from_api(cls, v: int) -> TaskCalculationType:
- return {0: cls.COMPLETES, 1: cls.STARTS}[v]
-
-
-class URLQueryKey(str, Enum, metaclass=ReprEnumMeta):
- PRODUCT_ID = "39057c8b"
- PRODUCT_USER_ID = "c184efc0"
- SESSION_ID = "0bb50182"
-
-
-MAX_INT32 = 2**31
diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py
index ebe839a..344c34a 100644
--- a/generalresearch/models/admin/__init__.py
+++ b/generalresearch/models/admin/__init__.py
@@ -1,14 +1,14 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pandas as pd
from dateutil import relativedelta
def get_date_list(start_datetime: datetime, end_datetime: datetime | None = None):
- start_datetime = start_datetime.replace(tzinfo=timezone.utc)
- end_datetime = end_datetime if end_datetime else datetime.now(tz=timezone.utc)
+ start_datetime = start_datetime.replace(tzinfo=UTC)
+ end_datetime = end_datetime if end_datetime else datetime.now(tz=UTC)
return (
pd.date_range(start_datetime, end_datetime, freq="1D")
.strftime("%Y-%m-%d")
@@ -22,7 +22,7 @@ def year_start(periods_ago: int = 6) -> datetime:
years. Goal is to provide a simple way to
know when to do filters from
"""
- n: datetime = datetime.now(tz=timezone.utc)
+ n: datetime = datetime.now(tz=UTC)
d: datetime = n - relativedelta.relativedelta(years=periods_ago)
return d.replace(month=1, day=1, hour=0, minute=0, second=0, microsecond=0)
@@ -33,7 +33,7 @@ def month_start(periods_ago: int = 6) -> datetime:
months. Goal is to provide a simple way to
know when to do filters from
"""
- n: datetime = datetime.now(tz=timezone.utc)
+ n: datetime = datetime.now(tz=UTC)
d: datetime = n - relativedelta.relativedelta(months=periods_ago)
return d.replace(day=1, hour=0, minute=0, second=0, microsecond=0)
@@ -44,7 +44,7 @@ def day_start(periods_ago: int = 6) -> datetime:
days. Goal is to provide a simple way to
know when to do filters from
"""
- n: datetime = datetime.now(tz=timezone.utc)
+ n: datetime = datetime.now(tz=UTC)
d: datetime = n - relativedelta.relativedelta(days=periods_ago)
return d.replace(hour=0, minute=0, second=0, microsecond=0)
@@ -55,6 +55,6 @@ def hour_start(periods_ago: int = 6) -> datetime:
hours. Goal is to provide a simple way to
know when to do filters from
"""
- n: datetime = datetime.now(tz=timezone.utc)
+ n: datetime = datetime.now(tz=UTC)
d: datetime = n - relativedelta.relativedelta(hours=periods_ago)
return d.replace(minute=0, second=0, microsecond=0)
diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py
index 67bd263..6112786 100644
--- a/generalresearch/models/admin/request.py
+++ b/generalresearch/models/admin/request.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from enum import Enum
from typing import Literal
@@ -25,9 +25,9 @@ class ReportRequest(BaseModel):
index1: str = Field(default="product_id")
start: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc) - timedelta(days=14)
+ default_factory=lambda: datetime.now(tz=UTC) - timedelta(days=14)
)
- end: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=timezone.utc))
+ end: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
interval: Literal["5min", "15min", "1h", "6h", "12h", "1d"] = "1h"
include_open_bucket: bool = Field(default=True)
@@ -35,7 +35,7 @@ class ReportRequest(BaseModel):
@computed_field(
title="Start floor",
description="The datetime that this report starts from",
- examples=[datetime(year=2025, month=5, day=1, tzinfo=timezone.utc)],
+ examples=[datetime(year=2025, month=5, day=1, tzinfo=UTC)],
return_type=datetime,
)
@property
@@ -60,7 +60,7 @@ class ReportRequest(BaseModel):
@model_validator(mode="after")
def check_start_end_tz(self):
- assert self.start.tzinfo == self.end.tzinfo == timezone.utc
+ assert self.start.tzinfo == self.end.tzinfo == UTC
return self
@model_validator(mode="after")
@@ -119,7 +119,7 @@ class ReportRequest(BaseModel):
@property
def end_naive(self) -> datetime:
- return datetime.now(tz=None)
+ return datetime.now(tz=None) # noqa
@property
def ts_start(self) -> pd.Timestamp:
@@ -150,7 +150,7 @@ class ReportRequest(BaseModel):
start=self.ts_start_floor,
end=self.ts_end,
freq=self.interval,
- tz=timezone.utc,
+ tz=UTC,
)
def bucket_ranges(self) -> list[tuple[pd.Timestamp, pd.Timestamp]]:
diff --git a/generalresearch/models/cint/__init__.py b/generalresearch/models/cint/__init__.py
index 2c1be7e..d2713ab 100644
--- a/generalresearch/models/cint/__init__.py
+++ b/generalresearch/models/cint/__init__.py
@@ -1,5 +1,6 @@
+from typing import Annotated
+
from pydantic import Field
-from typing_extensions import Annotated
CintQuestionIdType = Annotated[
str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py
index 1ac9eea..89c7871 100644
--- a/generalresearch/models/cint/question.py
+++ b/generalresearch/models/cint/question.py
@@ -1,29 +1,27 @@
from __future__ import annotations
import json
-from datetime import datetime, timezone
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Literal
+from datetime import UTC, datetime
+from enum import StrEnum
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import BaseModel, Field, field_validator, model_validator
-from typing_extensions import Self
-from generalresearch.models import Source, string_utils
from generalresearch.models.cint import CintQuestionIdType
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import Source
+from generalresearch.models.string_utils import remove_nbsp
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
MarketplaceUserQuestionAnswer,
)
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
-class CintQuestionType(str, Enum):
+class CintQuestionType(StrEnum):
SINGLE_SELECT = "s"
MULTI_SELECT = "m"
# Dummy means they're calculated
@@ -45,7 +43,7 @@ class CintQuestionType(str, Enum):
# This seems to be invalid as there are no options???
"Grid": None,
}
- return API_TYPE_MAP[a] if a in API_TYPE_MAP else None
+ return API_TYPE_MAP.get(a)
class CintUserQuestionAnswer(MarketplaceUserQuestionAnswer):
@@ -107,7 +105,7 @@ class CintQuestion(MarketplaceQuestion):
@field_validator("question_name", "question_text", mode="after")
def remove_nbsp(cls, s: str | None) -> str | None:
- return string_utils.remove_nbsp(s)
+ return remove_nbsp(s)
@model_validator(mode="after")
def check_type_options_agreement(self) -> Self:
@@ -151,7 +149,7 @@ class CintQuestion(MarketplaceQuestion):
options = None
created_at = datetime.strptime(
d["create_date"], "%Y-%m-%dT%H:%M:%S%z"
- ).astimezone(timezone.utc)
+ ).astimezone(UTC)
if d.get("question_options"):
options = [
@@ -189,7 +187,7 @@ class CintQuestion(MarketplaceQuestion):
]
if d.get("created_at"):
- d["created_at"] = d["created_at"].replace(tzinfo=timezone.utc)
+ d["created_at"] = d["created_at"].replace(tzinfo=UTC)
return cls(
question_id=d["question_id"],
diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py
index 21bc21a..b2a8935 100644
--- a/generalresearch/models/cint/survey.py
+++ b/generalresearch/models/cint/survey.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import json
import logging
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
-from typing import Annotated, Any, Literal
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
from pydantic import (
@@ -12,19 +12,19 @@ from pydantic import (
ConfigDict,
Field,
NonNegativeInt,
+ ValidationError,
computed_field,
model_validator,
)
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import Source, TaskCalculationType
from generalresearch.models.cint import CintQuestionIdType
from generalresearch.models.custom_types import (
AlphaNumStr,
AwareDatetimeISO,
CoercedStr,
)
+from generalresearch.models.definitions import Source, TaskCalculationType
from generalresearch.models.thl.demographics import Gender
from generalresearch.models.thl.survey import MarketplaceTask
from generalresearch.models.thl.survey.condition import (
@@ -68,7 +68,7 @@ class CintQuota(BaseModel):
condition_hashes: list[str] | None = Field(min_length=1, default=None)
def __hash__(self):
- return hash(tuple((tuple(self.condition_hashes), self.quota_id)))
+ return hash((tuple(self.condition_hashes), self.quota_id))
@model_validator(mode="after")
def validate_condition_len(self) -> Self:
@@ -318,7 +318,7 @@ class CintSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> Self | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@@ -371,8 +371,8 @@ class CintSurvey(MarketplaceTask):
d["mobile_conversion"] = None
d["revenue_per_click"] = None
- d["conditions"] = dict()
- d.setdefault("survey_qualifications", list())
+ d["conditions"] = {}
+ d.setdefault("survey_qualifications", [])
qualifications = [CintCondition.from_api(q) for q in d["survey_qualifications"]]
for q in qualifications:
d["conditions"][q.criterion_hash] = q
@@ -390,7 +390,7 @@ class CintSurvey(MarketplaceTask):
d["conditions"][q.criterion_hash] = q
d["quotas"] = quotas
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d["created_at"] = now
d["last_updated"] = now
@@ -418,8 +418,8 @@ class CintSurvey(MarketplaceTask):
@classmethod
def from_mysql(cls, d: dict[str, Any]) -> Self:
- d["created_at"] = d["created_at"].replace(tzinfo=timezone.utc)
- d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc)
+ d["created_at"] = d["created_at"].replace(tzinfo=UTC)
+ d["last_updated"] = d["last_updated"].replace(tzinfo=UTC)
d["qualifications"] = json.loads(d["qualifications"])
d["used_question_ids"] = json.loads(d["used_question_ids"])
d["quotas"] = json.loads(d["quotas"])
@@ -466,7 +466,7 @@ class CintSurvey(MarketplaceTask):
) -> tuple[bool | None, set[str]]:
# Many surveys have 0 quotas. Quotas are exclusionary.
# They can NOT match a quota where currently_open=0
- total_quota = [q for q in self.quotas if q.quota_type == "total"][0]
+ total_quota = next(q for q in self.quotas if q.quota_type == "total")
if not total_quota.is_open:
return False, set()
quotas = [q for q in self.quotas if q.quota_type != "total"]
@@ -475,7 +475,7 @@ class CintSurvey(MarketplaceTask):
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py
index 5d39090..4ae8de4 100644
--- a/generalresearch/models/cint/task_collection.py
+++ b/generalresearch/models/cint/task_collection.py
@@ -1,7 +1,5 @@
from __future__ import annotations
-from typing import List
-
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
@@ -35,8 +33,8 @@ CintTaskCollectionSchema = DataFrameSchema(
"bid_ir": Column(float, Check.between(0, 1), nullable=True),
"created_at": Column(dtype=pd.DatetimeTZDtype(tz="UTC")),
"last_updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")),
- "used_question_ids": Column(List[str]),
- "all_hashes": Column(List[str]), # set >> list for column support
+ "used_question_ids": Column(list[str]),
+ "all_hashes": Column(list[str]), # set >> list for column support
},
checks=[],
index=Index(
diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py
index 84bf8e3..680a99c 100644
--- a/generalresearch/models/custom_types.py
+++ b/generalresearch/models/custom_types.py
@@ -2,9 +2,8 @@ from __future__ import annotations
import json
import re
-import sys as _sys
-from datetime import datetime, timedelta, timezone
-from typing import Any, Literal
+from datetime import UTC, datetime, timedelta
+from typing import Annotated, Any, Literal
from uuid import UUID
from pydantic import (
@@ -20,9 +19,8 @@ from pydantic.functional_serializers import PlainSerializer
from pydantic.functional_validators import AfterValidator, BeforeValidator
from pydantic.networks import IPvAnyNetwork, UrlConstraints
from pydantic_core import MultiHostHost, Url
-from typing_extensions import Annotated
-from generalresearch.models import DeviceType, Source
+from generalresearch.models.definitions import DeviceType, Source
HOSTNAME_REGEX = re.compile(
r"^[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?(\.[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?)*$"
@@ -57,19 +55,17 @@ def convert_str_dt(v: Any) -> AwareDatetime | None:
# to parse a str that was dumped using the iso8601 format with Z suffix.
if v is not None and type(v) is str:
assert v.endswith("Z") and "T" in v, "invalid format"
- return datetime.strptime(v, "%Y-%m-%dT%H:%M:%S.%fZ").replace(
- tzinfo=timezone.utc
- )
+ return datetime.strptime(v, "%Y-%m-%dT%H:%M:%S.%fZ").replace(tzinfo=UTC)
return v
def assert_utc(v: AwareDatetime) -> AwareDatetime:
if isinstance(v, datetime):
# We need utcoffset b/c FastAPI parses datetimes using FixedTimezone
- assert v.tzinfo == timezone.utc or v.tzinfo.utcoffset(v) == timedelta(
+ assert v.tzinfo == UTC or v.tzinfo.utcoffset(v) == timedelta(
0
), "Timezone is not UTC"
- v = v.astimezone(timezone.utc)
+ v = v.astimezone(UTC)
return v
@@ -102,7 +98,7 @@ LanguageISOLike = Annotated[
def check_valid_uuid(v: str) -> str:
try:
assert UUID(v).hex == v
- except Exception:
+ except (ValueError, AssertionError):
raise ValueError("Invalid UUID")
return v
@@ -110,7 +106,7 @@ def check_valid_uuid(v: str) -> str:
def is_valid_uuid(v: str) -> bool:
try:
assert UUID(v).hex == v
- except Exception:
+ except (ValueError, AssertionError):
return False
return True
@@ -169,7 +165,7 @@ CoercedStr = Annotated[str, BeforeValidator(coerce_int_to_str)]
# Serializers that can transform a collection of str into a comma separated
# str bidirectionally
-to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(list(x))), return_type=str)
+to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(x)), return_type=str)
enum_to_comma_sep_str = PlainSerializer(
lambda x: ",".join(sorted([str(y.value) for y in x])), return_type=str
)
@@ -309,4 +305,4 @@ PropertyCode = Annotated[
def now_utc_factory():
- return datetime.now(tz=timezone.utc)
+ return datetime.now(tz=UTC)
diff --git a/generalresearch/models/definitions.py b/generalresearch/models/definitions.py
new file mode 100644
index 0000000..c0348d7
--- /dev/null
+++ b/generalresearch/models/definitions.py
@@ -0,0 +1,114 @@
+from __future__ import annotations
+
+from enum import IntEnum, StrEnum
+
+from generalresearch.utils.enum import ReprEnumMeta
+
+
+class Source(StrEnum, metaclass=ReprEnumMeta):
+ # The external marketplace, or the source of the survey / work.
+ # Max length of the value is 2.
+ GRS = "g"
+ CINT = "c"
+ DALIA = "a" # deprecated
+ DYNATA = "d"
+ ETX = "et"
+ FULL_CIRCLE = "f"
+ INNOVATE = "i"
+ LUCID = "l"
+ MORNING_CONSULT = "m"
+ OPEN_LABS = "n"
+ POLLFISH = "o"
+ PRECISION = "e"
+ PRODEGE_USER = "r" # deprecated
+ PRODEGE = "pr" # using 'r' for vendor_wall
+ PULLEY = "p" # deprecated
+ REPDATA = "rd" # using 'q' for vendor_wall
+ SAGO = "h"
+ SPECTRUM = "s"
+ TESTING = "t" # Used internally for testing
+ TESTING2 = "u" # Used internally for testing
+ WXET = "w"
+
+
+class DebitKey(IntEnum, metaclass=ReprEnumMeta):
+ # The debit key for marketplaces
+ CINT = 8
+ DALIA = 9
+ DYNATA = 6
+ # ETX = None
+ FULL_CIRCLE = 15
+ INNOVATE = 7
+ LUCID = 0
+ MORNING_CONSULT = 12
+ # OPEN_LABS = None
+ POLLFISH = 13
+ PRECISION = 14
+ PRODEGE = 11
+ SAGO = 10
+ SPECTRUM = 5
+ # WXET = None
+
+
+class DeviceType(IntEnum, metaclass=ReprEnumMeta):
+ UNKNOWN = 0
+ MOBILE = 1
+ DESKTOP = 2
+ TABLET = 3
+
+
+class LogicalOperator(StrEnum, metaclass=ReprEnumMeta):
+ OR = "OR"
+ AND = "AND"
+ # There is currently no use case for NOT. See MarketplaceCondition.explain_not
+ NOT = "NOT"
+
+
+class TaskStatus(StrEnum, metaclass=ReprEnumMeta):
+ # A survey is live if it is open and, given all conditions are met, is
+ # possible to send in traffic. All other statuses are just variants of
+ # NOT Live (not accepting traffic)
+ LIVE = "LIVE"
+
+ # This is a generic NOT Live status. A marketplace may use other more
+ # specific statuses but in practice they don't matter because all we care
+ # about is if the task is LIVE.
+ NOT_LIVE = "NOT_LIVE"
+
+ # We need a status to mark if a survey we thought was live does not come
+ # back from the API, we'll mark it as NOT_FOUND.
+ NOT_FOUND = "NOT_FOUND"
+
+
+class TaskCalculationType(StrEnum):
+ COMPLETES = "COMPLETES"
+ STARTS = "STARTS"
+
+ @classmethod
+ def from_api(cls, v: str) -> TaskCalculationType:
+ return {
+ "complete": cls.COMPLETES,
+ "completes": cls.COMPLETES,
+ "survey start": cls.STARTS,
+ "survey starts": cls.STARTS,
+ "start": cls.STARTS,
+ "prescreens": cls.STARTS,
+ "prescreen": cls.STARTS,
+ }[v.lower()]
+
+ @classmethod
+ def prodege_from_api(cls, v: int) -> TaskCalculationType:
+ return {1: cls.COMPLETES, 2: cls.STARTS}[v]
+
+ @classmethod
+ def innovate_from_api(cls, v: int) -> TaskCalculationType:
+ return {0: cls.COMPLETES, 1: cls.STARTS}[v]
+
+
+class URLQueryKey(StrEnum, metaclass=ReprEnumMeta):
+ PRODUCT_ID = "39057c8b"
+ PRODUCT_USER_ID = "c184efc0"
+ SESSION_ID = "0bb50182"
+
+
+MAX_INT32 = 2**31
diff --git a/generalresearch/models/device.py b/generalresearch/models/device.py
index cc15eee..432c897 100644
--- a/generalresearch/models/device.py
+++ b/generalresearch/models/device.py
@@ -1,6 +1,6 @@
from user_agents import parse as parse_ua
-from generalresearch.models import DeviceType
+from generalresearch.models.definitions import DeviceType
def parse_device_from_useragent(user_agent: str) -> DeviceType:
diff --git a/generalresearch/models/dynata/__init__.py b/generalresearch/models/dynata/__init__.py
index c6d3a67..0f3dfe7 100644
--- a/generalresearch/models/dynata/__init__.py
+++ b/generalresearch/models/dynata/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class DynataStatus(str, Enum):
+class DynataStatus(StrEnum):
OPEN = "OPEN"
PAUSED = "PAUSED"
CLOSED = "CLOSED"
diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py
index 7288588..cbd5d86 100644
--- a/generalresearch/models/dynata/question.py
+++ b/generalresearch/models/dynata/question.py
@@ -5,14 +5,14 @@ import json
import logging
import re
from datetime import timedelta
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from typing import Any, Literal
from pydantic import BaseModel, Field, PositiveInt, field_validator, model_validator
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion
logging.basicConfig()
@@ -51,7 +51,7 @@ class DynataQuestionOption(BaseModel):
return clean_text(s)
-class DynataQuestionType(str, Enum):
+class DynataQuestionType(StrEnum):
"""
From the API: {'geo', 'multi_select', 'multi_select_searchable', 'none',
'single_select', 'single_select_grid', 'single_select_searchable', 'zip'}
diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py
index e88491d..6ff397f 100644
--- a/generalresearch/models/dynata/survey.py
+++ b/generalresearch/models/dynata/survey.py
@@ -2,10 +2,10 @@ from __future__ import annotations
import json
import logging
-from datetime import timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal
+from typing import Any, Literal, Self
from more_itertools import flatten
from pydantic import (
@@ -17,10 +17,8 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import Source, TaskCalculationType
from generalresearch.models.custom_types import (
AlphaNumStr,
AlphaNumStrSet,
@@ -28,6 +26,7 @@ from generalresearch.models.custom_types import (
CoercedStr,
DeviceTypes,
)
+from generalresearch.models.definitions import Source, TaskCalculationType
from generalresearch.models.dynata import DynataStatus
from generalresearch.models.thl.demographics import (
Gender,
@@ -133,9 +132,7 @@ class DynataCondition(MarketplaceCondition):
if cell["kind"] == "RANGE":
d["values"] = [
- "{0}-{1}".format(
- cell["range"]["from"] or "inf", cell["range"]["to"] or "inf"
- )
+ f"{cell['range']['from'] or 'inf'}-{cell['range']['to'] or 'inf'}"
]
d["value_type"] = ConditionValueType.RANGE
return cls.model_validate(d)
@@ -171,7 +168,7 @@ class DynataQuota(BaseModel):
status: DynataStatus = Field()
def __hash__(self):
- return hash(tuple((tuple(self.condition_hashes), self.count, self.status)))
+ return hash((tuple(self.condition_hashes), self.count, self.status))
@property
def is_open(self) -> bool:
@@ -247,7 +244,7 @@ class DynataQuotaGroup(RootModel):
) -> tuple[bool | None, set[str]]:
# Qualify for ANY quota object within a quota group
obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root}
- evals = set(v[0] for v in obj_evals.values())
+ evals = {v[0] for v in obj_evals.values()}
# If we match 1 obj, then the others don't matter
if any(evals):
return True, set()
@@ -322,7 +319,7 @@ class DynataFilterGroup(RootModel):
) -> tuple[bool | None, set[str]]:
# Passes back "passes" (T/F/none) and a list of unknown criterion hashes
obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root}
- evals = set(v[0] for v in obj_evals.values())
+ evals = {v[0] for v in obj_evals.values()}
# If we match 1 obj, then the others don't matter
if any(evals):
return True, set()
@@ -552,8 +549,8 @@ class DynataSurvey(MarketplaceTask):
@classmethod
def from_db(cls, d: dict[str, Any]) -> Self:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc)
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["last_updated"] = d["last_updated"].replace(tzinfo=UTC)
d["filters"] = json.loads(d["filters"])
d["quotas"] = json.loads(d["quotas"])
d["used_question_ids"] = json.loads(d["used_question_ids"])
@@ -581,7 +578,7 @@ class DynataSurvey(MarketplaceTask):
group_eval = {
group: group.passes_soft(criteria_evaluation) for group in self.filters
}
- evals = set(g[0] for g in group_eval.values())
+ evals = {g[0] for g in group_eval.values()}
if False in evals:
return False, set()
elif None in evals:
@@ -617,7 +614,7 @@ class DynataSurvey(MarketplaceTask):
group_eval = {
quota: quota.passes_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in group_eval.values())
+ evals = {g[0] for g in group_eval.values()}
if False in evals:
return False, set()
elif None in evals:
diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py
index 2b82bfd..e6f0548 100644
--- a/generalresearch/models/dynata/task_collection.py
+++ b/generalresearch/models/dynata/task_collection.py
@@ -6,7 +6,7 @@ import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.locales import Localelator
-from generalresearch.models import TaskCalculationType
+from generalresearch.models.definitions import TaskCalculationType
from generalresearch.models.dynata import DynataStatus
from generalresearch.models.dynata.survey import DynataSurvey
from generalresearch.models.thl.survey.task_collection import (
diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py
index 63ed2a1..d7dbfd6 100644
--- a/generalresearch/models/events.py
+++ b/generalresearch/models/events.py
@@ -1,6 +1,6 @@
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from enum import StrEnum
-from typing import Dict, Literal, Optional, Union
+from typing import Annotated, Literal
from uuid import uuid4
from pydantic import (
@@ -12,14 +12,13 @@ from pydantic import (
TypeAdapter,
model_validator,
)
-from typing_extensions import Annotated
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
UUIDStr,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
SessionStatusCode2,
Status,
@@ -63,7 +62,7 @@ class TaskEnterPayload(BaseModel):
source: Source = Field()
survey_id: str = Field(min_length=1, max_length=32, examples=["127492892"])
- quota_id: Optional[str] = Field(
+ quota_id: str | None = Field(
default=None,
max_length=32,
description="The marketplace's internal quota id",
@@ -76,9 +75,9 @@ class TaskFinishPayload(TaskEnterPayload):
duration_sec: PositiveFloat = Field()
status: Status
- status_code_1: Optional[StatusCode1] = None
- status_code_2: Optional[WallStatusCode2] = None
- cpi: Optional[NonNegativeInt] = Field(le=4000, default=None)
+ status_code_1: StatusCode1 | None = None
+ status_code_2: WallStatusCode2 | None = None
+ cpi: NonNegativeInt | None = Field(le=4000, default=None)
class SessionEnterPayload(BaseModel):
@@ -91,18 +90,13 @@ class SessionFinishPayload(SessionEnterPayload):
duration_sec: PositiveFloat = Field()
status: Status
- status_code_1: Optional[StatusCode1] = None
- status_code_2: Optional[SessionStatusCode2] = None
- user_payout: Optional[NonNegativeInt] = Field(default=None, le=4000, ge=0)
+ status_code_1: StatusCode1 | None = None
+ status_code_2: SessionStatusCode2 | None = None
+ user_payout: NonNegativeInt | None = Field(default=None, le=4000, ge=0)
EventPayload = Annotated[
- Union[
- TaskEnterPayload,
- TaskFinishPayload,
- SessionEnterPayload,
- SessionFinishPayload,
- ],
+ TaskEnterPayload | TaskFinishPayload | SessionEnterPayload | SessionFinishPayload,
Field(discriminator="event_type"),
]
@@ -110,12 +104,10 @@ EventPayload = Annotated[
class EventEnvelope(BaseModel):
event_uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex)
event_type: EventType = Field()
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
version: int = 1
- product_user_id: Optional[str] = Field(
+ product_user_id: str | None = Field(
min_length=3,
max_length=128,
examples=["app-user-9329ebd"],
@@ -136,7 +128,7 @@ class EventEnvelope(BaseModel):
class AggregateBySource(BaseModel):
total: NonNegativeInt = Field(default=0)
- by_source: Dict[Source, NonNegativeInt] = Field(default_factory=dict)
+ by_source: dict[Source, NonNegativeInt] = Field(default_factory=dict)
@model_validator(mode="after")
def remove_zero(self):
@@ -145,8 +137,8 @@ class AggregateBySource(BaseModel):
class MaxGaugeBySource(BaseModel):
- value: Optional[NonNegativeInt] = Field(default=None)
- by_source: Dict[Source, NonNegativeInt] = Field(default_factory=dict)
+ value: NonNegativeInt | None = Field(default=None)
+ by_source: dict[Source, NonNegativeInt] = Field(default_factory=dict)
@model_validator(mode="after")
def remove_zero(self):
@@ -174,11 +166,9 @@ class StatsSnapshot(TaskStatsSnapshot):
model_config = ConfigDict(ser_json_timedelta="float")
# If this is set, then everything is scoped to this country.
- country_iso: Optional[CountryISOLike] = Field(default=None)
+ country_iso: CountryISOLike | None = Field(default=None)
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# Counts: User related
active_users_last_1h: NonNegativeInt = Field(
@@ -217,17 +207,17 @@ class StatsSnapshot(TaskStatsSnapshot):
)
# Rolling averages
- session_avg_payout_last_24h: Optional[NonNegativeInt] = Field(
+ session_avg_payout_last_24h: NonNegativeInt | None = Field(
description="Average (actual) payout of all tasks completed in the past 24 hrs"
)
- session_avg_user_payout_last_24h: Optional[NonNegativeInt] = Field(
+ session_avg_user_payout_last_24h: NonNegativeInt | None = Field(
description="Average (actual) user payout of all tasks completed in the past 24 hrs"
)
- session_fail_avg_loi_last_24h: Optional[timedelta] = Field(
+ session_fail_avg_loi_last_24h: timedelta | None = Field(
description="Average LOI of all tasks terminated in the past 24 hrs (excludes abandons)"
)
- session_complete_avg_loi_last_24h: Optional[timedelta] = Field(
+ session_complete_avg_loi_last_24h: timedelta | None = Field(
description="Average LOI of all tasks completed in the past 24 hrs"
)
@@ -246,34 +236,26 @@ class StatsSnapshot(TaskStatsSnapshot):
class EventMessage(BaseModel):
kind: Literal[MessageKind.EVENT] = Field(default=MessageKind.EVENT)
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
data: EventEnvelope
class StatsMessage(BaseModel):
kind: Literal[MessageKind.STATS] = Field(default=MessageKind.STATS)
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# The data/StatsSnapshot can optionally be scoped to a country
- country_iso: Optional[CountryISOLike] = Field(default=None)
+ country_iso: CountryISOLike | None = Field(default=None)
data: StatsSnapshot
class PingMessage(BaseModel):
kind: Literal[MessageKind.PING] = Field(default=MessageKind.PING)
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
class PongMessage(BaseModel):
kind: Literal[MessageKind.PONG] = Field(default=MessageKind.PONG)
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
class SubscribeMessage(BaseModel):
@@ -281,17 +263,14 @@ class SubscribeMessage(BaseModel):
product_id: UUIDStr = Field(examples=["4fe381fb7186416cb443a38fa66c6557"])
-ServerToClientMessage = Union[EventMessage, StatsMessage, PingMessage]
+ServerToClientMessage = EventMessage | StatsMessage | PingMessage
ServerToClientMessageField = Annotated[
ServerToClientMessage,
Field(discriminator="kind"),
]
ServerToClientMessageAdapter = TypeAdapter(ServerToClientMessageField)
-ClientToServerMessage = Union[
- SubscribeMessage,
- PongMessage,
-]
+ClientToServerMessage = SubscribeMessage | PongMessage
ClientToServerMessageField = Annotated[
ClientToServerMessage,
Field(discriminator="kind"),
diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py
index 4ee70f9..41c3eaa 100644
--- a/generalresearch/models/gr/authentication.py
+++ b/generalresearch/models/gr/authentication.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import binascii
import json
import os
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any
from pydantic import (
@@ -15,17 +15,16 @@ from pydantic import (
PositiveInt,
field_validator,
)
-from typing_extensions import Self
from generalresearch.decorators import LOG
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
if TYPE_CHECKING:
from generalresearch.models.gr.business import Business
from generalresearch.models.gr.team import Team
from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class Claims(BaseModel):
@@ -165,8 +164,8 @@ class GRUser(BaseModel):
self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config)
self.prefetch_teams(pg_config=pg_config, redis_config=redis_config)
- business_uuids = self.business_uuids
- team_uuids = self.team_uuids
+ business_uuids = self.business_uuids or []
+ team_uuids = self.team_uuids or []
if len(business_uuids + team_uuids) == 0:
self.products = []
@@ -182,7 +181,7 @@ class GRUser(BaseModel):
team_products = pm.fetch_uuids(team_uuids=team_uuids) if team_uuids else []
products = {p.id: p for p in business_products + team_products}
- self.products = sorted(products.values(), key=lambda x: getattr(x, "created"))
+ self.products = sorted(products.values(), key=lambda x: x.created)
def prefetch_token(self, pg_config: PostgresConfig):
from generalresearch.managers.gr.authentication import (
@@ -199,7 +198,7 @@ class GRUser(BaseModel):
@field_validator("date_joined")
@classmethod
def date_joined_utc(cls, v: datetime) -> datetime:
- return v.replace(tzinfo=timezone.utc)
+ return v.replace(tzinfo=UTC)
# --- Properties ---
@property
@@ -284,17 +283,15 @@ class GRUser(BaseModel):
ex=ex_secs,
)
- return None
-
# --- ORM ---
@classmethod
- def from_postgresql(cls, d: dict) -> Self:
- d["date_joined"] = d["date_joined"].replace(tzinfo=timezone.utc)
+ def from_postgresql(cls, d: dict[str, Any]) -> GRUser:
+ d["date_joined"] = d["date_joined"].replace(tzinfo=UTC)
return GRUser.model_validate(d)
@classmethod
- def from_redis(cls, d: str | dict[str, Any]) -> Self:
+ def from_redis(cls, d: str | dict[str, Any]) -> GRUser:
if isinstance(d, str):
d = json.loads(d)
assert isinstance(d, dict)
@@ -354,18 +351,18 @@ class GRToken(BaseModel):
@field_validator("created", mode="before")
@classmethod
def created_utc(cls, v: datetime) -> datetime:
- return v.replace(tzinfo=timezone.utc)
+ return v.replace(tzinfo=UTC)
# --- Properties ---
@property
- def auth_header(self, key_name: str = "Authorization") -> dict[str, str]:
- return {key_name: self.key}
+ def auth_header(self) -> dict[str, str]:
+ return {"Authorization": self.key}
# --- ORM ---
@classmethod
- def from_redis(cls, d: str | dict[str, Any]) -> Self:
+ def from_redis(cls, d: str | dict[str, Any]) -> GRToken:
if isinstance(d, str):
d = json.loads(d)
assert isinstance(d, dict)
diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py
index 70aafc6..2104650 100644
--- a/generalresearch/models/gr/business.py
+++ b/generalresearch/models/gr/business.py
@@ -3,24 +3,22 @@ from __future__ import annotations
import json
import logging
import os
-from datetime import datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
from pathlib import Path
from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
-from dask.distributed import Client
+import pyarrow as pa
+from dask.distributed import Client as DaskClient
from psycopg.cursor import Cursor
from psycopg.rows import dict_row
-from pydantic import BaseModel, ConfigDict, Field, PositiveInt
+from pydantic import BaseModel, ConfigDict, Field, PositiveInt, ValidationError
from pydantic.json_schema import SkipJsonSchema
from pydantic_extra_types.phone_numbers import PhoneNumber
-from typing_extensions import Self
from generalresearch.currency import USDCent
from generalresearch.decorators import LOG
-from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.incite.schemas.mergers.pop_ledger import (
numerical_col_names,
)
@@ -30,13 +28,20 @@ from generalresearch.models.custom_types import (
UUIDStr,
UUIDStrCoerce,
)
+from generalresearch.models.gr.definitions import BusinessType, TransferMethod
+from generalresearch.models.gr.team import Team
from generalresearch.models.thl.finance import BusinessBalances, POPFinancial
from generalresearch.models.thl.ledger import LedgerAccount, OrderBy
from generalresearch.models.thl.payout import BusinessPayoutEvent
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
from generalresearch.utils.aggregation import group_by_year
-from generalresearch.utils.enum import ReprEnumMeta
+
+if TYPE_CHECKING:
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
+
+logging.basicConfig()
+logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
@@ -46,6 +51,9 @@ if TYPE_CHECKING:
from generalresearch.incite.mergers.foundations.enriched_wall import (
EnrichedWallMerge,
)
+ from generalresearch.managers.gr.business import (
+ BusinessBankAccountManager,
+ )
from generalresearch.managers.thl.ledger_manager.ledger import (
LedgerManager,
)
@@ -55,20 +63,10 @@ if TYPE_CHECKING:
from generalresearch.managers.thl.payout import (
BusinessPayoutEventManager,
)
- from generalresearch.models.gr.team import Team
+ from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.thl.product import Product
-class TransferMethod(Enum, metaclass=ReprEnumMeta):
- ACH = 0
- WIRE = 1
-
-
-class BusinessType(str, Enum, metaclass=ReprEnumMeta):
- INDIVIDUAL = "i"
- COMPANY = "c"
-
-
class BusinessBankAccount(BaseModel):
model_config = ConfigDict(
use_enum_values=True,
@@ -204,16 +202,18 @@ class Business(BaseModel):
# Initialization is deferred until unless it's called
# (see .prebuild_***())
- balance: BusinessBalances | None = Field(default=None, name="Business Balance")
+ balance: BusinessBalances | None = Field(default=None, title="Business Balance")
payouts_total_str: str | None = Field(default=None)
payouts_total: USDCent | None = Field(default=None)
payouts: list[BusinessPayoutEvent] | None = Field(
default=None,
- name="Business Payouts",
- description="These are the ACH or Wire payments that were sent to the"
- "Business as a single amount, summed for all the Business"
- "child Products",
+ title="Business Payouts",
+ description=(
+ "These are the ACH or Wire payments that were sent to the"
+ "Business as a single amount, summed for all the Business"
+ "child Products"
+ ),
)
pop_financial: list[POPFinancial] | None = Field(default=None)
@@ -238,18 +238,19 @@ class Business(BaseModel):
# --- Prefetch ---
def prefetch_addresses(self, pg_config: PostgresConfig) -> None:
- with pg_config.make_connection() as conn:
- with conn.cursor(row_factory=dict_row) as c:
- c.execute(
- query="""
+ with pg_config.make_connection() as conn, conn.cursor(
+ row_factory=dict_row
+ ) as c:
+ c.execute(
+ query="""
SELECT *
FROM common_businessaddress AS ba
WHERE ba.business_id = %s
LIMIT 1
""",
- params=[self.id],
- )
- res = c.fetchall()
+ params=[self.id],
+ )
+ res = c.fetchall()
if len(res) == 0:
self.addresses = []
@@ -257,55 +258,54 @@ class Business(BaseModel):
self.addresses = [BusinessAddress.model_validate(i) for i in res]
def prefetch_teams(self, pg_config: PostgresConfig) -> None:
- from generalresearch.models.gr.team import Team
+ with pg_config.make_connection() as conn, conn.cursor(
+ row_factory=dict_row
+ ) as c:
+ c: Cursor
- with pg_config.make_connection() as conn:
- with conn.cursor(row_factory=dict_row) as c:
- c: Cursor
-
- c.execute(
- query="""
+ c.execute(
+ query="""
SELECT t.*
FROM common_team AS t
INNER JOIN common_team_businesses AS tb
ON tb.team_id = t.id
WHERE tb.business_id = %s
""",
- params=(self.id,),
- )
+ params=(self.id,),
+ )
- res = c.fetchall()
+ res = c.fetchall()
if len(res) == 0:
self.teams = []
self.teams = [Team.model_validate(i) for i in res]
- def prefetch_products(self, thl_pg_config: PostgresConfig) -> None:
+ def prefetch_products(self, product_manager: ProductManager) -> None:
"""
:return: All the Products for this Business
"""
- from generalresearch.managers.thl.product import ProductManager
- pm = ProductManager(pg_config=thl_pg_config)
- self.products = pm.fetch_uuids(business_uuids=[self.uuid])
+ self.products = product_manager.fetch_uuids(business_uuids=[self.uuid])
- def prefetch_bank_accounts(self, pg_config: PostgresConfig) -> None:
- from generalresearch.managers.gr.business import (
- BusinessBankAccountManager,
+ def prefetch_bank_accounts(
+ self, business_bank_account_manager: BusinessBankAccountManager
+ ) -> None:
+ self.bank_accounts = business_bank_account_manager.get_by_business_id(
+ business_id=self.id
)
- bam = BusinessBankAccountManager(pg_config=pg_config)
- self.bank_accounts = bam.get_by_business_id(business_id=self.id)
-
def prefetch_bp_accounts(
- self, thl_lm: ThlLedgerManager, thl_pg_config: PostgresConfig
+ self, thl_lm: ThlLedgerManager, product_manager: ProductManager
):
# We need to prefetch the Products everytime because there is no way
# of knowing if a new Product has been added since the last time it
# ran.
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
+ assert isinstance(self.products, list)
product_lookup = {p.uuid: p for p in self.products}
+ assert isinstance(self.product_uuids, list)
+ assert thl_lm.currency
accounts = thl_lm.get_accounts_if_exists(
qualified_names=[
@@ -320,11 +320,12 @@ class Business(BaseModel):
for product_uuid in self.product_uuids:
if product_uuid not in bp_account:
refresh = True
- logging.exception(
+ logger.exception(
f"Business {self.uuid} does not have a BP Wallet Account for Product {product_uuid}. Creating..."
)
product = product_lookup[product_uuid]
thl_lm.get_account_or_create_bp_wallet(product=product)
+
if refresh:
accounts = thl_lm.get_accounts_if_exists(
qualified_names=[
@@ -341,10 +342,10 @@ class Business(BaseModel):
def prebuild_balance(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
lm: LedgerManager,
ds: GRLDatasets,
- client: Client,
+ client: DaskClient,
pop_ledger: PopLedgerMerge | None = None,
at_timestamp: AwareDatetime | None = None,
) -> None:
@@ -371,8 +372,9 @@ class Business(BaseModel):
volume levels.
"""
LOG.debug(f"Business.prebuild_balance({self.uuid=})")
+ assert lm.currency
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
accounts: list[LedgerAccount] = lm.get_accounts_if_exists(
qualified_names=(
@@ -391,8 +393,8 @@ class Business(BaseModel):
pop_ledger = plm(ds=ds)
if at_timestamp is None:
- at_timestamp = datetime.now(tz=timezone.utc)
- assert at_timestamp.tzinfo == timezone.utc
+ at_timestamp = datetime.now(tz=UTC)
+ assert at_timestamp.tzinfo == UTC
ddf = pop_ledger.ddf(
force_rr_latest=False,
@@ -421,33 +423,27 @@ class Business(BaseModel):
# that is still valid. Don't attempt to build a balance, leave it
# as None rather than all zeros
LOG.warning(f"Business({self.uuid=}).prebuild_balance empty dataframe")
- return None
+ return
LOG.debug(f"Business.prebuild_balance.groupby() {df.head()}")
df = df.groupby("account_id").sum()
self.balance = BusinessBalances.from_pandas(
- input_data=df, accounts=accounts, thl_pg_config=thl_pg_config
+ input_data=df, accounts=accounts, product_manager=product_manager
)
return
def prebuild_payouts(
self,
- thl_pg_config: PostgresConfig,
- thl_lm: ThlLedgerManager,
bpem: BusinessPayoutEventManager,
) -> None:
LOG.debug(f"Business.prebuild_payouts({self.uuid=})")
- self.prefetch_products(thl_pg_config=thl_pg_config)
-
- self.payouts = bpem.get_business_payout_events_for_products(
- thl_ledger_manager=thl_lm,
- product_uuids=self.product_uuids,
+ self.payouts = bpem.get_business_payout_events_for_business(
+ business_uuid=self.uuid,
order_by=OrderBy.DESC,
)
-
self.prebuild_payouts_total()
def prebuild_payouts_total(self):
@@ -455,14 +451,12 @@ class Business(BaseModel):
self.payouts_total = USDCent(sum([po.amount for po in self.payouts]))
self.payouts_total_str = self.payouts_total.to_usd_str()
- return
-
def prebuild_pop_financial(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
thl_lm: ThlLedgerManager,
ds: GRLDatasets,
- client: Client,
+ client: DaskClient,
pop_ledger: PopLedgerMerge | None = None,
) -> None:
"""This is very similar to the Product POP Financial endpoint; however,
@@ -471,7 +465,8 @@ class Business(BaseModel):
financial activity within that time window.
"""
if self.bp_accounts is None:
- self.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_pg_config)
+ self.prefetch_bp_accounts(thl_lm=thl_lm, product_manager=product_manager)
+ assert isinstance(self.bp_accounts, list)
from generalresearch.models.admin.request import (
ReportRequest,
@@ -514,13 +509,13 @@ class Business(BaseModel):
def prebuild_enriched_session_parquet(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
ds: GRLDatasets,
- client: Client,
+ client: DaskClient,
mnt_gr_api: Path,
enriched_session: EnrichedSessionMerge | None = None,
) -> None:
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
if enriched_session is None:
from generalresearch.incite.defaults import (
@@ -551,21 +546,19 @@ class Business(BaseModel):
)
try:
- test = pd.read_parquet(path, engine="pyarrow")
- except Exception as e:
+ _ = pd.read_parquet(path, engine="pyarrow")
+ except (pa.ArrowException, OSError, ValueError) as e:
raise OSError(f"Parquet verification failed: {e}")
- return None
-
def prebuild_enriched_wall_parquet(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
ds: GRLDatasets,
- client: Client,
+ client: DaskClient,
mnt_gr_api: Path,
enriched_wall: EnrichedWallMerge | None = None,
) -> None:
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
if enriched_wall is None:
from generalresearch.incite.defaults import (
@@ -596,12 +589,10 @@ class Business(BaseModel):
)
try:
- test = pd.read_parquet(path, engine="pyarrow")
- except Exception as e:
+ _ = pd.read_parquet(path, engine="pyarrow")
+ except (pa.ArrowException, OSError, ValueError) as e:
raise OSError(f"Parquet verification failed: {e}")
- return None
-
@classmethod
def required_fields(cls) -> list[str]:
return [
@@ -633,9 +624,11 @@ class Business(BaseModel):
def set_cache(
self,
pg_config: PostgresConfig,
+ product_manager: ProductManager,
+ business_bank_account_manager: BusinessBankAccountManager,
thl_web_rr: PostgresConfig,
redis_config: RedisConfig,
- client: Client,
+ client: DaskClient,
ds: GRLDatasets,
lm: LedgerManager,
thl_lm: ThlLedgerManager,
@@ -651,20 +644,22 @@ class Business(BaseModel):
self.prefetch_addresses(pg_config=pg_config)
self.prefetch_teams(pg_config=pg_config)
- self.prefetch_products(thl_pg_config=thl_web_rr)
- self.prefetch_bank_accounts(pg_config=pg_config)
- self.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
+ self.prefetch_products(product_manager=product_manager)
+ self.prefetch_bank_accounts(
+ business_bank_account_manager=business_bank_account_manager
+ )
+ self.prefetch_bp_accounts(thl_lm=thl_lm, product_manager=product_manager)
self.prebuild_balance(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
lm=lm,
ds=ds,
client=client,
pop_ledger=pop_ledger,
)
- self.prebuild_payouts(thl_pg_config=thl_web_rr, thl_lm=thl_lm, bpem=bpem)
+ self.prebuild_payouts(bpem=bpem)
self.prebuild_pop_financial(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
thl_lm=thl_lm,
ds=ds,
client=client,
@@ -696,7 +691,7 @@ class Business(BaseModel):
enriched_session = es(ds=ds)
self.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
client=client,
ds=ds,
mnt_gr_api=mnt_gr_api,
@@ -709,7 +704,7 @@ class Business(BaseModel):
enriched_wall = ew(ds=ds)
self.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
client=client,
ds=ds,
mnt_gr_api=mnt_gr_api,
@@ -724,18 +719,18 @@ class Business(BaseModel):
uuid: UUIDStr,
fields: list[str],
gr_redis_config: RedisConfig,
- ) -> Self | None:
+ ) -> Business | None:
keys: list[str] = Business.required_fields() + fields
if "pop_financial" in keys:
# We should explicitly pass the pop_financial years we want. By default,
# at least get this year.
- year = datetime.now(tz=timezone.utc).year
+ year = datetime.now(tz=UTC).year
keys = list(set(keys) | {f"pop_financial:{year}"})
rc = gr_redis_config.create_redis_client()
try:
- res: list = rc.hmget(name=f"business:{uuid}", keys=keys)
+ res: list[str | bytes | None] = rc.hmget(name=f"business:{uuid}", keys=keys)
d = {
val: json.loads(res[idx]) if res[idx] is not None else None
for idx, val in enumerate(keys)
@@ -753,6 +748,5 @@ class Business(BaseModel):
result["pop_financial"] = pop_financial
return Business.model_validate(result)
- except Exception as e:
- logging.exception(e)
+ except ValidationError:
return None
diff --git a/generalresearch/models/gr/definitions.py b/generalresearch/models/gr/definitions.py
new file mode 100644
index 0000000..2e06c03
--- /dev/null
+++ b/generalresearch/models/gr/definitions.py
@@ -0,0 +1,13 @@
+from enum import Enum, StrEnum
+
+from generalresearch.utils.enum import ReprEnumMeta
+
+
+class TransferMethod(Enum, metaclass=ReprEnumMeta):
+ ACH = 0
+ WIRE = 1
+
+
+class BusinessType(StrEnum, metaclass=ReprEnumMeta):
+ INDIVIDUAL = "i"
+ COMPANY = "c"
diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py
index 8d60825..8d23bc5 100644
--- a/generalresearch/models/gr/team.py
+++ b/generalresearch/models/gr/team.py
@@ -2,46 +2,53 @@ from __future__ import annotations
import json
import os
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
+import pyarrow as pa
from dask.distributed import Client
from pydantic import (
BaseModel,
ConfigDict,
Field,
PositiveInt,
+ ValidationError,
field_validator,
)
from pydantic.json_schema import SkipJsonSchema
-from typing_extensions import Self
from generalresearch.decorators import LOG
-from generalresearch.incite.mergers.foundations.enriched_session import (
- EnrichedSessionMerge,
-)
-from generalresearch.incite.mergers.foundations.enriched_wall import (
- EnrichedWallMerge,
-)
from generalresearch.models.admin.request import ReportRequest, ReportType
from generalresearch.models.custom_types import (
AwareDatetimeISO,
UUIDStr,
UUIDStrCoerce,
)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
from generalresearch.utils.enum import ReprEnumMeta
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+ )
+ from generalresearch.managers.gr.authentication import (
+ GRUserManager,
+ )
+ from generalresearch.managers.gr.business import BusinessManager
+ from generalresearch.managers.gr.team import MembershipManager
+ from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.gr.authentication import GRUser
from generalresearch.models.gr.business import Business
from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class MembershipPrivilege(Enum, metaclass=ReprEnumMeta):
@@ -92,7 +99,7 @@ class Membership(BaseModel):
@classmethod
def created_utc(cls, v: datetime | str) -> datetime | str:
if isinstance(v, datetime):
- return v.replace(tzinfo=timezone.utc)
+ return v.replace(tzinfo=UTC)
return v
# --- prefetch methods ---
@@ -119,48 +126,29 @@ class Team(BaseModel):
# --- Prefetch Methods ---
- def prefetch_memberships(self, pg_config: PostgresConfig) -> None:
- from generalresearch.managers.gr.team import MembershipManager
+ def prefetch_memberships(self, gr_membership_manager: MembershipManager) -> None:
+ self.memberships = gr_membership_manager.get_by_team_id(team_id=self.id)
- mm = MembershipManager(pg_config=pg_config)
- self.memberships = mm.get_by_team_id(team_id=self.id)
+ def prefetch_gr_users(self, gr_user_manager: GRUserManager) -> None:
+ self.gr_users = gr_user_manager.get_by_team(team_id=self.id)
- def prefetch_gr_users(
- self, pg_config: PostgresConfig, redis_config: RedisConfig
- ) -> None:
- from generalresearch.managers.gr.authentication import (
- GRUserManager,
- )
+ def prefetch_businesses(self, gr_business_manager: BusinessManager) -> None:
+ self.businesses = gr_business_manager.get_by_team(team_id=self.id)
- gr_um = GRUserManager(pg_config=pg_config, redis_config=redis_config)
-
- self.gr_users = gr_um.get_by_team(team_id=self.id)
-
- def prefetch_businesses(
- self, pg_config: PostgresConfig, redis_config: RedisConfig
- ) -> None:
- from generalresearch.managers.gr.business import BusinessManager
-
- bm = BusinessManager(pg_config=pg_config, redis_config=redis_config)
- self.businesses = bm.get_by_team(team_id=self.id)
-
- def prefetch_products(self, thl_pg_config: PostgresConfig) -> None:
- from generalresearch.managers.thl.product import ProductManager
-
- pm = ProductManager(pg_config=thl_pg_config)
- self.products = pm.fetch_uuids(team_uuids=[self.uuid])
+ def prefetch_products(self, product_manager: ProductManager) -> None:
+ self.products = product_manager.fetch_uuids(team_uuids=[self.uuid])
# --- Prebuild Methods ---
def prebuild_enriched_session_parquet(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
ds: GRLDatasets,
client: Client,
mnt_gr_api: Path,
enriched_session: EnrichedSessionMerge | None = None,
) -> None:
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
if enriched_session is None:
from generalresearch.incite.defaults import (
@@ -192,20 +180,18 @@ class Team(BaseModel):
try:
_ = pd.read_parquet(path, engine="pyarrow")
- except Exception as e:
+ except (pa.ArrowException, OSError, ValueError) as e:
raise OSError(f"Parquet verification failed: {e}")
- return
-
def prebuild_enriched_wall_parquet(
self,
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
ds: GRLDatasets,
client: Client,
mnt_gr_api: Path,
enriched_wall: EnrichedWallMerge | None = None,
) -> None:
- self.prefetch_products(thl_pg_config=thl_pg_config)
+ self.prefetch_products(product_manager=product_manager)
if enriched_wall is None:
from generalresearch.incite.defaults import (
@@ -237,11 +223,9 @@ class Team(BaseModel):
try:
_ = pd.read_parquet(path, engine="pyarrow")
- except Exception as e:
+ except (pa.ArrowException, OSError, ValueError) as e:
raise OSError(f"Parquet verification failed: {e}")
- return None
-
@classmethod
def required_fields(cls) -> list[str]:
return [
@@ -275,8 +259,10 @@ class Team(BaseModel):
def set_cache(
self,
- pg_config: PostgresConfig,
- thl_web_rr: PostgresConfig,
+ product_manager: ProductManager,
+ gr_user_manager: GRUserManager,
+ gr_business_manager: BusinessManager,
+ gr_membership_manager: MembershipManager,
redis_config: RedisConfig,
client: Client,
ds: GRLDatasets,
@@ -284,12 +270,10 @@ class Team(BaseModel):
enriched_session: EnrichedSessionMerge | None = None,
enriched_wall: EnrichedWallMerge | None = None,
) -> None:
- ex_secs = 60 * 60 * 24 * 3 # 3 days
-
- self.prefetch_products(thl_pg_config=thl_web_rr)
- self.prefetch_gr_users(pg_config=pg_config, redis_config=redis_config)
- self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config)
- self.prefetch_memberships(pg_config=pg_config)
+ self.prefetch_products(product_manager=product_manager)
+ self.prefetch_gr_users(gr_user_manager=gr_user_manager)
+ self.prefetch_businesses(gr_business_manager=gr_business_manager)
+ self.prefetch_memberships(gr_membership_manager=gr_membership_manager)
rc = redis_config.create_redis_client()
mapping = self.model_dump(mode="json")
@@ -306,7 +290,7 @@ class Team(BaseModel):
enriched_session = es(ds=ds)
self.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
client=client,
ds=ds,
mnt_gr_api=mnt_gr_api,
@@ -319,15 +303,13 @@ class Team(BaseModel):
enriched_wall = ew(ds=ds)
self.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ product_manager=product_manager,
client=client,
ds=ds,
mnt_gr_api=mnt_gr_api,
enriched_wall=enriched_wall,
)
- return
-
# --- ORM ---
@classmethod
@@ -336,14 +318,14 @@ class Team(BaseModel):
uuid: UUIDStr,
fields: list[str],
gr_redis_config: RedisConfig,
- ) -> Self | None:
+ ) -> Team | None:
keys: list = Team.required_fields() + fields
rc = gr_redis_config.create_redis_client()
try:
- res: list = rc.hmget(name=f"team:{uuid}", keys=keys)
+ res: list[str | bytes | None] = rc.hmget(name=f"team:{uuid}", keys=keys)
d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)}
return Team.model_validate(d)
- except (Exception,) as e:
+ except ValidationError:
return None
diff --git a/generalresearch/models/innovate/__init__.py b/generalresearch/models/innovate/__init__.py
index 054c69d..9d7ea75 100644
--- a/generalresearch/models/innovate/__init__.py
+++ b/generalresearch/models/innovate/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
+from typing import Annotated
from pydantic import StringConstraints
-from typing_extensions import Annotated
# Note, this is called the KEY in the Question model
InnovateQuestionID = Annotated[
@@ -9,17 +9,17 @@ InnovateQuestionID = Annotated[
]
-class InnovateStatus(str, Enum):
+class InnovateStatus(StrEnum):
LIVE = "LIVE"
NOT_LIVE = "NOT_LIVE"
-class InnovateQuotaStatus(str, Enum):
+class InnovateQuotaStatus(StrEnum):
OPEN = "OPEN"
CLOSED = "CLOSED"
-class InnovateDuplicateCheckLevel(str, Enum):
+class InnovateDuplicateCheckLevel(StrEnum):
# How we should check for de-dupes / survey exclusions.
# https://innovatemr.stoplight.io/docs/supplier-api/ZG9jOjEzNzYxMTg2-statuses-term-reasons-and-categories
# #duplicatedtoken
diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py
index 306274f..4af0639 100644
--- a/generalresearch/models/innovate/question.py
+++ b/generalresearch/models/innovate/question.py
@@ -3,22 +3,20 @@ from __future__ import annotations
import json
import logging
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Literal
+from enum import StrEnum
+from typing import Any, Literal
-from pydantic import BaseModel, Field, field_validator, model_validator
+from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.innovate import InnovateQuestionID
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
MarketplaceUserQuestionAnswer,
)
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -50,7 +48,7 @@ class InnovateQuestionOption(BaseModel):
order: int = Field()
-class InnovateQuestionType(str, Enum):
+class InnovateQuestionType(StrEnum):
# API response: {'Multipunch', 'Numeric Open Ended', 'Single Punch'}
# "Numeric Open Ended" must be wrong... It can't be numeric, as UK's
# postcode question is marked as this, but it wants alphanumeric
@@ -71,7 +69,7 @@ class InnovateQuestionType(str, Enum):
@classmethod
def from_api(cls, a: int):
API_TYPE_MAP = cls.get_api_map()
- return API_TYPE_MAP[a] if a in API_TYPE_MAP else None
+ return API_TYPE_MAP.get(a)
class InnovateQuestion(MarketplaceQuestion):
@@ -141,7 +139,7 @@ class InnovateQuestion(MarketplaceQuestion):
@classmethod
def from_api(
- cls, d: dict, country_iso: str, language_iso: str
+ cls, d: dict[str, Any], country_iso: str, language_iso: str
) -> InnovateQuestion | None:
"""
:param d: Raw response from API
@@ -151,13 +149,13 @@ class InnovateQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d, country_iso, language_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
@classmethod
def _from_api(
- cls, d: dict, country_iso: str, language_iso: str
+ cls, d: dict[str, Any], country_iso: str, language_iso: str
) -> InnovateQuestion:
# Question AGE returns options even though its marked as a text entry (but only in some locales)
d["QuestionKey"] = d["QuestionKey"].lower()
diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py
index bcd50d3..3ca990c 100644
--- a/generalresearch/models/innovate/survey.py
+++ b/generalresearch/models/innovate/survey.py
@@ -2,14 +2,14 @@ from __future__ import annotations
import json
import logging
-from datetime import date, timezone
+from datetime import UTC, date
from decimal import Decimal
from functools import cached_property
from typing import (
Annotated,
Any,
Literal,
- Type,
+ Self,
)
from more_itertools import flatten
@@ -17,23 +17,23 @@ from pydantic import (
BaseModel,
ConfigDict,
Field,
+ ValidationError,
computed_field,
model_validator,
)
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import (
- LogicalOperator,
- Source,
- TaskCalculationType,
-)
from generalresearch.models.custom_types import (
AlphaNumStrSet,
AwareDatetimeISO,
CoercedStr,
DeviceTypes,
)
+from generalresearch.models.definitions import (
+ LogicalOperator,
+ Source,
+ TaskCalculationType,
+)
from generalresearch.models.innovate import (
InnovateDuplicateCheckLevel,
InnovateQuotaStatus,
@@ -70,7 +70,7 @@ class InnovateCondition(MarketplaceCondition):
d["logical_operator"] = LogicalOperator.OR
d["value_type"] = ConditionValueType.LIST
d["negate"] = False
- d["values"] = list(set(x.strip().lower() for x in d["values"]))
+ d["values"] = list({x.strip().lower() for x in d["values"]})
return cls.model_validate(d)
@@ -89,7 +89,7 @@ class InnovateQuota(BaseModel):
condition_hashes: list[str] = Field(min_length=0, default_factory=list)
def __hash__(self):
- return hash(tuple((tuple(self.condition_hashes), self.remaining_count)))
+ return hash((tuple(self.condition_hashes), self.remaining_count))
@property
def is_open(self) -> bool:
@@ -100,7 +100,7 @@ class InnovateQuota(BaseModel):
)
@classmethod
- def from_api(cls, d: dict):
+ def from_api(cls, d: dict[str, Any]):
return cls.model_validate(d)
def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
@@ -264,13 +264,13 @@ class InnovateSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> InnovateSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@classmethod
def _from_api(cls, d: dict[str, Any]) -> InnovateSurvey:
- d["conditions"] = dict()
+ d["conditions"] = {}
# If we haven't hit the "detail" endpoint, we won't get this
d.setdefault("qualifications", [])
@@ -290,7 +290,7 @@ class InnovateSurvey(MarketplaceTask):
return cls.model_validate(d)
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return InnovateCondition
@property
@@ -318,11 +318,14 @@ class InnovateSurvey(MarketplaceTask):
# Fancy repr that abbreviates exclude_pids and excluded_surveys
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k in {"exclude_pids", "include_pids", "excluded_surveys"}:
- if v and len(v) > 6:
- v = sorted(v)
- v = v[:3] + ["…"] + v[-3:]
- repr_args[n] = (k, v)
+ if (
+ k in {"exclude_pids", "include_pids", "excluded_surveys"}
+ and v
+ and len(v) > 6
+ ):
+ v = sorted(v)
+ v = v[:3] + ["…"] + v[-3:]
+ repr_args[n] = (k, v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -361,10 +364,10 @@ class InnovateSurvey(MarketplaceTask):
@classmethod
def from_db(cls, d: dict[str, Any]) -> Self:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
- d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc)
- d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc)
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
+ d["modified_api"] = d["modified_api"].replace(tzinfo=UTC)
+ d["created_api"] = d["created_api"].replace(tzinfo=UTC)
d["qualifications"] = json.loads(d["qualifications"])
d["used_question_ids"] = json.loads(d["used_question_ids"])
d["quotas"] = json.loads(d["quotas"])
@@ -381,14 +384,21 @@ class InnovateSurvey(MarketplaceTask):
"""
assert isinstance(att_survey_ids, set), "must pass a set"
assert isinstance(att_job_ids, set), "must pass a set"
+
if self.survey_id in att_survey_ids:
return False
- if self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB:
- if self.job_id in att_job_ids:
- return False
+
+ if (
+ self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB
+ and self.job_id in att_job_ids
+ ):
+ return False
+
if self.duplicate_check_level == InnovateDuplicateCheckLevel.EXCLUDED_SURVEYS:
+ assert self.excluded_surveys is not None
if self.excluded_surveys.intersection(att_survey_ids):
return False
+
return True
def passes_qualifications(
@@ -432,7 +442,7 @@ class InnovateSurvey(MarketplaceTask):
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py
index 2650b0b..23eeb38 100644
--- a/generalresearch/models/legacy/bucket.py
+++ b/generalresearch/models/legacy/bucket.py
@@ -4,7 +4,7 @@ import logging
import math
from datetime import timedelta
from decimal import Decimal
-from typing import Any, Literal
+from typing import Any, Literal, Self
from pydantic import (
BaseModel,
@@ -14,14 +14,13 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
HttpsUrl,
PropertyCode,
UUIDStr,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.stats import StatisticalSummary
logger = logging.getLogger()
@@ -121,8 +120,10 @@ class BucketBase(BaseModel):
)
uri: HttpsUrl = Field(
examples=[
- "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
- "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ (
+ "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
+ "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ )
],
description="The URL to send a respondent into. Must not edit this URL in any way",
)
@@ -440,6 +441,12 @@ class DurationSummary(StatisticalSummary):
@classmethod
def from_bucket(cls, bucket: Bucket) -> DurationSummary:
+ assert bucket.loi_min
+ assert bucket.loi_max
+ assert bucket.loi_q1
+ assert bucket.loi_q2
+ assert bucket.loi_q3
+
return cls(
min=bucket.loi_min.total_seconds(),
max=bucket.loi_max.total_seconds(),
@@ -466,12 +473,12 @@ class PayoutSummaryDecimal(StatisticalSummary):
class PayoutSummary(StatisticalSummary):
"""Payouts are in Integer USD Cents"""
- min: int = Field(gt=0, le=10000)
- max: int = Field(gt=0, le=10000)
- q1: int = Field(gt=0, le=10000)
- q2: int = Field(gt=0, le=10000)
- q3: int = Field(gt=0, le=10000)
- mean: int | None = Field(gt=0, le=10000, default=None)
+ min: int = Field(gt=0, le=10_000)
+ max: int = Field(gt=0, le=10_000)
+ q1: int = Field(gt=0, le=10_000)
+ q2: int = Field(gt=0, le=10_000)
+ q3: int = Field(gt=0, le=10_000)
+ mean: int | None = Field(gt=0, le=10_000, default=None)
model_config = {
"json_schema_extra": {
@@ -725,8 +732,10 @@ class OneShotOfferwallBucket(BaseModel):
)
uri: HttpsUrl = Field(
examples=[
- "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
- "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ (
+ "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
+ "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ )
],
description="The URL to send a respondent into. Must not edit this URL in any way",
)
@@ -760,8 +769,10 @@ class WXETOfferwallBucket(BaseModel):
)
uri: HttpsUrl = Field(
examples=[
- "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
- "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ (
+ "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
+ "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ )
],
description="The URL to send a respondent into. Must not edit this URL in any way",
)
diff --git a/generalresearch/models/legacy/definitions.py b/generalresearch/models/legacy/definitions.py
index 1755d2a..b3a6687 100644
--- a/generalresearch/models/legacy/definitions.py
+++ b/generalresearch/models/legacy/definitions.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class OfferwallReason(str, Enum):
+class OfferwallReason(StrEnum):
USER_BLOCKED = "USER_BLOCKED"
HIGH_RECON_RATE = "HIGH_RECON_RATE"
UNCOMMON_DEMOGRAPHICS = "UNCOMMON_DEMOGRAPHICS"
diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py
index 81e794c..caa6aae 100644
--- a/generalresearch/models/legacy/questions.py
+++ b/generalresearch/models/legacy/questions.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from typing import TYPE_CHECKING, Any
+from typing import TYPE_CHECKING, Annotated, Any
from pydantic import (
BaseModel,
@@ -14,7 +14,6 @@ from pydantic import (
model_validator,
)
from sentry_sdk import capture_exception
-from typing_extensions import Annotated, Self
from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.legacy.api_status import StatusResponse
@@ -25,9 +24,7 @@ from generalresearch.models.thl.session import Wall
from generalresearch.models.thl.user import User
if TYPE_CHECKING:
- from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
- )
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
from generalresearch.managers.thl.wall import WallManager
@@ -88,26 +85,17 @@ class UserQuestionAnswerIn(BaseModel):
fingerprint_tz = "a91cb1dea814480dba12d9b7b48696dd"
fingerprint_fingerprint = "1d1e2e8380ac474b87fb4e4c569b48df"
- if self.question_id in {
- user_agent_qid,
- fingerprint_langs,
- fingerprint_tz,
- fingerprint_fingerprint,
- }:
- if len(self.answer) != 1:
- raise ValueError("Too many answer values provided")
-
- return self
-
- @model_validator(mode="after")
- def user_agent_check(self) -> Self:
- # TODO: where / how do I want to pass in this Werz user_agent stuff?
- user_agent_qid = "2fbedb2b9f7647b09ff5e52fa119cc5e"
-
- if self.question_id == user_agent_qid:
- val = self.answer[0]
- # assert val == request.user_agent.to_header():
- pass
+ if (
+ self.question_id
+ in {
+ user_agent_qid,
+ fingerprint_langs,
+ fingerprint_tz,
+ fingerprint_fingerprint,
+ }
+ and len(self.answer) != 1
+ ):
+ raise ValueError("Too many answer values provided")
return self
@@ -218,7 +206,6 @@ class UserQuestionAnswers(BaseModel):
# --- Prefetch ---
def prefetch_user(self, um: UserManager) -> None:
- from generalresearch.models.thl.user import User
res: User | None = um.get_user_if_exists(
product_id=self.product_id, product_user_id=self.product_user_id
@@ -230,8 +217,7 @@ class UserQuestionAnswers(BaseModel):
self.user = res
def prefetch_wall(self, wm: WallManager) -> None:
- from generalresearch.models import Source
- from generalresearch.models.thl.session import Wall
+ from generalresearch.models.definitions import Source
res: Wall | None = wm.get_from_uuid_if_exists(wall_uuid=self.session_id)
diff --git a/generalresearch/models/lucid/__init__.py b/generalresearch/models/lucid/__init__.py
index c3365db..4653339 100644
--- a/generalresearch/models/lucid/__init__.py
+++ b/generalresearch/models/lucid/__init__.py
@@ -1,5 +1,6 @@
+from typing import Annotated
+
from pydantic import Field
-from typing_extensions import Annotated
LucidQuestionIdType = Annotated[
str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py
index 908ce70..288f0d2 100644
--- a/generalresearch/models/lucid/question.py
+++ b/generalresearch/models/lucid/question.py
@@ -1,22 +1,19 @@
from __future__ import annotations
import logging
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Literal
+from enum import StrEnum
+from typing import Any, Literal, Self
from pydantic import BaseModel, Field, field_validator, model_validator
-from typing_extensions import Self
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.lucid import LucidQuestionIdType
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
)
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -41,7 +38,7 @@ class LucidQuestionOption(BaseModel):
order: int = Field()
-class LucidQuestionType(str, Enum):
+class LucidQuestionType(StrEnum):
SINGLE_SELECT = "s"
MULTI_SELECT = "m"
TEXT_ENTRY = "t"
diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py
index 4b1bb98..bca471b 100644
--- a/generalresearch/models/lucid/survey.py
+++ b/generalresearch/models/lucid/survey.py
@@ -4,13 +4,13 @@ from typing import Any, Self
from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
BigAutoInteger,
CoercedStr,
UUIDStr,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.locales import CountryISO, LanguageISO
from generalresearch.models.thl.survey.condition import (
ConditionValueType,
diff --git a/generalresearch/models/marketplace/summary.py b/generalresearch/models/marketplace/summary.py
index f75c530..9551417 100644
--- a/generalresearch/models/marketplace/summary.py
+++ b/generalresearch/models/marketplace/summary.py
@@ -2,11 +2,10 @@ from __future__ import annotations
from abc import ABC
from collections.abc import Collection
-from typing import Literal
+from typing import Literal, Self
import numpy as np
from pydantic import BaseModel, ConfigDict, Field, computed_field
-from typing_extensions import Self
from generalresearch.models.thl.stats import StatisticalSummary
diff --git a/generalresearch/models/morning/__init__.py b/generalresearch/models/morning/__init__.py
index 2c61c49..e747586 100644
--- a/generalresearch/models/morning/__init__.py
+++ b/generalresearch/models/morning/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
+from typing import Annotated
from pydantic import StringConstraints
-from typing_extensions import Annotated
# This is text-based, in lowercase. e.g. 'age', 'household_income'
MorningQuestionID = Annotated[
@@ -9,7 +9,7 @@ MorningQuestionID = Annotated[
]
-class MorningStatus(str, Enum):
+class MorningStatus(StrEnum):
DRAFT = "draft"
ACTIVE = "active" # aka LIVE
PAUSED = "paused"
diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py
index 0ab5030..7c676fb 100644
--- a/generalresearch/models/morning/question.py
+++ b/generalresearch/models/morning/question.py
@@ -1,13 +1,12 @@
import json
-from enum import Enum
-from typing import Any, Literal, Dict, List, Optional
+from enum import StrEnum
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import BaseModel, Field, field_validator, model_validator
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.morning import MorningQuestionID
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
@@ -37,7 +36,7 @@ class MorningQuestionOption(BaseModel, frozen=True):
order: int = Field()
-class MorningQuestionType(str, Enum):
+class MorningQuestionType(StrEnum):
# The db stores these as a single letter
# Geographic questions represent geographic areas within a country.
@@ -54,7 +53,7 @@ class MorningQuestionType(str, Enum):
class MorningUserQuestionAnswer(MarketplaceUserQuestionAnswer):
question_id: MorningQuestionID = Field()
- question_type: Optional[MorningQuestionType] = Field(default=None)
+ question_type: MorningQuestionType | None = Field(default=None)
# Did this answer come from us asking, or was it passed back from the
# marketplace? Note, morning doesn't "pass back" answers, but we can
# retrieve a user's profile through API, so it is possible to populate
@@ -92,7 +91,7 @@ class MorningQuestion(MarketplaceQuestion):
frozen=True,
)
# API calls this "responses", but I think that is a confusing name
- options: Optional[List[MorningQuestionOption]] = Field(
+ options: list[MorningQuestionOption] | None = Field(
default=None, min_length=1, frozen=True
)
@@ -119,7 +118,7 @@ class MorningQuestion(MarketplaceQuestion):
return options
@classmethod
- def from_api(cls, d: Dict[str, Any], country_iso: str, language_iso: str):
+ def from_api(cls, d: dict[str, Any], country_iso: str, language_iso: str):
options = None
if d.get("responses"):
options = [
@@ -138,7 +137,7 @@ class MorningQuestion(MarketplaceQuestion):
)
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> Self:
+ def from_db(cls, d: dict[str, Any]) -> Self:
options = None
if d["options"]:
options = [
@@ -162,7 +161,7 @@ class MorningQuestion(MarketplaceQuestion):
),
)
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(mode="json", by_alias=True)
d["options"] = json.dumps(d["options"])
return d
diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py
index 255a63c..8738631 100644
--- a/generalresearch/models/morning/survey.py
+++ b/generalresearch/models/morning/survey.py
@@ -2,19 +2,14 @@ from __future__ import annotations
import json
import logging
-from datetime import timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
from typing import (
Annotated,
Any,
- Dict,
- List,
Literal,
- Optional,
- Set,
- Tuple,
- Type,
+ Self,
)
from pydantic import (
@@ -27,14 +22,13 @@ from pydantic import (
computed_field,
model_validator,
)
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
UUIDStrCoerce,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.morning import MorningQuestionID, MorningStatus
from generalresearch.models.morning.question import MorningQuestion
from generalresearch.models.thl.demographics import Gender
@@ -79,14 +73,14 @@ class MorningStatistics(BaseModel):
# bid bid_loi: int = Field(validation_alias="estimated_length_of_interview",
# le=120 * 60)
# If num_completes == 0 , this gets returned as 0. it should be None
- obs_median_loi: Optional[NonNegativeInt] = Field(
+ obs_median_loi: NonNegativeInt | None = Field(
validation_alias="median_length_of_interview", default=None, le=120 * 60
)
# API returns 100 until 5 completes! Should be None.
# This is calculated as the total completes divided by the total number of
# finished sessions that passed the prescreener.
- qualified_conversion: Optional[float] = Field(
+ qualified_conversion: float | None = Field(
ge=0, le=1, description="conversion rate of qualified respondents"
)
@@ -142,7 +136,7 @@ class MorningTaskStatistics(MorningStatistics):
# relevant to quotas.
# API returns 100 until 5 completes! Should be None ...
- system_conversion: Optional[float] = Field(
+ system_conversion: float | None = Field(
description="conversion rate of the system. completes divided by total number of entrants to the system",
ge=0,
le=1,
@@ -166,8 +160,8 @@ class MorningTaskStatistics(MorningStatistics):
class MorningCondition(MarketplaceCondition):
model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore")
- question_id: Optional[MorningQuestionID] = Field(validation_alias="id")
- values: List[Annotated[str, Field(max_length=128)]] = Field(
+ question_id: MorningQuestionID | None = Field(validation_alias="id")
+ values: list[Annotated[str, Field(max_length=128)]] = Field(
validation_alias="response_ids"
)
value_type: ConditionValueType = Field(default=ConditionValueType.LIST)
@@ -184,11 +178,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
max_digits=5,
validation_alias="cost_per_interview",
)
- condition_hashes: List[str] = Field(min_length=1, default_factory=list)
+ condition_hashes: list[str] = Field(min_length=1, default_factory=list)
# since the Quota is the MarketplaceTask, it needs these fields, copied from the Bid
source: Literal[Source.MORNING_CONSULT] = Field(default=Source.MORNING_CONSULT)
- used_question_ids: Set[MorningQuestionID] = Field(default_factory=set)
+ used_question_ids: set[MorningQuestionID] = Field(default_factory=set)
country_iso: CountryISO = Field(frozen=True)
country_isos: CountryISOs = Field()
language_isos: LanguageISOs = Field(frozen=True)
@@ -206,7 +200,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
data["country_isos"] = [data["country_iso"]]
if isinstance(data["language_isos"], str):
data["language_isos"] = set(data["language_isos"].split(","))
- data["language_iso"] = sorted(data["language_isos"])[0]
+ data["language_iso"] = min(data["language_isos"])
return data
@property
@@ -219,11 +213,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
@computed_field
@cached_property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
return set(self.condition_hashes)
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return MorningCondition
@property
@@ -233,7 +227,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
@property
def marketplace_genders(
self,
- ) -> Dict[Gender, Optional[MarketplaceCondition]]:
+ ) -> dict[Gender, MarketplaceCondition | None]:
return {
Gender.MALE: MorningCondition(
question_id="gender",
@@ -253,14 +247,14 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
# num_available includes in-progress (they're already deducted)
return self.num_available >= self._min_open_spots
- def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Passes means we 1) meet all conditions (aka "match") AND 2) the quota is open.
return self.is_open and self.matches(criteria_evaluation)
# TODO: I did some speed tests. This is faster than how this is implemented
# in sago/spectrum/dynata/etc. We should generalize this logic instead of
# copying/pasting it 7 times. (matches, matches_optional and _soft)
- def matches(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def matches(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Matches means we meet all conditions.
# In Morning, all quotas are mutually exclusive. so if it doesn't
# matter if we match a closed quota, b/c that means that we won't
@@ -268,8 +262,8 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
return self.matches_optional(criteria_evaluation) is True
def matches_optional(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Optional[bool]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> bool | None:
for c in self.condition_hashes:
eval_value = criteria_evaluation.get(c)
if eval_value is False:
@@ -279,14 +273,14 @@ class MorningQuota(MorningStatistics, MarketplaceTask):
return True
def matches_soft(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], List[str]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, list[str]]:
# Passes back "matches" (T/F/none) and a list of unknown criterion hashes
- unknowns = list()
+ unknowns = []
for c in self.condition_hashes:
eval_value = criteria_evaluation.get(c)
if eval_value is False:
- return False, list()
+ return False, []
if eval_value is None:
unknowns.append(c)
if unknowns:
@@ -321,22 +315,22 @@ class MorningBid(MorningTaskStatistics):
timeout: PositiveInt = Field(le=24 * 60 * 60)
topic_id: str = Field(min_length=1, max_length=64)
- exclusions: List[MorningExclusion] = Field(default_factory=list)
+ exclusions: list[MorningExclusion] = Field(default_factory=list)
- quotas: List[MorningQuota] = Field(default_factory=list)
+ quotas: list[MorningQuota] = Field(default_factory=list)
source: Literal[Source.MORNING_CONSULT] = Field(default=Source.MORNING_CONSULT)
- used_question_ids: Set[MorningQuestionID] = Field(default_factory=set)
+ used_question_ids: set[MorningQuestionID] = Field(default_factory=set)
# This is a "special" key to store all conditions that are used (as
# "condition_hashes") throughout this survey. In the reduced representation
# of this task (nearly always, for db i/o, in global_vars) this field will
# be null.
- conditions: Optional[Dict[str, MorningCondition]] = Field(default=None)
+ conditions: dict[str, MorningCondition] | None = Field(default=None)
# This doesn't get stored in the db directly
- experimental_single_use_qualifications: Optional[List[MorningQuestion]] = Field(
+ experimental_single_use_qualifications: list[MorningQuestion] | None = Field(
default=None
)
@@ -345,8 +339,8 @@ class MorningBid(MorningTaskStatistics):
created_api: AwareDatetimeISO = Field(validation_alias="published_at")
# This does not come from the API. We set it when we update this in the db.
- created: Optional[AwareDatetimeISO] = Field(default=None)
- updated: Optional[AwareDatetimeISO] = Field(default=None)
+ created: AwareDatetimeISO | None = Field(default=None)
+ updated: AwareDatetimeISO | None = Field(default=None)
# ignoring from API: closed_at
@@ -365,7 +359,7 @@ class MorningBid(MorningTaskStatistics):
@property
def language_iso_any(self):
- return sorted(self.language_isos)[0]
+ return min(self.language_isos)
@property
def locale(self):
@@ -373,7 +367,7 @@ class MorningBid(MorningTaskStatistics):
@computed_field
@cached_property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
s = set()
for q in self.quotas:
s.update(set(q.condition_hashes))
@@ -387,7 +381,7 @@ class MorningBid(MorningTaskStatistics):
@model_validator(mode="before")
@classmethod
- def setup_quota_fields(cls, data: Dict[str, Any]) -> Dict[str, Any]:
+ def setup_quota_fields(cls, data: dict[str, Any]) -> dict[str, Any]:
# These fields get "inherited" by each quota from its bid.
quota_fields = [
"country_iso",
@@ -419,11 +413,11 @@ class MorningBid(MorningTaskStatistics):
@model_validator(mode="before")
@classmethod
- def setup_conditions(cls, data: Dict[str, Any]) -> Dict[str, Any]:
+ def setup_conditions(cls, data: dict[str, Any]) -> dict[str, Any]:
if "conditions" in data:
return data
- data["conditions"] = dict()
+ data["conditions"] = {}
for quota in data["quotas"]:
if "qualifications" in quota:
quota_conditions = [
@@ -448,7 +442,7 @@ class MorningBid(MorningTaskStatistics):
@model_validator(mode="before")
@classmethod
- def clean_alias(cls, data: Dict[str, Any]) -> Dict[str, Any]:
+ def clean_alias(cls, data: dict[str, Any]) -> dict[str, Any]:
# Make sure fields are named certain ways, so we don't have to check
# aliases within other validators
if "estimated_length_of_interview" in data:
@@ -503,18 +497,16 @@ class MorningBid(MorningTaskStatistics):
return d
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> Self:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
- d["expected_end"] = d["expected_end"].replace(tzinfo=timezone.utc)
- d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc)
+ def from_db(cls, d: dict[str, Any]) -> Self:
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
+ d["expected_end"] = d["expected_end"].replace(tzinfo=UTC)
+ d["created_api"] = d["created_api"].replace(tzinfo=UTC)
d["used_question_ids"] = json.loads(d["used_question_ids"])
d["exclusions"] = json.loads(d["exclusions"])
return cls.model_validate(d)
- def passes_quotas(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Optional[str]:
+ def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> str | None:
# Quotas are mutually-exclusive. A user can only possibly match 1 quota.
# Returns the passing quota ID or None (if user doesn't pass any quota)
for q in self.quotas:
@@ -522,8 +514,8 @@ class MorningBid(MorningTaskStatistics):
return q.id
def passes_quotas_soft(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], Optional[List[str]], Optional[Set[str]]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, list[str] | None, set[str] | None]:
"""
Quotas are mutually-exclusive. A user can only possibly match 1
quota. As such, all unknown questions on any quota will be
@@ -547,15 +539,15 @@ class MorningBid(MorningTaskStatistics):
return False, None, None
def determine_eligibility(
- self, criteria_evaluation: dict[str, Optional[bool]]
- ) -> Optional[str]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> str | None:
if not self.is_open:
return None
return self.passes_quotas(criteria_evaluation)
def determine_eligibility_soft(
- self, criteria_evaluation: dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], Optional[List[str]], Optional[Set[str]]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, list[str] | None, set[str] | None]:
if not self.is_open:
return False, None, None
return self.passes_quotas_soft(criteria_evaluation)
diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py
index eb4cbd1..9303a2f 100644
--- a/generalresearch/models/morning/task_collection.py
+++ b/generalresearch/models/morning/task_collection.py
@@ -108,7 +108,7 @@ class MorningTaskCollection(TaskCollection):
]
quota_fields = list(quota_columns.keys())
rows = []
- bid_dict = dict()
+ bid_dict = {}
for k in bid_fields:
bid_dict[k] = getattr(bid, k)
bid_dict["bid.id"] = bid.id
diff --git a/generalresearch/models/network/__init__.py b/generalresearch/models/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/generalresearch/models/network/__init__.py
+++ /dev/null
diff --git a/generalresearch/models/network/definitions.py b/generalresearch/models/network/definitions.py
deleted file mode 100644
index 2e1ab91..0000000
--- a/generalresearch/models/network/definitions.py
+++ /dev/null
@@ -1,70 +0,0 @@
-from __future__ import annotations
-
-from enum import StrEnum
-from ipaddress import ip_address, ip_network
-
-CGNAT_NET = ip_network("100.64.0.0/10")
-
-
-class IPProtocol(StrEnum):
- TCP = "tcp"
- UDP = "udp"
- SCTP = "sctp"
- IP = "ip"
- ICMP = "icmp"
- ICMPv6 = "icmpv6"
-
- def to_number(self) -> int:
- # https://www.iana.org/assignments/protocol-numbers/protocol-numbers.xhtml
- return {
- self.TCP: 6,
- self.UDP: 17,
- self.SCTP: 132,
- self.IP: 4,
- self.ICMP: 1,
- self.ICMPv6: 58,
- }[self]
-
-
-class IPKind(StrEnum):
- PUBLIC = "public"
- PRIVATE = "private"
- CGNAT = "carrier_nat"
- LOOPBACK = "loopback"
- LINK_LOCAL = "link_local"
- MULTICAST = "multicast"
- RESERVED = "reserved"
- UNSPECIFIED = "unspecified"
-
-
-def get_ip_kind(ip: str | None) -> IPKind | None:
- if not ip:
- return None
-
- ip_obj = ip_address(ip)
-
- if ip_obj in CGNAT_NET:
- return IPKind.CGNAT
-
- if ip_obj.is_loopback:
- return IPKind.LOOPBACK
-
- if ip_obj.is_link_local:
- return IPKind.LINK_LOCAL
-
- if ip_obj.is_multicast:
- return IPKind.MULTICAST
-
- if ip_obj.is_unspecified:
- return IPKind.UNSPECIFIED
-
- if ip_obj.is_private:
- return IPKind.PRIVATE
-
- if ip_obj.is_reserved:
- return IPKind.RESERVED
-
- if ip_obj.is_global:
- return IPKind.PUBLIC
-
- return None
diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py
deleted file mode 100644
index e4ddd18..0000000
--- a/generalresearch/models/network/label.py
+++ /dev/null
@@ -1,125 +0,0 @@
-from __future__ import annotations
-
-import ipaddress
-from enum import StrEnum
-
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- IPvAnyNetwork,
- computed_field,
- field_validator,
-)
-
-from generalresearch.models.custom_types import (
- AwareDatetimeISO,
- now_utc_factory,
-)
-
-
-class IPTrustClass(StrEnum):
- TRUSTED = "trusted"
- UNTRUSTED = "untrusted"
- # Note: use case of unknown is for e.g. Spur says this IP is a residential proxy
- # on 2026-1-1, and then has no annotation a month later. It doesn't mean
- # the IP is TRUSTED, but we want to record that Spur now doesn't claim UNTRUSTED.
- UNKNOWN = "unknown"
-
-
-class IPLabelKind(StrEnum):
- # --- UNTRUSTED ---
- RESIDENTIAL_PROXY = "residential_proxy"
- DATACENTER_PROXY = "datacenter_proxy"
- ISP_PROXY = "isp_proxy"
- MOBILE_PROXY = "mobile_proxy"
- PROXY = "proxy"
- HOSTING = "hosting"
- VPN = "vpn"
- RELAY = "relay"
- TOR_EXIT = "tor_exit"
- BAD_ACTOR = "bad_actor"
- # --- TRUSTED ---
- TRUSTED_USER = "trusted_user"
- # --- UNKNOWN ---
- UNKNOWN = "unknown"
-
-
-class IPLabelSource(StrEnum):
- # We got this IP from our own use of a proxy service
- INTERNAL_USE = "internal_use"
-
- # An external "security" service flagged this IP
- SPUR = "spur"
- IPINFO = "ipinfo"
- MAXMIND = "maxmind"
-
- MANUAL = "manual"
-
-
-class IPLabel(BaseModel):
- """
- Stores *ground truth* about an IP at a specific time.
- To be used for model training and evaluation.
- """
-
- model_config = ConfigDict(validate_assignment=True)
-
- ip: IPvAnyNetwork = Field()
-
- labeled_at: AwareDatetimeISO = Field(default_factory=now_utc_factory)
- created_at: AwareDatetimeISO | None = Field(default=None)
-
- label_kind: IPLabelKind = Field()
- source: IPLabelSource = Field()
-
- confidence: float = Field(default=1.0, ge=0.0, le=1.0)
-
- # Optionally, if this is untrusted, which service is providing the proxy/vpn service
- provider: str | None = Field(
- default=None, examples=["geonode", "gecko"], max_length=128
- )
-
- metadata: IPLabelMetadata | None = Field(default=None)
-
- @field_validator("ip", mode="before")
- @classmethod
- def normalize_and_validate_network(cls, v):
- net = ipaddress.ip_network(v, strict=False)
-
- if isinstance(net, ipaddress.IPv6Network):
- if net.prefixlen > 64:
- raise ValueError("IPv6 network must be /64 or larger")
-
- return net
-
- @field_validator("provider", mode="before")
- @classmethod
- def provider_format(cls, v: str | None) -> str | None:
- if v is None:
- return v
- return v.lower().strip()
-
- @computed_field()
- @property
- def trust_class(self) -> IPTrustClass:
- if self.label_kind == IPLabelKind.UNKNOWN:
- return IPTrustClass.UNKNOWN
- if self.label_kind == IPLabelKind.TRUSTED_USER:
- return IPTrustClass.TRUSTED
- return IPTrustClass.UNTRUSTED
-
- def model_dump_postgres(self):
- d = self.model_dump(mode="json")
- d["metadata"] = self.metadata.model_dump_json() if self.metadata else None
- return d
-
-
-class IPLabelMetadata(BaseModel):
- """
- To be expanded. Just for storing some things from Spur for now
- """
-
- model_config = ConfigDict(validate_assignment=True, extra="allow")
-
- services: list[str] | None = Field(min_length=1, examples=[["RDP"]])
diff --git a/generalresearch/models/network/mtr/__init__.py b/generalresearch/models/network/mtr/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/generalresearch/models/network/mtr/__init__.py
+++ /dev/null
diff --git a/generalresearch/models/network/mtr/command.py b/generalresearch/models/network/mtr/command.py
deleted file mode 100644
index fd5a8d0..0000000
--- a/generalresearch/models/network/mtr/command.py
+++ /dev/null
@@ -1,73 +0,0 @@
-from __future__ import annotations
-
-import subprocess
-
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.mtr.parser import parse_mtr_output
-from generalresearch.models.network.mtr.result import MTRResult
-from generalresearch.models.network.tool_run_command import MTRRunCommand
-
-SUPPORTED_PROTOCOLS = {
- IPProtocol.TCP,
- IPProtocol.UDP,
- IPProtocol.SCTP,
- IPProtocol.ICMP,
-}
-PROTOCOLS_W_PORT = {IPProtocol.TCP, IPProtocol.UDP, IPProtocol.SCTP}
-
-
-def build_mtr_command(
- ip: str,
- protocol: IPProtocol | None = None,
- port: int | None = None,
- report_cycles: int | None = 10,
-) -> str:
- # https://manpages.ubuntu.com/manpages/focal/man8/mtr.8.html
- # e.g. "mtr -r -c 2 -b -z -j -T -P 443 74.139.70.149"
- args = ["mtr", "--report", "--show-ips", "--aslookup", "--json"]
- if report_cycles is not None:
- args.extend(["-c", str(int(report_cycles))])
- if port is not None:
- if protocol is None:
- protocol = IPProtocol.TCP
- assert protocol in PROTOCOLS_W_PORT, "port only allowed for TCP/SCTP/UDP traces"
- args.extend(["--port", str(int(port))])
- if protocol:
- assert protocol in SUPPORTED_PROTOCOLS, f"unsupported protocol: {protocol}"
- # default is ICMP (no args)
- arg_map = {
- IPProtocol.TCP: "--tcp",
- IPProtocol.UDP: "--udp",
- IPProtocol.SCTP: "--sctp",
- }
- if protocol in arg_map:
- args.append(arg_map[protocol])
- args.append(ip)
- return " ".join(args)
-
-
-def get_mtr_version() -> str:
- proc = subprocess.run(
- ["mtr", "-v"],
- capture_output=True,
- text=True,
- check=False,
- )
- # e.g. mtr 0.95
- ver_str = proc.stdout.strip()
- return ver_str.split(" ", 1)[1]
-
-
-def run_mtr(config: MTRRunCommand) -> MTRResult:
- cmd = config.to_command_str()
- args = cmd.split(" ")
- proc = subprocess.run(
- args,
- capture_output=True,
- text=True,
- check=False,
- )
- raw = proc.stdout.strip()
- return parse_mtr_output(
- raw, protocol=config.options.protocol, port=config.options.port
- )
diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py
deleted file mode 100644
index d77e814..0000000
--- a/generalresearch/models/network/mtr/execute.py
+++ /dev/null
@@ -1,55 +0,0 @@
-from __future__ import annotations
-
-from datetime import datetime, timezone
-from uuid import uuid4
-
-from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.mtr.command import (
- get_mtr_version,
- run_mtr,
-)
-from generalresearch.models.network.tool_run import MTRRun, Status, ToolClass, ToolName
-from generalresearch.models.network.tool_run_command import (
- MTRRunCommand,
- MTRRunCommandOptions,
-)
-from generalresearch.models.network.utils import get_source_ip
-
-
-def execute_mtr(
- ip: str,
- scan_group_id: UUIDStr | None = None,
- protocol: IPProtocol | None = IPProtocol.ICMP,
- port: int | None = None,
- report_cycles: int = 10,
-) -> MTRRun:
- config = MTRRunCommand(
- options=MTRRunCommandOptions(
- ip=ip,
- report_cycles=report_cycles,
- protocol=protocol,
- port=port,
- ),
- )
-
- started_at = datetime.now(tz=timezone.utc)
- tool_version = get_mtr_version()
- result = run_mtr(config)
- finished_at = datetime.now(tz=timezone.utc)
-
- return MTRRun(
- tool_name=ToolName.MTR,
- tool_class=ToolClass.TRACEROUTE,
- tool_version=tool_version,
- status=Status.SUCCESS,
- ip=ip,
- started_at=started_at,
- finished_at=finished_at,
- raw_command=config.to_command_str(),
- scan_group_id=scan_group_id or uuid4().hex,
- config=config,
- parsed=result,
- source_ip=get_source_ip(),
- facility_id=1,
- )
diff --git a/generalresearch/models/network/mtr/parser.py b/generalresearch/models/network/mtr/parser.py
deleted file mode 100644
index c29439e..0000000
--- a/generalresearch/models/network/mtr/parser.py
+++ /dev/null
@@ -1,18 +0,0 @@
-import json
-
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.mtr.result import MTRResult
-
-
-def parse_mtr_output(raw: str, port: int, protocol: IPProtocol) -> MTRResult:
- data = parse_mtr_raw_output(raw)
- data["port"] = port
- data["protocol"] = protocol
- return MTRResult.model_validate(data)
-
-
-def parse_mtr_raw_output(raw: str) -> dict:
- data = json.loads(raw)["report"]
- data.update(data.pop("mtr"))
- data["hops"] = data.pop("hubs")
- return data
diff --git a/generalresearch/models/network/mtr/result.py b/generalresearch/models/network/mtr/result.py
deleted file mode 100644
index 34de845..0000000
--- a/generalresearch/models/network/mtr/result.py
+++ /dev/null
@@ -1,168 +0,0 @@
-from __future__ import annotations
-
-import re
-from functools import cached_property
-from ipaddress import ip_address
-
-import tldextract
-from pydantic import (
- BaseModel,
- ConfigDict,
- Field,
- computed_field,
- field_validator,
- model_validator,
-)
-
-from generalresearch.models.network.definitions import IPKind, IPProtocol, get_ip_kind
-
-HOST_RE = re.compile(r"^(?P<hostname>.+?) \((?P<ip>[^)]+)\)$")
-
-
-class MTRHop(BaseModel):
- model_config = ConfigDict(populate_by_name=True)
-
- hop: int = Field(alias="count")
- host: str
- asn: int | None = Field(default=None, alias="ASN")
-
- loss_pct: float = Field(alias="Loss%")
- sent: int = Field(alias="Snt")
-
- last_ms: float = Field(alias="Last")
- avg_ms: float = Field(alias="Avg")
- best_ms: float = Field(alias="Best")
- worst_ms: float = Field(alias="Wrst")
- stdev_ms: float = Field(alias="StDev")
-
- hostname: str | None = Field(
- default=None, examples=["fixed-187-191-8-145.totalplay.net"]
- )
- ip: str | None = None
-
- @field_validator("asn", mode="before")
- @classmethod
- def normalize_asn(cls, v: str):
- if v is None or v == "AS???":
- return None
- if type(v) is int:
- return v
- return int(v.replace("AS", ""))
-
- @model_validator(mode="after")
- def parse_host(self):
- host = self.host.strip()
-
- # hostname (ip)
- m = HOST_RE.match(host)
- if m:
- self.hostname = m.group("hostname")
- self.ip = m.group("ip")
- return self
-
- # ip only
- try:
- ip_address(host)
- self.ip = host
- self.hostname = None
- return self
- except ValueError:
- pass
-
- # hostname only
- self.hostname = host
- self.ip = None
- return self
-
- @cached_property
- def ip_kind(self) -> IPKind | None:
- return get_ip_kind(self.ip)
-
- @cached_property
- def icmp_rate_limited(self):
- if self.avg_ms == 0:
- return False
- return self.stdev_ms > self.avg_ms or self.worst_ms > self.best_ms * 10
-
- @computed_field(examples=["totalplay.net"])
- @cached_property
- def domain(self) -> str | None:
- if self.hostname:
- return tldextract.extract(self.hostname).top_domain_under_public_suffix
-
- def model_dump_postgres(self, run_id: int):
- # Writes for the network_mtrhop table
- d = {"mtr_run_id": run_id}
- data = self.model_dump(
- mode="json",
- include={
- "hop",
- "ip",
- "domain",
- "asn",
- },
- )
- d.update(data)
- return d
-
-
-class MTRResult(BaseModel):
- model_config = ConfigDict(populate_by_name=True)
-
- source: str = Field(description="Hostname of the system running mtr.", alias="src")
- destination: str = Field(
- description="Destination hostname or IP being traced.", alias="dst"
- )
- tos: int = Field(description="IP Type-of-Service (TOS) value used for probes.")
- tests: int = Field(description="Number of probes sent per hop.")
- psize: int = Field(description="Probe packet size in bytes.")
- bitpattern: str = Field(description="Payload byte pattern used in probes (hex).")
-
- # Protocol used for the traceroute
- protocol: IPProtocol = Field(default=IPProtocol.ICMP)
- # The target port number for TCP/SCTP/UDP traces
- port: int | None = Field(default=None)
-
- hops: list[MTRHop] = Field()
-
- def model_dump_postgres(self):
- # Writes for the network_mtr table
- d = self.model_dump(
- mode="json",
- include={"port"},
- )
- d["protocol"] = self.protocol.to_number()
- d["parsed"] = self.model_dump_json(indent=0)
- return d
-
- def print_report(self) -> None:
- print(
- f"MTR Report → {self.destination} {self.protocol.name} {self.port or ''}\n"
- )
- host_max_len = max(len(h.host) for h in self.hops)
-
- header = (
- f"{'Hop':>3} "
- f"{'Host':<{host_max_len}} "
- f"{'Kind':<10} "
- f"{'ASN':<8} "
- f"{'Loss%':>6} {'Sent':>5} "
- f"{'Last':>7} {'Avg':>7} {'Best':>7} {'Worst':>7} {'StDev':>7}"
- )
- print(header)
- print("-" * len(header))
-
- for hop in self.hops:
- print(
- f"{hop.hop:>3} "
- f"{hop.host:<{host_max_len}} "
- f"{hop.ip_kind or '???':<10} "
- f"{hop.asn or '???':<8} "
- f"{hop.loss_pct:6.1f} "
- f"{hop.sent:5d} "
- f"{hop.last_ms:7.1f} "
- f"{hop.avg_ms:7.1f} "
- f"{hop.best_ms:7.1f} "
- f"{hop.worst_ms:7.1f} "
- f"{hop.stdev_ms:7.1f}"
- )
diff --git a/generalresearch/models/network/nmap/__init__.py b/generalresearch/models/network/nmap/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/generalresearch/models/network/nmap/__init__.py
+++ /dev/null
diff --git a/generalresearch/models/network/nmap/command.py b/generalresearch/models/network/nmap/command.py
deleted file mode 100644
index 6a524a8..0000000
--- a/generalresearch/models/network/nmap/command.py
+++ /dev/null
@@ -1,48 +0,0 @@
-from __future__ import annotations
-
-import subprocess
-
-from generalresearch.models.network.nmap.parser import parse_nmap_xml
-from generalresearch.models.network.nmap.result import NmapResult
-from generalresearch.models.network.tool_run_command import NmapRunCommand
-
-
-def build_nmap_command(
- ip: str,
- no_ping: bool = True,
- enable_advanced: bool = True,
- timing: int = 4,
- ports: str | None = None,
- top_ports: int | None = None,
-) -> str:
- # e.g. "nmap -Pn -T4 -A --top-ports 1000 -oX - scanme.nmap.org"
- # https://linux.die.net/man/1/nmap
- args = ["nmap"]
- assert 0 <= timing <= 5
- args.append(f"-T{timing}")
- if no_ping:
- args.append("-Pn")
- if enable_advanced:
- args.append("-A")
- if ports is not None:
- assert top_ports is None
- args.extend(["-p", ports])
- if top_ports is not None:
- assert ports is None
- args.extend(["--top-ports", str(top_ports)])
-
- args.extend(["-oX", "-", ip])
- return " ".join(args)
-
-
-def run_nmap(config: NmapRunCommand) -> NmapResult:
- cmd = config.to_command_str()
- args = cmd.split(" ")
- proc = subprocess.run(
- args,
- capture_output=True,
- text=True,
- check=False,
- )
- raw = proc.stdout.strip()
- return parse_nmap_xml(raw)
diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py
deleted file mode 100644
index 8a73307..0000000
--- a/generalresearch/models/network/nmap/execute.py
+++ /dev/null
@@ -1,51 +0,0 @@
-from __future__ import annotations
-
-from uuid import uuid4
-
-from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.network.nmap.command import run_nmap
-from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName
-from generalresearch.models.network.tool_run_command import (
- NmapRunCommand,
- NmapRunCommandOptions,
-)
-
-
-def execute_nmap(
- ip: str,
- top_ports: int | None = 1000,
- ports: str | None = None,
- no_ping: bool = True,
- enable_advanced: bool = True,
- timing: int = 4,
- scan_group_id: UUIDStr | None = None,
-):
- config = NmapRunCommand(
- options=NmapRunCommandOptions(
- top_ports=top_ports,
- ports=ports,
- no_ping=no_ping,
- enable_advanced=enable_advanced,
- timing=timing,
- ip=ip,
- )
- )
- result = run_nmap(config)
- assert result.exit_status == "success"
- assert result.target_ip == ip, f"{result.target_ip=}, {ip=}"
- assert result.command_line == config.to_command_str()
-
- run = NmapRun(
- tool_name=ToolName.NMAP,
- tool_class=ToolClass.PORT_SCAN,
- tool_version=result.version,
- status=Status.SUCCESS,
- ip=ip,
- started_at=result.started_at,
- finished_at=result.finished_at,
- raw_command=result.command_line,
- scan_group_id=scan_group_id or uuid4().hex,
- config=config,
- parsed=result,
- )
- return run
diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py
deleted file mode 100644
index e946e5f..0000000
--- a/generalresearch/models/network/nmap/parser.py
+++ /dev/null
@@ -1,414 +0,0 @@
-from __future__ import annotations
-
-import xml.etree.ElementTree as ET
-from datetime import datetime, timezone
-from typing import Any
-
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.nmap.result import (
- NmapHostname,
- NmapHostScript,
- NmapHostState,
- NmapHostStatusReason,
- NmapOSClass,
- NmapOSMatch,
- NmapPort,
- NmapPortStats,
- NmapResult,
- NmapScanInfo,
- NmapScanType,
- NmapScript,
- NmapService,
- NmapTrace,
- NmapTraceHop,
- PortState,
- PortStateReason,
-)
-
-
-class NmapParserException(Exception):
- def __init__(self, msg):
- self.msg = msg
-
- def __str__(self):
- return self.msg
-
-
-class NmapXmlParser:
- """
- Example: https://nmap.org/book/output-formats-xml-output.html
- Full DTD: https://nmap.org/book/nmap-dtd.html
- """
-
- @classmethod
- def parse_xml(cls, nmap_data: str) -> NmapResult:
- """
- Expects a full nmap scan report.
- """
-
- try:
- root = ET.fromstring(nmap_data)
- except Exception as e:
- emsg = f"Wrong XML structure: cannot parse data: {e}"
- raise NmapParserException(emsg)
-
- if root.tag != "nmaprun":
- raise NmapParserException("Unpexpected data structure for XML " "root node")
- return cls._parse_xml_nmaprun(root)
-
- @classmethod
- def _parse_xml_nmaprun(cls, root: ET.Element) -> NmapResult:
- """
- This method parses out a full nmap scan report from its XML root
- node: <nmaprun>. We expect there is only 1 host in this report!
-
- :param root: Element from xml.ElementTree (top of XML the document)
- """
- cls._validate_nmap_root(root)
- host_count = len(root.findall(".//host"))
- assert host_count == 1, f"Expected 1 host, got {host_count}"
-
- xml_str = ET.tostring(root, encoding="unicode").replace("\n", "")
- nmap_data = {"raw_xml": xml_str}
- nmap_data.update(cls._parse_nmaprun(root))
-
- nmap_data["scan_infos"] = [
- cls._parse_scaninfo(scaninfo_el)
- for scaninfo_el in root.findall(".//scaninfo")
- ]
-
- nmap_data.update(cls._parse_runstats(root))
-
- nmap_data.update(cls._parse_xml_host(root.find(".//host")))
-
- return NmapResult.model_validate(nmap_data)
-
- @classmethod
- def _validate_nmap_root(cls, root: ET.Element) -> None:
- allowed = {
- "scaninfo",
- "host",
- "runstats",
- "verbose",
- "debugging",
- "taskprogress",
- }
-
- found = {child.tag for child in root}
- unexpected = found - allowed
- if unexpected:
- raise ValueError(
- f"Unexpected top-level tags in nmap XML: {sorted(unexpected)}"
- )
-
- @classmethod
- def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo:
- data = dict()
- data["type"] = NmapScanType(scaninfo_el.attrib["type"])
- data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"])
- data["num_services"] = scaninfo_el.attrib["numservices"]
- data["services"] = scaninfo_el.attrib["services"]
- return NmapScanInfo.model_validate(data)
-
- @classmethod
- def _parse_runstats(cls, root: ET.Element) -> dict:
- runstats = root.find("runstats")
- if runstats is None:
- return {}
-
- finished = runstats.find("finished")
- if finished is None:
- return {}
-
- finished_at = None
- ts = finished.attrib.get("time")
- if ts:
- finished_at = datetime.fromtimestamp(int(ts), tz=timezone.utc)
-
- return {
- "finished_at": finished_at,
- "exit_status": finished.attrib.get("exit"),
- }
-
- @classmethod
- def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict:
- nmap_data = dict()
- nmaprun = dict(nmaprun_el.attrib)
- nmap_data["command_line"] = nmaprun["args"]
- nmap_data["started_at"] = datetime.fromtimestamp(
- float(nmaprun["start"]), tz=timezone.utc
- )
- nmap_data["version"] = nmaprun["version"]
- nmap_data["xmloutputversion"] = nmaprun["xmloutputversion"]
- return nmap_data
-
- @classmethod
- def _parse_xml_host(cls, host_el: ET.Element) -> dict:
- """
- Receives a <host> XML tag representing a scanned host with
- its services.
- """
- data = dict()
-
- # <status state="up" reason="user-set" reason_ttl="0"/>
- status_el = host_el.find("status")
- data["host_state"] = NmapHostState(status_el.attrib["state"])
- data["host_state_reason"] = NmapHostStatusReason(status_el.attrib["reason"])
- host_state_reason_ttl = status_el.attrib.get("reason_ttl")
- if host_state_reason_ttl:
- data["host_state_reason_ttl"] = int(host_state_reason_ttl)
-
- # <address addr="108.171.53.1" addrtype="ipv4"/>
- address_el = host_el.find("address")
- data["target_ip"] = address_el.attrib["addr"]
-
- data["hostnames"] = cls._parse_hostnames(host_el.find("hostnames"))
-
- data["ports"], data["port_stats"] = cls._parse_xml_ports(host_el.find("ports"))
-
- uptime = host_el.find("uptime")
- if uptime is not None:
- data["uptime_seconds"] = int(uptime.attrib["seconds"])
-
- distance = host_el.find("distance")
- if distance is not None:
- data["distance"] = int(distance.attrib["value"])
-
- tcpsequence = host_el.find("tcpsequence")
- if tcpsequence is not None:
- data["tcp_sequence_index"] = int(tcpsequence.attrib["index"])
- data["tcp_sequence_difficulty"] = tcpsequence.attrib["difficulty"]
- ipidsequence = host_el.find("ipidsequence")
- if ipidsequence is not None:
- data["ipid_sequence_class"] = ipidsequence.attrib["class"]
- tcptssequence = host_el.find("tcptssequence")
- if tcptssequence is not None:
- data["tcp_timestamp_class"] = tcptssequence.attrib["class"]
-
- times_elem = host_el.find("times")
- if times_elem is not None:
- data.update(
- {
- "srtt_us": int(times_elem.attrib.get("srtt", 0)) or None,
- "rttvar_us": int(times_elem.attrib.get("rttvar", 0)) or None,
- "timeout_us": int(times_elem.attrib.get("to", 0)) or None,
- }
- )
-
- hostscripts_el = host_el.find("hostscript")
- if hostscripts_el is not None:
- data["host_scripts"] = [
- NmapHostScript(id=el.attrib["id"], output=el.attrib.get("output"))
- for el in hostscripts_el.findall("script")
- ]
-
- data["os_matches"] = cls._parse_os_matches(host_el)
-
- data["trace"] = cls._parse_trace(host_el)
-
- return data
-
- @classmethod
- def _parse_os_matches(cls, host_el: ET.Element) -> list[NmapOSMatch] | None:
- os_elem = host_el.find("os")
- if os_elem is None:
- return None
-
- matches: list[NmapOSMatch] = []
-
- for m in os_elem.findall("osmatch"):
- classes: list[NmapOSClass] = []
-
- for c in m.findall("osclass"):
- cpes = [e.text.strip() for e in c.findall("cpe") if e.text]
-
- classes.append(
- NmapOSClass(
- vendor=c.attrib.get("vendor"),
- osfamily=c.attrib.get("osfamily"),
- osgen=c.attrib.get("osgen"),
- accuracy=(
- int(c.attrib["accuracy"])
- if "accuracy" in c.attrib
- else None
- ),
- cpe=cpes or None,
- )
- )
-
- matches.append(
- NmapOSMatch(
- name=m.attrib["name"],
- accuracy=int(m.attrib["accuracy"]),
- classes=classes,
- )
- )
-
- return matches or None
-
- @classmethod
- def _parse_hostnames(cls, hostnames_el: ET.Element) -> list[NmapHostname]:
- """
- Parses the hostnames element.
- e.g. <hostnames>
- <hostname name="108-171-53-1.aceips.com" type="PTR"/>
- </hostnames>
- """
- return [
- cls._parse_hostname(hname) for hname in hostnames_el.findall("hostname")
- ]
-
- @classmethod
- def _parse_hostname(cls, hostname_el: ET.Element) -> NmapHostname:
- """
- Parses the hostname element.
- e.g. <hostname name="108-171-53-1.aceips.com" type="PTR"/>
-
- :param hostname_el: <hostname> XML tag from a nmap scan
- """
- return NmapHostname.model_validate(dict(hostname_el.attrib))
-
- @classmethod
- def _parse_xml_ports(
- cls, ports_elem: ET.Element
- ) -> tuple[list[NmapPort], NmapPortStats]:
- """
- Parses the list of scanned services from a targeted host.
- """
- ports: list[NmapPort] = []
- stats = NmapPortStats()
-
- # handle extraports first
- for e in ports_elem.findall("extraports"):
- state = PortState(e.attrib["state"])
- count = int(e.attrib["count"])
-
- key = state.value.replace("|", "_")
- setattr(stats, key, getattr(stats, key) + count)
-
- for port_elem in ports_elem.findall("port"):
- port = cls._parse_xml_port(port_elem)
- ports.append(port)
- key = port.state.value.replace("|", "_")
- setattr(stats, key, getattr(stats, key) + 1)
- return ports, stats
-
- @classmethod
- def _parse_xml_service(cls, service_elem: ET.Element) -> NmapService:
- svc = {
- "name": service_elem.attrib.get("name"),
- "product": service_elem.attrib.get("product"),
- "version": service_elem.attrib.get("version"),
- "extrainfo": service_elem.attrib.get("extrainfo"),
- "method": service_elem.attrib.get("method"),
- "conf": (
- int(service_elem.attrib["conf"])
- if "conf" in service_elem.attrib
- else None
- ),
- "cpe": [e.text.strip() for e in service_elem.findall("cpe")],
- }
-
- return NmapService.model_validate(svc)
-
- @classmethod
- def _parse_xml_script(cls, script_elem: ET.Element) -> NmapScript:
- output = script_elem.attrib.get("output")
- if output:
- output = output.strip()
- script = {
- "id": script_elem.attrib["id"],
- "output": output,
- }
-
- elements: dict[str, Any] = {}
-
- # handle <elem key="...">value</elem>
- for elem in script_elem.findall(".//elem"):
- key = elem.attrib.get("key")
- if key:
- elements[key.strip()] = elem.text.strip()
-
- script["elements"] = elements
- return NmapScript.model_validate(script)
-
- @classmethod
- def _parse_xml_port(cls, port_elem: ET.Element) -> NmapPort:
- """
- <port protocol="tcp" portid="61232">
- <state state="open" reason="syn-ack" reason_ttl="47"/>
- <service name="socks5" extrainfo="Username/password authentication required" method="probed" conf="10"/>
- <script id="socks-auth-info" output="&#xa; Username and password">
- <table>
- <elem key="name">Username and password</elem>
- <elem key="method">2</elem>
- </table>
- </script>
- </port>
- """
- state_elem = port_elem.find("state")
-
- port = {
- "port": int(port_elem.attrib["portid"]),
- "protocol": port_elem.attrib["protocol"],
- "state": PortState(state_elem.attrib["state"]),
- "reason": (
- PortStateReason(state_elem.attrib["reason"])
- if "reason" in state_elem.attrib
- else None
- ),
- "reason_ttl": (
- int(state_elem.attrib["reason_ttl"])
- if "reason_ttl" in state_elem.attrib
- else None
- ),
- }
-
- service_elem = port_elem.find("service")
- if service_elem is not None:
- port["service"] = cls._parse_xml_service(service_elem)
-
- port["scripts"] = []
- for script_elem in port_elem.findall("script"):
- port["scripts"].append(cls._parse_xml_script(script_elem))
-
- return NmapPort.model_validate(port)
-
- @classmethod
- def _parse_trace(cls, host_elem: ET.Element) -> NmapTrace | None:
- trace_elem = host_elem.find("trace")
- if trace_elem is None:
- return None
-
- port_attr = trace_elem.attrib.get("port")
- proto_attr = trace_elem.attrib.get("proto")
-
- hops: list[NmapTraceHop] = []
-
- for hop_elem in trace_elem.findall("hop"):
- ttl = hop_elem.attrib.get("ttl")
- if ttl is None:
- continue # ttl is required by the DTD but guard anyway
-
- rtt = hop_elem.attrib.get("rtt")
- ipaddr = hop_elem.attrib.get("ipaddr")
- host = hop_elem.attrib.get("host")
-
- hops.append(
- NmapTraceHop(
- ttl=int(ttl),
- ipaddr=ipaddr,
- rtt_ms=float(rtt) if rtt is not None else None,
- host=host,
- )
- )
-
- return NmapTrace(
- port=int(port_attr) if port_attr is not None else None,
- protocol=IPProtocol(proto_attr) if proto_attr is not None else None,
- hops=hops,
- )
-
-
-def parse_nmap_xml(raw) -> NmapResult:
- return NmapXmlParser.parse_xml(raw)
diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py
deleted file mode 100644
index 3f9cae6..0000000
--- a/generalresearch/models/network/nmap/result.py
+++ /dev/null
@@ -1,434 +0,0 @@
-from __future__ import annotations
-
-import json
-from datetime import timedelta
-from enum import StrEnum
-from functools import cached_property
-from typing import Any, Literal, Set
-
-from pydantic import BaseModel, Field, computed_field
-
-from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr
-from generalresearch.models.network.definitions import IPProtocol
-
-
-class PortState(StrEnum):
- OPEN = "open"
- CLOSED = "closed"
- FILTERED = "filtered"
- UNFILTERED = "unfiltered"
- OPEN_FILTERED = "open|filtered"
- CLOSED_FILTERED = "closed|filtered"
- # Added by me, does not get returned. Used for book-keeping
- NOT_SCANNED = "not_scanned"
-
-
-class PortStateReason(StrEnum):
- SYN_ACK = "syn-ack"
- RESET = "reset"
- CONN_REFUSED = "conn-refused"
- NO_RESPONSE = "no-response"
- SYN = "syn"
- FIN = "fin"
-
- ICMP_NET_UNREACH = "net-unreach"
- ICMP_HOST_UNREACH = "host-unreach"
- ICMP_PROTO_UNREACH = "proto-unreach"
- ICMP_PORT_UNREACH = "port-unreach"
-
- ADMIN_PROHIBITED = "admin-prohibited"
- HOST_PROHIBITED = "host-prohibited"
- NET_PROHIBITED = "net-prohibited"
-
- ECHO_REPLY = "echo-reply"
- TIME_EXCEEDED = "time-exceeded"
-
-
-class NmapScanType(StrEnum):
- SYN = "syn"
- CONNECT = "connect"
- ACK = "ack"
- WINDOW = "window"
- MAIMON = "maimon"
- FIN = "fin"
- NULL = "null"
- XMAS = "xmas"
- UDP = "udp"
- SCTP_INIT = "sctpinit"
- SCTP_COOKIE_ECHO = "sctpcookieecho"
-
-
-class NmapHostState(StrEnum):
- UP = "up"
- DOWN = "down"
- UNKNOWN = "unknown"
-
-
-class NmapHostStatusReason(StrEnum):
- USER_SET = "user-set"
- SYN_ACK = "syn-ack"
- RESET = "reset"
- ECHO_REPLY = "echo-reply"
- ARP_RESPONSE = "arp-response"
- NO_RESPONSE = "no-response"
- NET_UNREACH = "net-unreach"
- HOST_UNREACH = "host-unreach"
- PROTO_UNREACH = "proto-unreach"
- PORT_UNREACH = "port-unreach"
- ADMIN_PROHIBITED = "admin-prohibited"
- LOCALHOST_RESPONSE = "localhost-response"
-
-
-class NmapOSClass(BaseModel):
- vendor: str = None
- osfamily: str = None
- osgen: str | None = None
- accuracy: int = None
- cpe: list[str] | None = None
-
-
-class NmapOSMatch(BaseModel):
- name: str
- accuracy: int
- classes: list[NmapOSClass] = Field(default_factory=list)
-
- @property
- def best_class(self) -> NmapOSClass | None:
- if not self.classes:
- return None
- return max(self.classes, key=lambda m: m.accuracy)
-
-
-class NmapScript(BaseModel):
- """
- <script id="socks-auth-info" output="&#xa; Username and password">
- <table>
- <elem key="name">Username and password</elem>
- <elem key="method">2</elem>
- </table>
- </script>
- """
-
- id: str
- output: str | None = None
- elements: dict[str, Any] = Field(default_factory=dict)
-
-
-class NmapService(BaseModel):
- # <service name="socks5" extrainfo="Username/password authentication required" method="probed" conf="10"/>
- name: str | None = None
- product: str | None = None
- version: str | None = None
- extrainfo: str | None = None
- method: str | None = None
- conf: int | None = None
- cpe: list[str] = Field(default_factory=list)
-
- def model_dump_postgres(self):
- d = self.model_dump(mode="json")
- d["service_name"] = self.name
- return d
-
-
-class NmapPort(BaseModel):
- port: int = Field()
- protocol: IPProtocol = Field()
- # Closed ports will not have a NmapPort record
- state: PortState = Field()
- reason: PortStateReason | None = Field(default=None)
- reason_ttl: int | None = Field(default=None)
-
- service: NmapService | None = None
- scripts: list[NmapScript] = Field(default_factory=list)
-
- def model_dump_postgres(self, run_id: int):
- # Writes for the network_portscanport table
- d = {"port_scan_id": run_id}
- data = self.model_dump(
- mode="json",
- include={
- "port",
- "state",
- "reason",
- "reason_ttl",
- },
- )
- d.update(data)
- d["protocol"] = self.protocol.to_number()
- if self.service:
- d.update(self.service.model_dump_postgres())
- return d
-
-
-class NmapHostScript(BaseModel):
- id: str = Field()
- output: str | None = Field(default=None)
-
-
-class NmapTraceHop(BaseModel):
- """
- One hop observed during Nmap's traceroute.
-
- Example XML:
- <hop ttl="7" ipaddr="62.115.192.20" rtt="17.17" host="gdl-b2-link.ip.twelve99.net"/>
- """
-
- ttl: int = Field()
-
- ipaddr: str | None = Field(
- default=None,
- description="IP address of the responding router or host",
- )
-
- rtt_ms: float | None = Field(
- default=None,
- description="Round-trip time in milliseconds for the probe reaching this hop.",
- )
-
- host: str | None = Field(
- default=None,
- description="Reverse DNS hostname for the hop if Nmap resolved one.",
- )
-
-
-class NmapTrace(BaseModel):
- """
- Traceroute information collected by Nmap.
-
- Nmap performs a single traceroute per host using probes matching the scan
- type (typically TCP) directed at a chosen destination port.
-
- Example XML:
- <trace port="61232" proto="tcp">
- <hop ttl="1" ipaddr="192.168.86.1" rtt="3.83"/>
- ...
- </trace>
- """
-
- port: int | None = Field(
- default=None,
- description="Destination port used for traceroute probes (may be absent depending on scan type).",
- )
- protocol: IPProtocol | None = Field(
- default=None,
- description="Transport protocol used for the traceroute probes (tcp, udp, etc.).",
- )
-
- hops: list[NmapTraceHop] = Field(
- default_factory=list,
- description="Ordered list of hops observed during the traceroute.",
- )
-
- @property
- def destination(self) -> NmapTraceHop | None:
- return self.hops[-1] if self.hops else None
-
-
-class NmapHostname(BaseModel):
- # <hostname name="108-171-53-1.aceips.com" type="PTR"/>
- name: str
- type: Literal["PTR", "user"] | None = None
-
-
-class NmapPortStats(BaseModel):
- """
- This is counts across all protocols scanned (tcp/udp)
- """
-
- open: int = 0
- closed: int = 0
- filtered: int = 0
- unfiltered: int = 0
- open_filtered: int = 0
- closed_filtered: int = 0
-
-
-class NmapScanInfo(BaseModel):
- """
- We could have multiple protocols in one run.
- <scaninfo type="syn" protocol="tcp" numservices="983" services="22-1000,1100,3389,11000,61232"/>
- <scaninfo type="syn" protocol="udp" numservices="983" services="1100"/>
- """
-
- type: NmapScanType = Field()
- protocol: IPProtocol = Field()
- num_services: int = Field()
- services: str = Field()
-
- @cached_property
- def port_set(self) -> Set[int]:
- """
- Expand the Nmap services string into a set of port numbers.
- Example:
- "22-25,80,443" -> {22,23,24,25,80,443}
- """
- ports: Set[int] = set()
- for part in self.services.split(","):
- if "-" in part:
- start, end = part.split("-", 1)
- ports.update(range(int(start), int(end) + 1))
- else:
- ports.add(int(part))
- return ports
-
-
-class NmapResult(BaseModel):
- """
- A Nmap Run. Expects that we've only scanned ONE host.
- """
-
- command_line: str = Field()
- started_at: AwareDatetimeISO = Field()
- version: str = Field()
- xmloutputversion: str = Field()
-
- scan_infos: list[NmapScanInfo] = Field(min_length=1)
-
- # comes from <runstats>
- finished_at: AwareDatetimeISO | None = Field(default=None)
- exit_status: Literal["success", "error"] | None = Field(default=None)
-
- #####
- # Everything below here is from within the *single* host we've scanned
- #####
-
- # <status state="up" reason="user-set" reason_ttl="0"/>
- host_state: NmapHostState = Field()
- host_state_reason: NmapHostStatusReason = Field()
- host_state_reason_ttl: int | None = None
-
- # <address addr="108.171.53.1" addrtype="ipv4"/>
- target_ip: IPvAnyAddressStr = Field()
-
- hostnames: list[NmapHostname] = Field()
-
- ports: list[NmapPort] = []
- port_stats: NmapPortStats = Field()
-
- # <uptime seconds="4063775" lastboot="Fri Jan 16 12:12:06 2026"/>
- uptime_seconds: int | None = Field(default=None)
- # <distance value="11"/>
- distance: int | None = Field(description="approx number of hops", default=None)
-
- # <tcpsequence index="263" difficulty="Good luck!">
- tcp_sequence_index: int | None = None
- tcp_sequence_difficulty: str | None = None
-
- # <ipidsequence class="All zeros">
- ipid_sequence_class: str | None = None
-
- # <tcptssequence class="1000HZ" >
- tcp_timestamp_class: str | None = None
-
- # <times srtt="54719" rttvar="23423" to="148411"/>
- srtt_us: int | None = Field(
- default=None, description="smoothed RTT estimate (microseconds µs)"
- )
- rttvar_us: int | None = Field(
- default=None, description="RTT variance (microseconds µs)"
- )
- timeout_us: int | None = Field(
- default=None, description="probe timeout (microseconds µs)"
- )
-
- os_matches: list[NmapOSMatch] | None = Field(default=None)
-
- host_scripts: list[NmapHostScript] = Field(default_factory=list)
-
- trace: NmapTrace | None = Field(default=None)
-
- raw_xml: str | None = None
-
- @computed_field
- @property
- def last_boot(self) -> AwareDatetimeISO | None:
- if self.uptime_seconds:
- return self.started_at - timedelta(seconds=self.uptime_seconds)
-
- @property
- def scan_info_tcp(self):
- return next(
- filter(lambda x: x.protocol == IPProtocol.TCP, self.scan_infos), None
- )
-
- @property
- def scan_info_udp(self):
- return next(
- filter(lambda x: x.protocol == IPProtocol.UDP, self.scan_infos), None
- )
-
- @property
- def latency_ms(self) -> float | None:
- return self.srtt_us / 1000 if self.srtt_us is not None else None
-
- @property
- def best_os_match(self) -> NmapOSMatch | None:
- if not self.os_matches:
- return None
- return max(self.os_matches, key=lambda m: m.accuracy)
-
- def filter_ports(self, protocol: IPProtocol, state: PortState) -> list[NmapPort]:
- return [p for p in self.ports if p.protocol == protocol and p.state == state]
-
- @property
- def tcp_open_ports(self) -> list[int]:
- """
- Returns a list of open TCP port numbers.
- """
- return [
- p.port
- for p in self.filter_ports(protocol=IPProtocol.TCP, state=PortState.OPEN)
- ]
-
- @property
- def udp_open_ports(self) -> list[int]:
- """
- Returns a list of open UDP port numbers.
- """
- return [
- p.port
- for p in self.filter_ports(protocol=IPProtocol.UDP, state=PortState.OPEN)
- ]
-
- @cached_property
- def _port_index(self) -> dict[tuple[IPProtocol, int], NmapPort]:
- return {(p.protocol, p.port): p for p in self.ports}
-
- def get_port_state(
- self, port: int, protocol: IPProtocol = IPProtocol.TCP
- ) -> PortState:
- # Explicit (only if scanned and not closed)
- if (protocol, port) in self._port_index:
- return self._port_index[(protocol, port)].state
-
- # Check if we even scanned it
- scaninfo = next((s for s in self.scan_infos if s.protocol == protocol), None)
- if scaninfo and port in scaninfo.port_set:
- return PortState.CLOSED
-
- # We didn't scan it
- return PortState.NOT_SCANNED
-
- def model_dump_postgres(self):
- # Writes for the network_portscan table
- d = dict()
- data = self.model_dump(
- mode="json",
- include={
- "started_at",
- "host_state",
- "host_state_reason",
- "distance",
- "uptime_seconds",
- "raw_xml",
- },
- )
- d.update(data)
- d["ip"] = self.target_ip
- d["xml_version"] = self.xmloutputversion
- d["latency_ms"] = self.latency_ms
- d["last_boot"] = self.last_boot
- d["parsed"] = self.model_dump_json(indent=0)
- d["open_tcp_ports"] = json.dumps(self.tcp_open_ports)
- d["open_udp_ports"] = json.dumps(self.udp_open_ports)
- return d
diff --git a/generalresearch/models/network/rdns/__init__.py b/generalresearch/models/network/rdns/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/generalresearch/models/network/rdns/__init__.py
+++ /dev/null
diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py
deleted file mode 100644
index e88a84d..0000000
--- a/generalresearch/models/network/rdns/command.py
+++ /dev/null
@@ -1,35 +0,0 @@
-import subprocess
-
-from generalresearch.models.network.rdns.parser import parse_rdns_output
-from generalresearch.models.network.rdns.result import RDNSResult
-from generalresearch.models.network.tool_run_command import RDNSRunCommand
-
-
-def run_rdns(config: RDNSRunCommand) -> RDNSResult:
- cmd = config.to_command_str()
- args = cmd.split(" ")
- proc = subprocess.run(
- args,
- capture_output=True,
- text=True,
- check=False,
- )
- raw = proc.stdout.strip()
- return parse_rdns_output(ip=config.options.ip, raw=raw)
-
-
-def build_rdns_command(ip: str) -> str:
- # e.g. dig +noall +answer -x 1.2.3.4
- return " ".join(["dig", "+noall", "+answer", "-x", ip])
-
-
-def get_dig_version() -> str:
- proc = subprocess.run(
- ["dig", "-v"],
- capture_output=True,
- text=True,
- check=False,
- )
- # e.g. DiG 9.18.39-0ubuntu0.22.04.2-Ubuntu
- ver_str = proc.stderr.strip() + proc.stdout.strip()
- return ver_str.split("-", 1)[0].split(" ", 1)[1]
diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py
deleted file mode 100644
index cabd13c..0000000
--- a/generalresearch/models/network/rdns/execute.py
+++ /dev/null
@@ -1,44 +0,0 @@
-from __future__ import annotations
-
-from datetime import datetime, timezone
-from uuid import uuid4
-
-from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.network.rdns.command import (
- get_dig_version,
- run_rdns,
-)
-from generalresearch.models.network.tool_run import (
- RDNSRun,
- Status,
- ToolClass,
- ToolName,
-)
-from generalresearch.models.network.tool_run_command import (
- RDNSRunCommand,
- RDNSRunCommandOptions,
-)
-
-
-def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None):
- started_at = datetime.now(tz=timezone.utc)
- tool_version = get_dig_version()
- config = RDNSRunCommand(options=RDNSRunCommandOptions(ip=ip))
- result = run_rdns(config)
- finished_at = datetime.now(tz=timezone.utc)
-
- run = RDNSRun(
- tool_name=ToolName.DIG,
- tool_class=ToolClass.RDNS,
- tool_version=tool_version,
- status=Status.SUCCESS,
- ip=ip,
- started_at=started_at,
- finished_at=finished_at,
- raw_command=config.to_command_str(),
- scan_group_id=scan_group_id or uuid4().hex,
- config=config,
- parsed=result,
- )
-
- return run
diff --git a/generalresearch/models/network/rdns/parser.py b/generalresearch/models/network/rdns/parser.py
deleted file mode 100644
index 31a5ed6..0000000
--- a/generalresearch/models/network/rdns/parser.py
+++ /dev/null
@@ -1,21 +0,0 @@
-import ipaddress
-import re
-
-from generalresearch.models.custom_types import IPvAnyAddressStr
-from generalresearch.models.network.rdns.result import RDNSResult
-
-PTR_RE = re.compile(r"\sPTR\s+([^\s]+)\.")
-
-
-def parse_rdns_output(ip: IPvAnyAddressStr, raw: str) -> RDNSResult:
- hostnames: list[str] = []
-
- for line in raw.splitlines():
- m = PTR_RE.search(line)
- if m:
- hostnames.append(m.group(1))
-
- return RDNSResult(
- ip=ipaddress.ip_address(ip),
- hostnames=hostnames,
- )
diff --git a/generalresearch/models/network/rdns/result.py b/generalresearch/models/network/rdns/result.py
deleted file mode 100644
index 46af643..0000000
--- a/generalresearch/models/network/rdns/result.py
+++ /dev/null
@@ -1,52 +0,0 @@
-from __future__ import annotations
-
-import json
-from functools import cached_property
-
-import tldextract
-from pydantic import BaseModel, Field, computed_field, model_validator
-
-from generalresearch.models.custom_types import IPvAnyAddressStr
-
-
-class RDNSResult(BaseModel):
-
- ip: IPvAnyAddressStr = Field()
-
- hostnames: list[str] = Field(default_factory=list)
-
- @model_validator(mode="after")
- def validate_hostname_prop(self):
- assert len(self.hostnames) == self.hostname_count
- if self.hostnames:
- assert self.hostnames[0] == self.primary_hostname
- assert self.primary_domain in self.primary_hostname
- return self
-
- @computed_field(examples=["fixed-187-191-8-145.totalplay.net"])
- @cached_property
- def primary_hostname(self) -> str | None:
- if self.hostnames:
- return self.hostnames[0]
-
- @computed_field(examples=[1])
- @cached_property
- def hostname_count(self) -> int:
- return len(self.hostnames)
-
- @computed_field(examples=["totalplay.net"])
- @cached_property
- def primary_domain(self) -> str | None:
- if self.primary_hostname:
- return tldextract.extract(
- self.primary_hostname
- ).top_domain_under_public_suffix
-
- def model_dump_postgres(self):
- # Writes for the network_rdnsresult table
- d = self.model_dump(
- mode="json",
- include={"primary_hostname", "primary_domain", "hostname_count"},
- )
- d["hostnames"] = json.dumps(self.hostnames)
- return d
diff --git a/generalresearch/models/network/tool_run.py b/generalresearch/models/network/tool_run.py
deleted file mode 100644
index c49ffc0..0000000
--- a/generalresearch/models/network/tool_run.py
+++ /dev/null
@@ -1,118 +0,0 @@
-from __future__ import annotations
-
-from enum import StrEnum
-from typing import Literal
-from uuid import uuid4
-
-from pydantic import BaseModel, Field, PositiveInt
-
-from generalresearch.models.custom_types import (
- AwareDatetimeISO,
- IPvAnyAddressStr,
- UUIDStr,
-)
-from generalresearch.models.network.mtr.result import MTRResult
-from generalresearch.models.network.nmap.result import NmapResult
-from generalresearch.models.network.rdns.result import RDNSResult
-from generalresearch.models.network.tool_run_command import (
- MTRRunCommand,
- NmapRunCommand,
- RDNSRunCommand,
- ToolRunCommand,
-)
-
-
-class ToolClass(StrEnum):
- PORT_SCAN = "port_scan"
- RDNS = "rdns"
- PING = "ping"
- TRACEROUTE = "traceroute"
-
-
-class ToolName(StrEnum):
- NMAP = "nmap"
- RUSTMAP = "rustmap"
- DIG = "dig"
- PING = "ping"
- TRACEROUTE = "traceroute"
- MTR = "mtr"
-
-
-class Status(StrEnum):
- SUCCESS = "success"
- FAILED = "failed"
- TIMEOUT = "timeout"
- ERROR = "error"
-
-
-class ToolRun(BaseModel):
- """
- A run of a networking tool against one host/ip.
- """
-
- id: PositiveInt | None = Field(default=None)
-
- ip: IPvAnyAddressStr = Field()
- scan_group_id: UUIDStr = Field(default_factory=lambda: uuid4().hex)
- tool_class: ToolClass = Field()
- tool_name: ToolName = Field()
- tool_version: str = Field()
-
- started_at: AwareDatetimeISO = Field()
- finished_at: AwareDatetimeISO | None = Field(default=None)
- status: Status | None = Field(default=None)
-
- raw_command: str = Field()
-
- config: ToolRunCommand = Field()
-
- def model_dump_postgres(self):
- d = self.model_dump(mode="json", exclude={"config"})
- d["config"] = self.config.model_dump_json()
- return d
-
-
-class NmapRun(ToolRun):
- tool_class: Literal[ToolClass.PORT_SCAN] = Field(default=ToolClass.PORT_SCAN)
- tool_name: Literal[ToolName.NMAP] = Field(default=ToolName.NMAP)
- config: NmapRunCommand = Field()
-
- parsed: NmapResult = Field()
-
- def model_dump_postgres(self):
- d = super().model_dump_postgres()
- d["run_id"] = self.id
- d.update(self.parsed.model_dump_postgres())
- return d
-
-
-class RDNSRun(ToolRun):
- tool_class: Literal[ToolClass.RDNS] = Field(default=ToolClass.RDNS)
- tool_name: Literal[ToolName.DIG] = Field(default=ToolName.DIG)
- config: RDNSRunCommand = Field()
-
- parsed: RDNSResult = Field()
-
- def model_dump_postgres(self):
- d = super().model_dump_postgres()
- d["run_id"] = self.id
- d.update(self.parsed.model_dump_postgres())
- return d
-
-
-class MTRRun(ToolRun):
- tool_class: Literal[ToolClass.TRACEROUTE] = Field(default=ToolClass.TRACEROUTE)
- tool_name: Literal[ToolName.MTR] = Field(default=ToolName.MTR)
- config: MTRRunCommand = Field()
-
- facility_id: int = Field(default=1)
- source_ip: IPvAnyAddressStr = Field()
- parsed: MTRResult = Field()
-
- def model_dump_postgres(self):
- d = super().model_dump_postgres()
- d["run_id"] = self.id
- d["source_ip"] = self.source_ip
- d["facility_id"] = self.facility_id
- d.update(self.parsed.model_dump_postgres())
- return d
diff --git a/generalresearch/models/network/tool_run_command.py b/generalresearch/models/network/tool_run_command.py
deleted file mode 100644
index 6f22d6b..0000000
--- a/generalresearch/models/network/tool_run_command.py
+++ /dev/null
@@ -1,66 +0,0 @@
-from __future__ import annotations
-
-from typing import Literal
-
-from pydantic import BaseModel, Field
-
-from generalresearch.models.custom_types import IPvAnyAddressStr
-from generalresearch.models.network.definitions import IPProtocol
-
-
-class ToolRunCommand(BaseModel):
- command: str = Field()
- options: dict[str, str | int | None] = Field(default_factory=dict)
-
-
-class NmapRunCommandOptions(BaseModel):
- ip: IPvAnyAddressStr
- top_ports: int | None = Field(default=1000)
- ports: str | None = Field(default=None)
- no_ping: bool = Field(default=True)
- enable_advanced: bool = Field(default=True)
- timing: int = Field(default=4)
-
-
-class NmapRunCommand(ToolRunCommand):
- command: Literal["nmap"] = Field(default="nmap")
- options: NmapRunCommandOptions = Field()
-
- def to_command_str(self):
- from generalresearch.models.network.nmap.command import build_nmap_command
-
- options = self.options
- return build_nmap_command(**options.model_dump())
-
-
-class RDNSRunCommandOptions(BaseModel):
- ip: IPvAnyAddressStr
-
-
-class RDNSRunCommand(ToolRunCommand):
- command: Literal["dig"] = Field(default="dig")
- options: RDNSRunCommandOptions = Field()
-
- def to_command_str(self):
- from generalresearch.models.network.rdns.command import build_rdns_command
-
- options = self.options
- return build_rdns_command(**options.model_dump())
-
-
-class MTRRunCommandOptions(BaseModel):
- ip: IPvAnyAddressStr = Field()
- protocol: IPProtocol = Field(default=IPProtocol.ICMP)
- port: int | None = Field(default=None)
- report_cycles: int = Field(default=10)
-
-
-class MTRRunCommand(ToolRunCommand):
- command: Literal["mtr"] = Field(default="mtr")
- options: MTRRunCommandOptions = Field()
-
- def to_command_str(self):
- from generalresearch.models.network.mtr.command import build_mtr_command
-
- options = self.options
- return build_mtr_command(**options.model_dump())
diff --git a/generalresearch/models/network/utils.py b/generalresearch/models/network/utils.py
deleted file mode 100644
index fee9b80..0000000
--- a/generalresearch/models/network/utils.py
+++ /dev/null
@@ -1,5 +0,0 @@
-import requests
-
-
-def get_source_ip():
- return requests.get("https://icanhazip.com?").text.strip()
diff --git a/generalresearch/models/pollfish/question.py b/generalresearch/models/pollfish/question.py
index ae0dc5a..aadf71e 100644
--- a/generalresearch/models/pollfish/question.py
+++ b/generalresearch/models/pollfish/question.py
@@ -3,18 +3,16 @@ from __future__ import annotations
# https://wss.pollfish.com/mediation/documentation
import json
import logging
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Literal, Self
+from enum import StrEnum
+from typing import Any, Literal, Self
from pydantic import BaseModel, Field, model_validator
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -39,7 +37,7 @@ class PollfishQuestionOption(BaseModel):
order: int = Field()
-class PollfishQuestionType(str, Enum):
+class PollfishQuestionType(StrEnum):
"""
From the API: {'single_punch', 'multi_punch', 'open_ended'}
"""
diff --git a/generalresearch/models/precision/__init__.py b/generalresearch/models/precision/__init__.py
index 4bb2c6a..07deb8e 100644
--- a/generalresearch/models/precision/__init__.py
+++ b/generalresearch/models/precision/__init__.py
@@ -1,10 +1,10 @@
-from enum import Enum
+from enum import StrEnum
+from typing import Annotated
from pydantic import StringConstraints
-from typing_extensions import Annotated
-class PrecisionStatus(str, Enum):
+class PrecisionStatus(StrEnum):
# I made this up. They use isactive: "Yes" or "no", which I think is stupid
OPEN = "open"
CLOSED = "closed"
diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py
index 4030509..97f6ca1 100644
--- a/generalresearch/models/precision/question.py
+++ b/generalresearch/models/precision/question.py
@@ -3,22 +3,21 @@ from __future__ import annotations
# https://integrations.precisionsample.com/api.html#Get%20Questions
import json
import logging
-from enum import Enum
-from typing import TYPE_CHECKING, Any, Literal
+from enum import StrEnum
+from typing import Any, Literal
-from pydantic import BaseModel, Field, field_validator, model_validator
+from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator
-from generalresearch.models import Source, string_utils
+from generalresearch.models.definitions import Source
from generalresearch.models.precision import PrecisionQuestionID
+from generalresearch.models.string_utils import remove_nbsp
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
MarketplaceUserQuestionAnswer,
)
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -43,7 +42,7 @@ class PrecisionQuestionOption(BaseModel):
order: int = Field()
-class PrecisionQuestionType(str, Enum):
+class PrecisionQuestionType(StrEnum):
"""
From the API: {'Drop Down', 'Multi Select', 'Single Select', 'Single Select Matrix', 'Vertical Question'}
Of course undocumented. And there doesn't seem to be a text entry option?
@@ -54,15 +53,15 @@ class PrecisionQuestionType(str, Enum):
TEXT_ENTRY = "t"
@classmethod
- def from_api(cls, a: int):
- API_TYPE_MAP = {
+ def from_api(cls, a: int) -> PrecisionQuestionType | None:
+ api_type_map: dict[str, PrecisionQuestionType] = {
"Drop Down": PrecisionQuestionType.SINGLE_SELECT,
"Multi Select": PrecisionQuestionType.MULTI_SELECT,
"Single Select": PrecisionQuestionType.SINGLE_SELECT,
"Single Select Matrix": PrecisionQuestionType.SINGLE_SELECT,
"Vertical Question": PrecisionQuestionType.SINGLE_SELECT,
}
- return API_TYPE_MAP[a] if a in API_TYPE_MAP else None
+ return api_type_map.get(a, None)
class PrecisionUserQuestionAnswer(MarketplaceUserQuestionAnswer):
@@ -94,7 +93,7 @@ class PrecisionQuestion(MarketplaceQuestion):
@field_validator("question_text", mode="after")
def remove_nbsp(cls, s: str | None):
- return string_utils.remove_nbsp(s)
+ return remove_nbsp(s)
@model_validator(mode="after")
def check_type_options_agreement(self):
@@ -112,7 +111,7 @@ class PrecisionQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py
index 646d60e..cebe155 100644
--- a/generalresearch/models/precision/survey.py
+++ b/generalresearch/models/precision/survey.py
@@ -1,9 +1,9 @@
from __future__ import annotations
import json
-from datetime import timezone
+from datetime import UTC
from functools import cached_property
-from typing import Any, Dict, List, Literal, Optional, Self, Set, Tuple, Type
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
from pydantic import (
@@ -14,9 +14,7 @@ from pydantic import (
computed_field,
model_validator,
)
-from typing_extensions import Annotated
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AlphaNumStrSet,
AwareDatetimeISO,
@@ -24,6 +22,7 @@ from generalresearch.models.custom_types import (
DeviceTypes,
UUIDStrCoerce,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.precision import PrecisionQuestionID, PrecisionStatus
from generalresearch.models.thl.demographics import Gender
from generalresearch.models.thl.survey import MarketplaceTask
@@ -34,8 +33,8 @@ from generalresearch.models.thl.survey.condition import (
class PrecisionCondition(MarketplaceCondition):
- question_id: Optional[PrecisionQuestionID] = Field()
- values: List[Annotated[str, Field(max_length=128)]] = Field()
+ question_id: PrecisionQuestionID | None = Field()
+ values: list[Annotated[str, Field(max_length=128)]] = Field()
value_type: ConditionValueType = Field(default=ConditionValueType.LIST)
_CONVERT_LIST_TO_RANGE = ["age"]
@@ -54,7 +53,7 @@ class PrecisionQuota(BaseModel):
termination_count: int = Field(ge=0)
overquota_count: int = Field(ge=0)
- condition_hashes: List[str] = Field(min_length=1, default_factory=list)
+ condition_hashes: list[str] = Field(min_length=1, default_factory=list)
# Min spots a quota should have open to be OPEN
_min_open_spots: int = PrivateAttr(default=3)
@@ -78,7 +77,7 @@ class PrecisionQuota(BaseModel):
# TODO: I did some speed tests. This is faster than how this is implemented
# in sago/spectrum/dynata/etc. We should generalize this logic instead of
# copying/pasting it 7 times. (matches, matches_optional and _soft)
- def matches(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def matches(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Matches means we meet all conditions.
# In Morning, all quotas are mutually exclusive. so if it doesn't
# matter if we match a closed quota, b/c that means that we won't
@@ -86,8 +85,8 @@ class PrecisionQuota(BaseModel):
return self.matches_optional(criteria_evaluation) is True
def matches_optional(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Optional[bool]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> bool | None:
for c in self.condition_hashes:
eval_value = criteria_evaluation.get(c)
if eval_value is False:
@@ -97,14 +96,14 @@ class PrecisionQuota(BaseModel):
return True
def matches_soft(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], List[str]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, list[str]]:
# Passes back "matches" (T/F/none) and a list of unknown criterion hashes
- unknowns = list()
+ unknowns = []
for c in self.condition_hashes:
eval_value = criteria_evaluation.get(c)
if eval_value is False:
- return False, list()
+ return False, []
if eval_value is None:
unknowns.append(c)
if unknowns:
@@ -128,7 +127,7 @@ class PrecisionSurvey(MarketplaceTask):
name: str = Field(validation_alias="prj_name")
survey_guid: UUIDStrCoerce = Field(validation_alias="prj_guid")
- category_id: Optional[str] = Field(validation_alias="sc_id", default=None)
+ category_id: str | None = Field(validation_alias="sc_id", default=None)
buyer_id: CoercedStr = Field(max_length=16)
# This seems to always be 0 ... ?
@@ -141,7 +140,7 @@ class PrecisionSurvey(MarketplaceTask):
bid_ir: float = Field(ge=0, le=1, validation_alias="ir")
# Be careful with this, it doesn't make any sense. See survey 452481, has 12 completes with a 100% live_ir,
# but the only quotas have 0 completes and 1052 terms. .... ??
- global_conversion: Optional[float] = Field(
+ global_conversion: float | None = Field(
ge=0,
le=1,
default=None,
@@ -156,31 +155,31 @@ class PrecisionSurvey(MarketplaceTask):
allowed_devices: DeviceTypes = Field(min_length=1)
entry_link: str = Field(validation_alias="url")
- excluded_surveys: Optional[AlphaNumStrSet] = Field(
+ excluded_surveys: AlphaNumStrSet | None = Field(
description="list of excluded survey ids",
default=None,
validation_alias="exclusion_project_id",
)
- quotas: List[PrecisionQuota] = Field(default_factory=list)
+ quotas: list[PrecisionQuota] = Field(default_factory=list)
source: Literal[Source.PRECISION] = Field(default=Source.PRECISION)
- used_question_ids: Set[PrecisionQuestionID] = Field(default_factory=set)
+ used_question_ids: set[PrecisionQuestionID] = Field(default_factory=set)
# This is a "special" key to store all conditions that are used (as "condition_hashes") throughout
# this survey. In the reduced representation of this task (nearly always, for db i/o, in global_vars)
# this field will be null.
- conditions: Optional[Dict[str, PrecisionCondition]] = Field(default=None)
+ conditions: dict[str, PrecisionCondition] | None = Field(default=None)
# This comes from the API
- expected_end_date: Optional[AwareDatetimeISO] = Field(
+ expected_end_date: AwareDatetimeISO | None = Field(
default=None, validation_alias="end_date"
)
# This does not come from the API. We set it when we update this in the db.
- created: Optional[AwareDatetimeISO] = Field(default=None)
- updated: Optional[AwareDatetimeISO] = Field(default=None)
+ created: AwareDatetimeISO | None = Field(default=None)
+ updated: AwareDatetimeISO | None = Field(default=None)
@property
def internal_id(self) -> str:
@@ -199,7 +198,7 @@ class PrecisionSurvey(MarketplaceTask):
@computed_field
@cached_property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
s = set()
for q in self.quotas:
s.update(set(q.condition_hashes))
@@ -219,7 +218,7 @@ class PrecisionSurvey(MarketplaceTask):
return data
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return PrecisionCondition
@property
@@ -227,7 +226,7 @@ class PrecisionSurvey(MarketplaceTask):
return "age"
@property
- def marketplace_genders(self) -> Dict[Gender, Optional[MarketplaceCondition]]:
+ def marketplace_genders(self) -> dict[Gender, MarketplaceCondition | None]:
return {
Gender.MALE: PrecisionCondition(
question_id="gender",
@@ -246,11 +245,10 @@ class PrecisionSurvey(MarketplaceTask):
# Fancy repr that abbreviates exclude_pids and excluded_surveys
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k in {"excluded_surveys"}:
- if v and len(v) > 6:
- v = sorted(v)
- v = v[:3] + ["…"] + v[-3:]
- repr_args[n] = (k, v)
+ if k in {"excluded_surveys"} and v and len(v) > 6:
+ v = sorted(v)
+ v = v[:3] + ["…"] + v[-3:]
+ repr_args[n] = (k, v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -262,7 +260,7 @@ class PrecisionSurvey(MarketplaceTask):
exclude={"updated", "conditions", "created"}
) == other.model_dump(exclude={"updated", "conditions", "created"})
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(
mode="json",
exclude={
@@ -283,11 +281,11 @@ class PrecisionSurvey(MarketplaceTask):
return d
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> Self:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
+ def from_db(cls, d: dict[str, Any]) -> Self:
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
d["expected_end_date"] = (
- d["expected_end_date"].replace(tzinfo=timezone.utc)
+ d["expected_end_date"].replace(tzinfo=UTC)
if d["expected_end_date"]
else None
)
@@ -295,7 +293,7 @@ class PrecisionSurvey(MarketplaceTask):
d["used_question_ids"] = json.loads(d["used_question_ids"])
return cls.model_validate(d)
- def passes_quotas(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# We have to match 1 or more quota.
# Quotas are exclusionary: they can NOT match a quota where currently_open=0
any_pass = False
@@ -308,13 +306,13 @@ class PrecisionSurvey(MarketplaceTask):
return any_pass
def passes_quotas_soft(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], Set[str]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, set[str]]:
# Quotas are exclusionary. They can NOT match a quota where currently_open=0
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
@@ -345,19 +343,19 @@ class PrecisionSurvey(MarketplaceTask):
return False, set()
def determine_eligibility(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
return self.is_open and self.passes_quotas(criteria_evaluation)
def determine_eligibility_soft(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Tuple[Optional[bool], Optional[Set[str]]]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> tuple[bool | None, set[str] | None]:
if not self.is_open:
return False, None
return self.passes_quotas_soft(criteria_evaluation)
def participation_allowed(
- self, att_survey_ids: Set[str], att_group_ids: Set[str]
+ self, att_survey_ids: set[str], att_group_ids: set[str]
) -> bool:
"""
Checks if this user can participate in this survey
@@ -370,6 +368,4 @@ class PrecisionSurvey(MarketplaceTask):
return False
if self.group_id in att_group_ids:
return False
- if self.excluded_surveys & att_survey_ids:
- return False
- return True
+ return not self.excluded_surveys & att_survey_ids
diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py
index 233d329..c8db2af 100644
--- a/generalresearch/models/precision/task_collection.py
+++ b/generalresearch/models/precision/task_collection.py
@@ -1,4 +1,4 @@
-from typing import Any, Dict, List
+from typing import Any
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
@@ -36,8 +36,8 @@ PrecisionTaskCollectionSchema = DataFrameSchema(
"expected_end_date": Column(dtype=pd.DatetimeTZDtype(tz="UTC"), nullable=True),
"created": Column(dtype=pd.DatetimeTZDtype(tz="UTC")),
"updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")),
- "used_question_ids": Column(List[str]),
- "all_hashes": Column(List[str]), # set >> list for column support
+ "used_question_ids": Column(list[str]),
+ "all_hashes": Column(list[str]), # set >> list for column support
},
checks=[],
index=Index(
@@ -53,10 +53,10 @@ PrecisionTaskCollectionSchema = DataFrameSchema(
class PrecisionTaskCollection(TaskCollection):
- items: List[PrecisionSurvey]
+ items: list[PrecisionSurvey]
_schema = PrecisionTaskCollectionSchema
- def to_row(self, s: PrecisionSurvey) -> Dict[str, Any]:
+ def to_row(self, s: PrecisionSurvey) -> dict[str, Any]:
d = s.model_dump(
mode="json",
exclude={
diff --git a/generalresearch/models/prodege/__init__.py b/generalresearch/models/prodege/__init__.py
index d419c0c..64b148a 100644
--- a/generalresearch/models/prodege/__init__.py
+++ b/generalresearch/models/prodege/__init__.py
@@ -1,15 +1,14 @@
-from enum import Enum
-from typing import Literal
+from enum import StrEnum
+from typing import Annotated, Literal
from pydantic import Field
-from typing_extensions import Annotated
ProdegeQuestionIdType = Annotated[
str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
]
-class ProdegeStatus(str, Enum):
+class ProdegeStatus(StrEnum):
LIVE = "LIVE"
# We need another status to mark if a survey we thought was live does not come back
# from the API, we'll mark it as NOT_FOUND
@@ -19,7 +18,7 @@ class ProdegeStatus(str, Enum):
INELIGIBLE = "INELIGIBLE"
-class ProdegePastParticipationType(str, Enum):
+class ProdegePastParticipationType(StrEnum):
# These come from the "participation_types" key in the survey API response
# which is how we filter by users' past_participation.
CLICK = "click"
diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py
index 1c61ab9..c0160bd 100644
--- a/generalresearch/models/prodege/question.py
+++ b/generalresearch/models/prodege/question.py
@@ -3,23 +3,28 @@ from __future__ import annotations
import json
import logging
-from datetime import datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from functools import cached_property
-from typing import TYPE_CHECKING, Any, Literal
+from typing import Any, Literal
-from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator
+from pydantic import (
+ BaseModel,
+ ConfigDict,
+ Field,
+ PositiveInt,
+ ValidationError,
+ model_validator,
+)
from generalresearch.locales import Localelator
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.prodege import ProdegeQuestionIdType
from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -43,9 +48,7 @@ class ProdegeUserQuestionAnswer(BaseModel):
# This may be a pipe-separated string if the question_type is multi. regex means any chars except capital letters
option_id: str = Field(pattern=r"^[^A-Z]*$")
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# ISO 3166-1 alpha-2 (two-letter codes, lowercase)
country_iso: str = Field(
@@ -92,7 +95,7 @@ class ProdegeQuestionOption(BaseModel):
is_exclusive: bool = Field(default=False)
-class ProdegeQuestionType(str, Enum):
+class ProdegeQuestionType(StrEnum):
"""
{'Derived', 'Multi Punch', 'Numeric - Open End', 'Single Punch', 'Zip Code'}
"""
@@ -145,7 +148,7 @@ class ProdegeQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d, country_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py
index c12f130..27034b0 100644
--- a/generalresearch/models/prodege/survey.py
+++ b/generalresearch/models/prodege/survey.py
@@ -4,22 +4,22 @@ from __future__ import annotations
import json
import logging
from collections import defaultdict
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal, Type
+from typing import Any, Literal
from pydantic import (
BaseModel,
ConfigDict,
Field,
+ ValidationError,
computed_field,
field_validator,
model_validator,
)
from generalresearch.locales import Localelator
-from generalresearch.models import LogicalOperator, Source, TaskCalculationType
from generalresearch.models.custom_types import (
AlphaNumStrSet,
AwareDatetimeISO,
@@ -27,6 +27,11 @@ from generalresearch.models.custom_types import (
InclExcl,
UUIDStr,
)
+from generalresearch.models.definitions import (
+ LogicalOperator,
+ Source,
+ TaskCalculationType,
+)
from generalresearch.models.prodege import (
ProdegePastParticipationType,
ProdegeQuestionIdType,
@@ -127,7 +132,7 @@ class ProdegeQuota(BaseModel):
return self.remaining_count >= min_open_spots
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return ProdegeCondition
@property
@@ -276,7 +281,7 @@ class ProdegeUserPastParticipation(BaseModel):
raise ValueError(f"Unknown ext_status_code_1: {self.ext_status_code_1}")
def days_ago(self) -> float:
- now = datetime.now(timezone.utc)
+ now = datetime.now(UTC)
return (now - self.started).total_seconds() / (3600 * 24)
@@ -486,7 +491,7 @@ class ProdegeSurvey(MarketplaceTask):
return data
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return ProdegeCondition
@property
@@ -513,7 +518,7 @@ class ProdegeSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> ProdegeSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@@ -538,7 +543,7 @@ class ProdegeSurvey(MarketplaceTask):
d["country_isos"] = [
locale_helper.get_country_iso(d.pop("country_code").lower())
]
- d["country_iso"] = sorted(d["country_isos"])[0]
+ d["country_iso"] = min(d["country_isos"])
# No languages are returned anywhere for anything
d["language_isos"] = [
locale_helper.get_default_lang_from_country(d["country_isos"][0])
@@ -551,7 +556,7 @@ class ProdegeSurvey(MarketplaceTask):
d["past_participation"] = ProdegePastParticipation.from_api(
d["past_participation"]
)
- d["conditions"] = dict()
+ d["conditions"] = {}
for quota in d["quotas"]:
quota["condition_hashes"] = []
for c in quota["targeting_criteria"]:
@@ -562,7 +567,7 @@ class ProdegeSurvey(MarketplaceTask):
d["quotas"] = [ProdegeQuota.from_api(q) for q in d["quotas"]]
countries = {q.country_iso for q in d["quotas"] if q.country_iso}
if countries:
- d["country_iso"] = sorted(countries)[0]
+ d["country_iso"] = min(countries)
d["country_isos"] = countries
d["language_iso"] = locale_helper.get_default_lang_from_country(
d["country_iso"]
@@ -656,8 +661,8 @@ class ProdegeSurvey(MarketplaceTask):
@classmethod
def from_db(cls, d: dict[str, Any]) -> ProdegeSurvey:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
d["quotas"] = json.loads(d["quotas"])
for k in [
"max_clicks_settings",
diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py
index 4544050..9f6a81b 100644
--- a/generalresearch/models/prodege/task_collection.py
+++ b/generalresearch/models/prodege/task_collection.py
@@ -76,7 +76,7 @@ class ProdegeTaskCollection(TaskCollection):
"used_question_ids",
"all_hashes",
]
- d = dict()
+ d = {}
for k in fields:
d[k] = getattr(s, k)
d["cpi"] = float(d["cpi"])
diff --git a/generalresearch/models/repdata/__init__.py b/generalresearch/models/repdata/__init__.py
index 9706b28..4a3d8d0 100644
--- a/generalresearch/models/repdata/__init__.py
+++ b/generalresearch/models/repdata/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class RepDataStatus(str, Enum):
+class RepDataStatus(StrEnum):
LIVE = "LIVE"
DRAFT = "DRAFT"
PAUSED = "PAUSED"
diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py
index 8d0da13..5115426 100644
--- a/generalresearch/models/repdata/question.py
+++ b/generalresearch/models/repdata/question.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import json
import logging
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
-from typing import TYPE_CHECKING, Any, Literal
+from typing import Any, Literal
from uuid import UUID
from pydantic import (
@@ -12,18 +12,17 @@ from pydantic import (
ConfigDict,
Field,
PositiveInt,
+ ValidationError,
field_validator,
model_validator,
)
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -86,7 +85,7 @@ class RepDataQuestionOption(BaseModel):
order: int = Field()
-class RepDataQuestionType(str, Enum):
+class RepDataQuestionType(StrEnum):
"""
{'Derived', 'Multi Punch', 'Numeric - Open End', 'Single Punch', 'Zip Code'}
"""
@@ -142,6 +141,7 @@ class RepDataQuestion(MarketplaceQuestion):
@property
def internal_id(self) -> str:
+ assert self.lucid_id
return self.lucid_id
@field_validator("question_id", mode="before")
@@ -167,7 +167,7 @@ class RepDataQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d, country_iso, language_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py
index 2290ca6..a69a54b 100644
--- a/generalresearch/models/repdata/survey.py
+++ b/generalresearch/models/repdata/survey.py
@@ -3,35 +3,35 @@ from __future__ import annotations
import json
import logging
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal, Type
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import (
BaseModel,
ConfigDict,
Field,
+ ValidationError,
computed_field,
field_validator,
model_validator,
)
-from typing_extensions import Self
from generalresearch.grpc import timestamp_from_datetime
from generalresearch.locales import Localelator
-from generalresearch.models import (
- DeviceType,
- LogicalOperator,
- Source,
- TaskCalculationType,
-)
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CoercedStr,
UUIDStr,
)
+from generalresearch.models.definitions import (
+ DeviceType,
+ LogicalOperator,
+ Source,
+ TaskCalculationType,
+)
from generalresearch.models.repdata import RepDataStatus
from generalresearch.models.thl.demographics import Gender
from generalresearch.models.thl.survey import MarketplaceTask
@@ -304,7 +304,7 @@ class RepDataStream(MarketplaceTask):
return self.stream_status == RepDataStatus.LIVE
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return RepDataCondition
@property
@@ -460,7 +460,7 @@ class RepDataSurvey(BaseModel):
@property
def all_conditions(self) -> list[RepDataCondition]:
- cs = list()
+ cs = []
for stream in self.streams:
cs.extend(stream.all_conditions)
# dedupe by criterion_hash
@@ -478,7 +478,7 @@ class RepDataSurvey(BaseModel):
"""
try:
return cls._from_api(survey_response)
- except Exception as e:
+ except ValidationError as e:
survey_id = survey_response.get("survey_id") or survey_response.get(
"SurveyNumber"
)
@@ -486,7 +486,7 @@ class RepDataSurvey(BaseModel):
return None
@classmethod
- def _from_api(cls, survey_response) -> RepDataSurvey:
+ def _from_api(cls, survey_response: dict[str, Any]) -> RepDataSurvey:
d = survey_response.copy()
d["country_iso"] = locale_helper.get_country_iso(d["SurveyCountry"].lower())
d["language_iso"] = locale_helper.get_language_iso(d["SurveyLanguage"].lower())
@@ -538,8 +538,8 @@ class RepDataSurveyHashed(RepDataSurvey):
DeviceType(int(x)) for x in res["allowed_devices"].split(",")
]
if res["created"] is not None:
- res["created"] = res["created"].replace(tzinfo=timezone.utc)
- res["last_updated"] = res["last_updated"].replace(tzinfo=timezone.utc)
+ res["created"] = res["created"].replace(tzinfo=UTC)
+ res["last_updated"] = res["last_updated"].replace(tzinfo=UTC)
return cls.model_validate(res)
def to_mysql(self) -> dict[str, Any]:
@@ -553,7 +553,7 @@ class RepDataSurveyHashed(RepDataSurvey):
return d
def to_grpc(self, repdata_pb2):
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
timestamp = timestamp_from_datetime(now)
return repdata_pb2.RepDataOpportunity(
diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py
index d625349..aa591a9 100644
--- a/generalresearch/models/repdata/task_collection.py
+++ b/generalresearch/models/repdata/task_collection.py
@@ -6,7 +6,7 @@ import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.locales import Localelator
-from generalresearch.models import TaskCalculationType
+from generalresearch.models.definitions import TaskCalculationType
from generalresearch.models.repdata import RepDataStatus
from generalresearch.models.repdata.survey import RepDataSurveyHashed
from generalresearch.models.thl.survey.task_collection import (
@@ -110,7 +110,7 @@ class RepDataTaskCollection(TaskCollection):
"remaining_count",
]
rows = []
- d = dict()
+ d = {}
for k in survey_fields:
d[k] = getattr(s, k)
d["allowed_devices"] = s.allowed_devices_str
diff --git a/generalresearch/models/sago/__init__.py b/generalresearch/models/sago/__init__.py
index 292f0f2..1059a6f 100644
--- a/generalresearch/models/sago/__init__.py
+++ b/generalresearch/models/sago/__init__.py
@@ -1,13 +1,13 @@
-from enum import Enum
+from enum import StrEnum
+from typing import Annotated
from pydantic import Field
-from typing_extensions import Annotated
SagoQuestionIdType = Annotated[
str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
]
-class SagoStatus(str, Enum):
+class SagoStatus(StrEnum):
LIVE = "LIVE"
NOT_LIVE = "NOT_LIVE"
diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py
index f911854..148b015 100644
--- a/generalresearch/models/sago/question.py
+++ b/generalresearch/models/sago/question.py
@@ -4,27 +4,27 @@ from __future__ import annotations
# -answers-lanaguge-languageid
import json
import logging
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
-from typing import TYPE_CHECKING, Any, Literal
+from typing import Any, Literal
from pydantic import (
BaseModel,
ConfigDict,
Field,
PositiveInt,
+ ValidationError,
field_validator,
model_validator,
)
-from generalresearch.models import MAX_INT32, Source, string_utils
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import MAX_INT32, Source
+from generalresearch.models.string_utils import remove_nbsp
from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -54,10 +54,10 @@ class SagoQuestionOption(BaseModel):
@field_validator("text", mode="after")
def remove_nbsp(cls, s: str):
- return string_utils.remove_nbsp(s)
+ return remove_nbsp(s)
-class SagoQuestionType(str, Enum):
+class SagoQuestionType(StrEnum):
"""
From the API:
{1: 'Single Punch', 2: 'Multi Punch', 3: 'Open Ended', 4: 'Dummy',
@@ -86,7 +86,7 @@ class SagoQuestionType(str, Enum):
6: SagoQuestionType.TEXT_ENTRY,
7: SagoQuestionType.TEXT_ENTRY,
}
- return API_TYPE_MAP[a] if a in API_TYPE_MAP else None
+ return API_TYPE_MAP.get(a, None)
class SagoUserQuestionAnswer(BaseModel):
@@ -168,7 +168,7 @@ class SagoQuestion(MarketplaceQuestion):
@field_validator("question_name", "question_text", "tags", mode="after")
def remove_nbsp(cls, s: str | None):
- return string_utils.remove_nbsp(s)
+ return remove_nbsp(s)
@classmethod
def from_api(
@@ -182,7 +182,7 @@ class SagoQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d, country_iso, language_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py
index e73ddee..5330ace 100644
--- a/generalresearch/models/sago/survey.py
+++ b/generalresearch/models/sago/survey.py
@@ -2,17 +2,22 @@ from __future__ import annotations
import json
import logging
-from datetime import timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
-from typing import Annotated, Any, Literal, Type
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
-from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator
-from typing_extensions import Self
+from pydantic import (
+ BaseModel,
+ ConfigDict,
+ Field,
+ ValidationError,
+ computed_field,
+ model_validator,
+)
from generalresearch.locales import Localelator
-from generalresearch.models import LogicalOperator, Source
from generalresearch.models.custom_types import (
AlphaNumStr,
AlphaNumStrSet,
@@ -21,6 +26,7 @@ from generalresearch.models.custom_types import (
DeviceTypes,
IPLikeStrSet,
)
+from generalresearch.models.definitions import LogicalOperator, Source
from generalresearch.models.sago import SagoStatus
from generalresearch.models.thl.demographics import Gender
from generalresearch.models.thl.survey import MarketplaceTask
@@ -72,7 +78,7 @@ class SagoQuota(BaseModel):
# There is no explicit status. The quota is closed if the count is 0
def __hash__(self) -> int:
- return hash(tuple((tuple(self.condition_hashes), self.remaining_count)))
+ return hash((tuple(self.condition_hashes), self.remaining_count))
@property
def is_open(self) -> bool:
@@ -235,7 +241,7 @@ class SagoSurvey(MarketplaceTask):
return data
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return SagoCondition
@property
@@ -262,7 +268,7 @@ class SagoSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> SagoSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@@ -274,11 +280,10 @@ class SagoSurvey(MarketplaceTask):
# Fancy repr that abbreviates ip_exclusions and survey_exclusions
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k in {"ip_exclusions", "survey_exclusions"}:
- if v and len(v) > 6:
- v = sorted(v)
- v = v[:3] + ["…"] + v[-3:]
- repr_args[n] = (k, v)
+ if k in {"ip_exclusions", "survey_exclusions"} and v and len(v) > 6:
+ v = sorted(v)
+ v = v[:3] + ["…"] + v[-3:]
+ repr_args[n] = (k, v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -314,9 +319,9 @@ class SagoSurvey(MarketplaceTask):
@classmethod
def from_db(cls, d: dict[str, Any]):
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
- d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc)
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
+ d["modified_api"] = d["modified_api"].replace(tzinfo=UTC)
d["qualifications"] = json.loads(d["qualifications"])
d["used_question_ids"] = json.loads(d["used_question_ids"])
d["quotas"] = json.loads(d["quotas"])
@@ -363,7 +368,7 @@ class SagoSurvey(MarketplaceTask):
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
diff --git a/generalresearch/models/spectrum/__init__.py b/generalresearch/models/spectrum/__init__.py
index b62c089..0040551 100644
--- a/generalresearch/models/spectrum/__init__.py
+++ b/generalresearch/models/spectrum/__init__.py
@@ -1,7 +1,7 @@
from enum import Enum
+from typing import Annotated
from pydantic import Field
-from typing_extensions import Annotated
SpectrumQuestionIdType = Annotated[
str, Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py
index db8a55d..839f7a8 100644
--- a/generalresearch/models/spectrum/question.py
+++ b/generalresearch/models/spectrum/question.py
@@ -3,32 +3,31 @@ from __future__ import annotations
import json
import logging
-from datetime import datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import IntEnum, StrEnum
from functools import cached_property
-from typing import TYPE_CHECKING, Any, Literal
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import (
BaseModel,
Field,
PositiveInt,
+ ValidationError,
field_validator,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.models import MAX_INT32, Source, string_utils
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.spectrum import SpectrumQuestionIdType
+from generalresearch.models.string_utils import remove_nbsp
from generalresearch.models.thl.profiling.marketplace import (
MarketplaceQuestion,
)
-
-if TYPE_CHECKING:
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
+from generalresearch.models.thl.profiling.upk_question import (
+ UpkQuestion,
+)
logging.basicConfig()
logger = logging.getLogger()
@@ -50,9 +49,7 @@ class SpectrumUserQuestionAnswer(BaseModel):
# This may be a pipe-separated string if the question_type is multi. regex
# means any chars except capital letters
option_id: str = Field(pattern=r"^[^A-Z]*$")
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# ISO 3166-1 alpha-2 (two-letter codes, lowercase)
country_iso: str = Field(
max_length=2, min_length=2, pattern=r"^[a-z]{2}$", frozen=True
@@ -93,10 +90,12 @@ class SpectrumQuestionOption(BaseModel):
@field_validator("text", mode="after")
def remove_nbsp(cls, s: str) -> str:
- return string_utils.remove_nbsp(s)
+ res = remove_nbsp(s)
+ assert isinstance(res, str), "Spectrum Question Option text must be str"
+ return res
-class SpectrumQuestionType(str, Enum):
+class SpectrumQuestionType(StrEnum):
# The documentation defines 4 types (1,2,3,4), however 2 is the same as 1
# and never comes back in the api, and we also get back 5, 6, and 7,
# which are all undocumented.
@@ -135,10 +134,10 @@ class SpectrumQuestionType(str, Enum):
@classmethod
def from_api(cls, a: int):
api_type_map = cls.get_api_map()
- return api_type_map[a] if a in api_type_map else None
+ return api_type_map.get(a, None)
-class SpectrumQuestionClass(int, Enum):
+class SpectrumQuestionClass(IntEnum):
CORE = 1
EXTENDED = 2
CUSTOM = 3
@@ -208,7 +207,7 @@ class SpectrumQuestion(MarketplaceQuestion):
@field_validator("question_name", "question_text", "tags", mode="after")
def remove_nbsp(cls, s: str | None):
- return string_utils.remove_nbsp(s)
+ return remove_nbsp(s)
@model_validator(mode="before")
@classmethod
@@ -263,7 +262,7 @@ class SpectrumQuestion(MarketplaceQuestion):
return None
try:
return cls._from_api(d, country_iso, language_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
@@ -283,7 +282,9 @@ class SpectrumQuestion(MarketplaceQuestion):
]
created = (
- datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=timezone.utc)
+ datetime.fromtimestamp(timestamp=d["crtd_on"] / 1000, tz=UTC).replace(
+ tzinfo=UTC
+ )
if d.get("crtd_on")
else None
)
@@ -308,9 +309,7 @@ class SpectrumQuestion(MarketplaceQuestion):
SpectrumQuestionOption(id=r["id"], text=r["text"], order=r["order"])
for r in d["options"]
]
- d["created"] = (
- d["created"].replace(tzinfo=timezone.utc) if d["created"] else None
- )
+ d["created"] = d["created"].replace(tzinfo=UTC) if d["created"] else None
return cls(
question_id=d["question_id"],
diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py
index a591445..5b330a8 100644
--- a/generalresearch/models/spectrum/survey.py
+++ b/generalresearch/models/spectrum/survey.py
@@ -2,16 +2,14 @@ from __future__ import annotations
import json
import logging
-from datetime import timezone
+from datetime import UTC
from decimal import Decimal
-from typing import Any, Literal, Type
+from typing import Any, Literal, Self
from more_itertools import flatten
from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator
-from typing_extensions import Self
from generalresearch.locales import Localelator
-from generalresearch.models import Source, TaskCalculationType
from generalresearch.models.custom_types import (
AlphaNumStr,
AlphaNumStrSet,
@@ -19,6 +17,7 @@ from generalresearch.models.custom_types import (
CoercedStr,
UUIDStrSet,
)
+from generalresearch.models.definitions import Source, TaskCalculationType
from generalresearch.models.spectrum import SpectrumStatus
from generalresearch.models.thl.demographics import Gender
from generalresearch.models.thl.survey import MarketplaceTask
@@ -56,7 +55,7 @@ class SpectrumCondition(MarketplaceCondition):
try:
values = [tuple(map(int, v.split("-"))) for v in self.values]
assert all(len(x) == 2 for x in values)
- except (ValueError, AssertionError):
+ except ValueError, AssertionError:
return self
self.values = sorted(
{str(val) for tupl in values for val in range(tupl[0], tupl[1] + 1)}
@@ -76,8 +75,7 @@ class SpectrumCondition(MarketplaceCondition):
rs["from"] = round(rs["from"] / 12)
rs["to"] = round(rs["to"] / 12)
d["values"] = [
- "{0}-{1}".format(rs["from"] or "inf", rs["to"] or "inf")
- for rs in d["range_sets"]
+ f"{rs['from'] or 'inf'}-{rs['to'] or 'inf'}" for rs in d["range_sets"]
]
d["value_type"] = ConditionValueType.RANGE
return cls.model_validate(d)
@@ -104,7 +102,7 @@ class SpectrumQuota(BaseModel):
# There is no explicit status. The quota is closed if the count is 0
def __hash__(self) -> int:
- return hash(tuple((tuple(self.condition_hashes), self.remaining_count)))
+ return hash((tuple(self.condition_hashes), self.remaining_count))
@property
def is_open(self) -> bool:
@@ -114,7 +112,7 @@ class SpectrumQuota(BaseModel):
return self.remaining_count >= min_open_spots
@classmethod
- def from_api(cls, d: dict) -> Self:
+ def from_api(cls, d: dict[str, Any]) -> Self:
d["remaining_count"] = d["quantities"]["currently_open"]
return cls.model_validate(d)
@@ -297,7 +295,7 @@ class SpectrumSurvey(MarketplaceTask):
return data
@property
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
return SpectrumCondition
@property
@@ -324,7 +322,7 @@ class SpectrumSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> SpectrumSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except (AssertionError, ValueError) as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@@ -337,7 +335,7 @@ class SpectrumSurvey(MarketplaceTask):
else TaskCalculationType.COMPLETES
)
- d["conditions"] = dict()
+ d["conditions"] = {}
# If we haven't hit the "detail" endpoint, we won't get this
d.setdefault("qualifications", [])
@@ -389,16 +387,14 @@ class SpectrumSurvey(MarketplaceTask):
@classmethod
def from_db(cls, d: dict[str, Any]) -> Self:
- d["created_api"] = d["created_api"].replace(tzinfo=timezone.utc)
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
- d["modified_api"] = d["modified_api"].replace(tzinfo=timezone.utc)
+ d["created_api"] = d["created_api"].replace(tzinfo=UTC)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
+ d["modified_api"] = d["modified_api"].replace(tzinfo=UTC)
d["field_end_date"] = (
- d["field_end_date"].replace(tzinfo=timezone.utc)
- if d["field_end_date"]
- else None
+ d["field_end_date"].replace(tzinfo=UTC) if d["field_end_date"] else None
)
d["project_last_complete_date"] = (
- d["project_last_complete_date"].replace(tzinfo=timezone.utc)
+ d["project_last_complete_date"].replace(tzinfo=UTC)
if d["project_last_complete_date"]
else None
)
@@ -457,7 +453,7 @@ class SpectrumSurvey(MarketplaceTask):
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py
index 6715378..609114e 100644
--- a/generalresearch/models/spectrum/task_collection.py
+++ b/generalresearch/models/spectrum/task_collection.py
@@ -4,7 +4,7 @@ import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.locales import Localelator
-from generalresearch.models import TaskCalculationType
+from generalresearch.models.definitions import TaskCalculationType
from generalresearch.models.spectrum import SpectrumStatus
from generalresearch.models.spectrum.survey import SpectrumSurvey
from generalresearch.models.thl.survey.task_collection import (
@@ -91,7 +91,7 @@ class SpectrumTaskCollection(TaskCollection):
"survey_id",
]
rows = []
- d = dict()
+ d = {}
for k in fields:
d[k] = getattr(s, k) if hasattr(s, k) else None
d["used_question_ids"] = list(s.used_question_ids)
diff --git a/generalresearch/models/string_utils.py b/generalresearch/models/string_utils.py
index 23c1017..dff2f4d 100644
--- a/generalresearch/models/string_utils.py
+++ b/generalresearch/models/string_utils.py
@@ -1,8 +1,7 @@
import unicodedata
-from typing import Optional
-def remove_nbsp(s: Optional[str]) -> Optional[str]:
+def remove_nbsp(s: str | None) -> str | None:
# Some text comes back from the API with lots of (copied from excel or
# something), and random unicode...
if s:
diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py
index abc129b..4a29349 100644
--- a/generalresearch/models/thl/__init__.py
+++ b/generalresearch/models/thl/__init__.py
@@ -1,33 +1,17 @@
-from decimal import Decimal
-
from generalresearch.models.thl.finance import (
POPFinancial,
ProductBalances,
)
+from generalresearch.models.thl.ledger import LedgerAccount
from generalresearch.models.thl.payout import (
BrokerageProductPayoutEvent,
PayoutEvent,
)
from generalresearch.models.thl.product import Product
-_ = (
- Product,
- PayoutEvent,
- BrokerageProductPayoutEvent,
- ProductBalances,
- POPFinancial,
-)
-
Product.model_rebuild()
PayoutEvent.model_rebuild()
BrokerageProductPayoutEvent.model_rebuild()
-
-
-def decimal_to_int_cents(usd: Decimal | None) -> int | None:
- return round(usd * 100) if usd is not None else None
-
-
-def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None:
- if value is None:
- return None
- return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals)
+ProductBalances.model_rebuild()
+POPFinancial.model_rebuild()
+LedgerAccount.model_rebuild()
diff --git a/generalresearch/models/thl/category.py b/generalresearch/models/thl/category.py
index 1ed436a..ebfc840 100644
--- a/generalresearch/models/thl/category.py
+++ b/generalresearch/models/thl/category.py
@@ -1,10 +1,9 @@
from __future__ import annotations
-from typing import Any
+from typing import Any, Self
from uuid import uuid4
from pydantic import BaseModel, Field, PositiveInt, model_validator
-from typing_extensions import Self
from generalresearch.models.custom_types import UUIDStr
diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py
index c02acbe..0d7ace5 100644
--- a/generalresearch/models/thl/contest/__init__.py
+++ b/generalresearch/models/thl/contest/__init__.py
@@ -1,7 +1,7 @@
from __future__ import annotations
-from datetime import datetime, timezone
-from typing import Any
+from datetime import UTC, datetime
+from typing import Any, Self
from uuid import uuid4
from pydantic import (
@@ -11,7 +11,6 @@ from pydantic import (
computed_field,
model_validator,
)
-from typing_extensions import Self
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
@@ -107,7 +106,7 @@ class ContestWinner(BaseModel):
uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex)
created_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When this user won this prize",
)
diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py
index 6889dcb..6fc60f6 100644
--- a/generalresearch/models/thl/contest/contest.py
+++ b/generalresearch/models/thl/contest/contest.py
@@ -2,8 +2,8 @@ from __future__ import annotations
import json
from abc import ABC, abstractmethod
-from datetime import datetime, timezone
-from typing import Any
+from datetime import UTC, datetime
+from typing import Any, Self
from uuid import uuid4
from pydantic import (
@@ -14,7 +14,6 @@ from pydantic import (
NonNegativeInt,
model_validator,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
from generalresearch.models.thl.contest import (
@@ -57,7 +56,7 @@ class ContestBase(BaseModel, ABC):
starts_at: AwareDatetimeISO = Field(
description="When the contest starts",
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
)
terms_and_conditions: HttpUrl | None = Field(default=None)
@@ -78,6 +77,10 @@ class ContestBase(BaseModel, ABC):
self.model_config["validate_assignment"] = True
self.__class__.model_validate(self)
+ @classmethod
+ def example_json_schema_extra(cls, schema: dict[str, Any]) -> None:
+ schema["examples"] = [cls.example().model_dump(mode="json")]
+
class Contest(ContestBase):
id: int | None = Field(
@@ -91,11 +94,11 @@ class Contest(ContestBase):
product_id: UUIDStr = Field(description="Contest applies only to a single BP")
created_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When this contest was created",
)
updated_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When this contest was last modified. Does not include "
"entries being created/modified",
)
@@ -137,10 +140,12 @@ class Contest(ContestBase):
# return False
def should_end(self) -> tuple[bool, ContestEndReason | None]:
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.ends_at:
- if datetime.now(tz=timezone.utc) >= self.end_condition.ends_at:
- return True, ContestEndReason.ENDS_AT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.ends_at
+ and datetime.now(tz=UTC) >= self.end_condition.ends_at
+ ):
+ return True, ContestEndReason.ENDS_AT
return False, None
@@ -158,17 +163,17 @@ class Contest(ContestBase):
if winners is not None:
self.update(
status=ContestStatus.COMPLETED,
- ended_at=datetime.now(tz=timezone.utc),
+ ended_at=datetime.now(tz=UTC),
end_reason=reason,
all_winners=winners,
)
else:
self.update(
status=ContestStatus.COMPLETED,
- ended_at=datetime.now(tz=timezone.utc),
+ ended_at=datetime.now(tz=UTC),
end_reason=reason,
)
- return None
+ return
def model_dump_mysql(self, **kwargs) -> dict[str, Any]:
d = self.model_dump(mode="json", **kwargs)
@@ -185,7 +190,7 @@ class Contest(ContestBase):
@classmethod
def model_validate_mysql(cls, data: dict[str, Any]) -> Self:
- data = {k: v for k, v in data.items() if k in cls.model_fields.keys()}
+ data = {k: v for k, v in data.items() if k in cls.model_fields}
if isinstance(data["end_condition"], dict):
data["end_condition"] = ContestEndCondition.model_validate(
data["end_condition"]
@@ -211,7 +216,7 @@ class ContestUserView(Contest):
)
def is_user_eligible(self, country_iso: str) -> tuple[bool, str]:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
assert country_iso.lower() == country_iso
if now < self.starts_at:
diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py
index 2a9ecde..1a36fa9 100644
--- a/generalresearch/models/thl/contest/contest_entry.py
+++ b/generalresearch/models/thl/contest/contest_entry.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import Any
from uuid import uuid4
@@ -13,7 +13,9 @@ from pydantic import (
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-from generalresearch.models.thl.contest.definitions import ContestEntryType
+from generalresearch.models.thl.contest.definitions import (
+ ContestEntryType,
+)
from generalresearch.models.thl.user import User
@@ -40,12 +42,8 @@ class ContestEntryCreate(BaseModel):
class ContestEntry(BaseModel):
uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex)
- created_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(timezone.utc)
- )
- updated_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(timezone.utc)
- )
+ created_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC))
+ updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC))
# entry_type and amount are the same as on ContestEntryCreate
entry_type: ContestEntryType = Field()
@@ -57,22 +55,21 @@ class ContestEntry(BaseModel):
)
# user_id used internally, for DB joins/index
+ # todo: this should be a UserRef
user: User = Field(exclude=True)
@model_validator(mode="before")
@classmethod
- def validate_amount_type(cls, data: dict) -> dict:
- from generalresearch.models.thl.contest.definitions import (
- ContestEntryType,
- )
+ def validate_amount_type(cls, data: dict[str, Any]) -> dict[str, Any]:
amount = data.get("amount")
entry_type = data.get("entry_type")
if entry_type == ContestEntryType.COUNT:
- assert isinstance(amount, int) and not isinstance(
- amount, USDCent
- ), "amount must be int in ContestEntryType.COUNT"
+ assert isinstance(amount, int) and not isinstance(amount, USDCent), (
+ "amount must be int in ContestEntryType.COUNT"
+ )
+
elif entry_type == ContestEntryType.CASH:
# This may be coming from the DB, in which case it is an int.
data["amount"] = USDCent(data["amount"])
@@ -81,9 +78,6 @@ class ContestEntry(BaseModel):
@computed_field()
def amount_str(self) -> str:
- from generalresearch.models.thl.contest.definitions import (
- ContestEntryType,
- )
if self.entry_type == ContestEntryType.COUNT:
return str(self.amount)
diff --git a/generalresearch/models/thl/contest/definitions.py b/generalresearch/models/thl/contest/definitions.py
index 1a71408..cc0ee05 100644
--- a/generalresearch/models/thl/contest/definitions.py
+++ b/generalresearch/models/thl/contest/definitions.py
@@ -1,17 +1,17 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class ContestStatus(str, Enum):
+class ContestStatus(StrEnum):
ACTIVE = "active"
COMPLETED = "completed"
CANCELLED = "cancelled"
-class ContestType(str, Enum, metaclass=ReprEnumMeta):
+class ContestType(StrEnum, metaclass=ReprEnumMeta):
"""There are 3 contest types. They have a common base, with some unique
configurations and behaviors for each.
"""
@@ -26,7 +26,7 @@ class ContestType(str, Enum, metaclass=ReprEnumMeta):
MILESTONE = "milestone"
-class ContestEndReason(str, Enum):
+class ContestEndReason(StrEnum):
"""
Defines why a contest ended
"""
@@ -44,7 +44,7 @@ class ContestEndReason(str, Enum):
MAX_WINNERS = "max_winners"
-class ContestPrizeKind(str, Enum, metaclass=ReprEnumMeta):
+class ContestPrizeKind(StrEnum, metaclass=ReprEnumMeta):
# A physical prize (e.g. a iPhone, cash in the mail, dinner with Max)
PHYSICAL = "physical"
@@ -56,7 +56,7 @@ class ContestPrizeKind(str, Enum, metaclass=ReprEnumMeta):
CASH = "cash"
-class ContestEntryTrigger(str, Enum):
+class ContestEntryTrigger(StrEnum):
"""
Defines what action/event triggers a (possible) entry into the contest (automatically).
This only is valid on milestone contests
@@ -67,7 +67,7 @@ class ContestEntryTrigger(str, Enum):
REFERRAL = "referral"
-class ContestEntryType(str, Enum, metaclass=ReprEnumMeta):
+class ContestEntryType(StrEnum, metaclass=ReprEnumMeta):
"""
All entries into a contest must be of the same type, and match
the entry_type of the Contest itself.
@@ -83,7 +83,7 @@ class ContestEntryType(str, Enum, metaclass=ReprEnumMeta):
CASH = "cash"
-class LeaderboardTieBreakStrategy(str, Enum):
+class LeaderboardTieBreakStrategy(StrEnum):
"""
Strategies for resolving ties in leaderboard-based contests.
"""
diff --git a/generalresearch/models/thl/contest/examples.py b/generalresearch/models/thl/contest/examples.py
deleted file mode 100644
index 1810e63..0000000
--- a/generalresearch/models/thl/contest/examples.py
+++ /dev/null
@@ -1,416 +0,0 @@
-from __future__ import annotations
-
-from typing import Any
-
-from pydantic import HttpUrl
-
-from generalresearch.config import EXAMPLE_PRODUCT_ID
-from generalresearch.currency import USDCent
-
-
-def _example_raffle_create(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestEndCondition,
- ContestEntryRule,
- ContestPrize,
- )
- from generalresearch.models.thl.contest.contest_entry import (
- ContestEntryType,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.raffle import (
- RaffleContestCreate,
- )
-
- schema["example"] = RaffleContestCreate(
- name="Win an iPhone",
- description="iPhone winner will be drawn in proportion to entry "
- "amount. Contest ends once $800 has been entered.",
- contest_type=ContestType.RAFFLE,
- end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PHYSICAL,
- name="iPhone 16",
- estimated_cash_value=USDCent(800_00),
- )
- ],
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=None,
- entry_rule=ContestEntryRule(
- max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
- ),
- country_isos={"us", "ca"},
- entry_type=ContestEntryType.CASH,
- ).model_dump(mode="json")
-
-
-def _example_raffle(schema: dict) -> None:
- from generalresearch.models.thl.contest import (
- ContestEndCondition,
- ContestEntryRule,
- ContestPrize,
- )
- from generalresearch.models.thl.contest.contest_entry import (
- ContestEntryType,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestStatus,
- ContestType,
- )
- from generalresearch.models.thl.contest.raffle import RaffleContest
-
- schema["example"] = RaffleContest(
- name="Win an iPhone",
- description="iPhone winner will be drawn in proportion to entry "
- "amount. Contest ends once $800 has been entered.",
- contest_type=ContestType.RAFFLE,
- end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PHYSICAL,
- name="iPhone 16",
- estimated_cash_value=USDCent(800_00),
- )
- ],
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=None,
- entry_rule=ContestEntryRule(
- max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
- ),
- country_isos={"us", "ca"},
- entry_type=ContestEntryType.CASH,
- status=ContestStatus.ACTIVE,
- uuid="ce3968b8e18a4b96af62007f262ed7f7",
- created_at="2025-06-12T21:12:58.061205Z",
- updated_at="2025-06-12T21:12:58.061205Z",
- current_amount=4723,
- current_participants=12,
- product_id=EXAMPLE_PRODUCT_ID,
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_raffle_user_view(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestEndCondition,
- ContestEntryRule,
- ContestPrize,
- )
- from generalresearch.models.thl.contest.contest_entry import (
- ContestEntryType,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestStatus,
- ContestType,
- )
- from generalresearch.models.thl.contest.raffle import RaffleUserView
-
- schema["example"] = RaffleUserView(
- name="Win an iPhone",
- description="iPhone winner will be drawn in proportion to entry "
- "amount. Contest ends once $800 has been entered.",
- contest_type=ContestType.RAFFLE,
- end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PHYSICAL,
- name="iPhone 16",
- estimated_cash_value=USDCent(800_00),
- )
- ],
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=None,
- entry_rule=ContestEntryRule(
- max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
- ),
- country_isos={"us", "ca"},
- entry_type=ContestEntryType.CASH,
- status=ContestStatus.ACTIVE,
- uuid="ce3968b8e18a4b96af62007f262ed7f7",
- created_at="2025-06-12T21:12:58.061205Z",
- updated_at="2025-06-12T21:12:58.061205Z",
- current_amount=4723,
- current_participants=12,
- product_id=EXAMPLE_PRODUCT_ID,
- user_amount=420,
- user_amount_today=0,
- product_user_id="test-user",
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_milestone_create(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.milestone import (
- ContestEntryTrigger,
- MilestoneContestCreate,
- MilestoneContestEndCondition,
- )
-
- schema["example"] = MilestoneContestCreate(
- name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
- description="Only valid for the first 50 users",
- contest_type=ContestType.MILESTONE,
- end_condition=MilestoneContestEndCondition(max_winners=50),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PROMOTION,
- name="50% bonus on completes for 7 days",
- estimated_cash_value=USDCent(0),
- ),
- ContestPrize(
- kind=ContestPrizeKind.CASH,
- name="$5.00 Bonus",
- cash_amount=USDCent(5_00),
- estimated_cash_value=USDCent(5_00),
- ),
- ],
- entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
- target_amount=10,
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=HttpUrl("https://www.example.com"),
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_milestone(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.milestone import (
- ContestEntryTrigger,
- MilestoneContest,
- MilestoneContestEndCondition,
- )
-
- schema["example"] = MilestoneContest(
- name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
- description="Only valid for the first 50 users",
- contest_type=ContestType.MILESTONE,
- end_condition=MilestoneContestEndCondition(max_winners=50),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PROMOTION,
- name="50% bonus on completes for 7 days",
- estimated_cash_value=USDCent(0),
- ),
- ContestPrize(
- kind=ContestPrizeKind.CASH,
- name="$5.00 Bonus",
- cash_amount=USDCent(5_00),
- estimated_cash_value=USDCent(5_00),
- ),
- ],
- entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
- target_amount=10,
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=HttpUrl("https://www.example.com"),
- product_id=EXAMPLE_PRODUCT_ID,
- uuid="747fe3b709ae460e816821dcb81aebb9",
- created_at="2025-06-12T21:12:58.061205Z",
- updated_at="2025-06-12T21:12:58.061205Z",
- win_count=12,
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_milestone_user_view(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import ContestPrize
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.milestone import (
- ContestEntryTrigger,
- MilestoneContestEndCondition,
- MilestoneUserView,
- )
-
- schema["example"] = MilestoneUserView(
- name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
- description="Only valid for the first 50 users",
- contest_type=ContestType.MILESTONE,
- end_condition=MilestoneContestEndCondition(max_winners=50),
- prizes=[
- ContestPrize(
- kind=ContestPrizeKind.PROMOTION,
- name="50% bonus on completes for 7 days",
- estimated_cash_value=USDCent(0),
- ),
- ContestPrize(
- kind=ContestPrizeKind.CASH,
- name="$5.00 Bonus",
- cash_amount=USDCent(5_00),
- estimated_cash_value=USDCent(5_00),
- ),
- ],
- entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
- target_amount=10,
- starts_at="2025-06-12T21:12:58.061170Z",
- terms_and_conditions=HttpUrl("https://www.example.com"),
- product_id=EXAMPLE_PRODUCT_ID,
- uuid="747fe3b709ae460e816821dcb81aebb9",
- created_at="2025-06-12T21:12:58.061205Z",
- updated_at="2025-06-12T21:12:58.061205Z",
- win_count=12,
- user_amount=8,
- product_user_id="test-user",
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.leaderboard import (
- LeaderboardContestCreate,
- )
-
- schema["example"] = LeaderboardContestCreate(
- name="Prizes for top survey takers this week",
- description="$15 1st place, $10 2nd, $5 3rd place US weekly",
- contest_type=ContestType.LEADERBOARD,
- prizes=[
- ContestPrize(
- name="$15 Cash",
- estimated_cash_value=USDCent(15_00),
- cash_amount=USDCent(15_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=1,
- ),
- ContestPrize(
- name="$10 Cash",
- estimated_cash_value=USDCent(10_00),
- cash_amount=USDCent(10_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=2,
- ),
- ContestPrize(
- name="$5 Cash",
- estimated_cash_value=USDCent(5_00),
- cash_amount=USDCent(5_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=3,
- ),
- ],
- leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count",
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_leaderboard_contest(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.leaderboard import (
- LeaderboardContest,
- )
-
- schema["example"] = LeaderboardContest(
- name="Prizes for top survey takers this week",
- description="$15 1st place, $10 2nd, $5 3rd place US weekly",
- contest_type=ContestType.LEADERBOARD,
- prizes=[
- ContestPrize(
- name="$15 Cash",
- estimated_cash_value=USDCent(15_00),
- cash_amount=USDCent(15_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=1,
- ),
- ContestPrize(
- name="$10 Cash",
- estimated_cash_value=USDCent(10_00),
- cash_amount=USDCent(10_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=2,
- ),
- ContestPrize(
- name="$5 Cash",
- estimated_cash_value=USDCent(5_00),
- cash_amount=USDCent(5_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=3,
- ),
- ],
- leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count",
- product_id=EXAMPLE_PRODUCT_ID,
- ).model_dump(mode="json")
-
- return None
-
-
-def _example_leaderboard_contest_user_view(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.contest import (
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.leaderboard import (
- LeaderboardContestUserView,
- )
-
- schema["example"] = LeaderboardContestUserView(
- name="Prizes for top survey takers this week",
- description="$15 1st place, $10 2nd, $5 3rd place US weekly",
- contest_type=ContestType.LEADERBOARD,
- prizes=[
- ContestPrize(
- name="$15 Cash",
- estimated_cash_value=USDCent(15_00),
- cash_amount=USDCent(15_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=1,
- ),
- ContestPrize(
- name="$10 Cash",
- estimated_cash_value=USDCent(10_00),
- cash_amount=USDCent(10_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=2,
- ),
- ContestPrize(
- name="$5 Cash",
- estimated_cash_value=USDCent(5_00),
- cash_amount=USDCent(5_00),
- kind=ContestPrizeKind.CASH,
- leaderboard_rank=3,
- ),
- ],
- leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count",
- product_id=EXAMPLE_PRODUCT_ID,
- product_user_id="test-user",
- ).model_dump(mode="json")
diff --git a/generalresearch/models/thl/contest/io.py b/generalresearch/models/thl/contest/io.py
index e68f76e..d11080f 100644
--- a/generalresearch/models/thl/contest/io.py
+++ b/generalresearch/models/thl/contest/io.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from uuid import uuid4
from generalresearch.models.thl.contest.definitions import ContestType
@@ -37,7 +37,7 @@ from generalresearch.models.thl.contest.contest import Contest
def contest_create_to_contest(
product_id: str, contest_create: ContestCreate
) -> Contest:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
d = contest_create.model_dump(mode="json")
d["uuid"] = uuid4().hex
d["product_id"] = product_id
diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py
index 8167f46..064c0d1 100644
--- a/generalresearch/models/thl/contest/leaderboard.py
+++ b/generalresearch/models/thl/contest/leaderboard.py
@@ -1,7 +1,7 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
-from typing import Any, Literal
+from datetime import UTC, datetime, timedelta
+from typing import Any, Literal, Self
from pydantic import (
ConfigDict,
@@ -11,8 +11,8 @@ from pydantic import (
model_validator,
)
from redis import Redis
-from typing_extensions import Self
+from generalresearch.currency import USDCent
from generalresearch.decorators import LOG
from generalresearch.managers.leaderboard import country_timezone
from generalresearch.managers.leaderboard.manager import LeaderboardManager
@@ -21,6 +21,7 @@ from generalresearch.managers.thl.user_manager.user_manager import (
)
from generalresearch.models.thl.contest import (
ContestEndCondition,
+ ContestPrize,
ContestWinner,
)
from generalresearch.models.thl.contest.contest import (
@@ -35,11 +36,6 @@ from generalresearch.models.thl.contest.definitions import (
ContestType,
LeaderboardTieBreakStrategy,
)
-from generalresearch.models.thl.contest.examples import (
- _example_leaderboard_contest,
- _example_leaderboard_contest_create,
- _example_leaderboard_contest_user_view,
-)
from generalresearch.models.thl.leaderboard import (
Leaderboard,
LeaderboardCode,
@@ -51,7 +47,6 @@ class LeaderboardContestCreate(ContestBase):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_leaderboard_contest_create,
)
contest_type: Literal[ContestType.LEADERBOARD] = Field(
@@ -78,9 +73,9 @@ class LeaderboardContestCreate(ContestBase):
ranks = {x.leaderboard_rank for x in self.prizes}
assert None not in ranks, "Must have leaderboard_rank defined"
assert min(ranks) == 1, "Must start with rank 1"
- assert ranks == set(
- range(min(ranks), max(ranks) + 1)
- ), "cannot skip prize leaderboard_ranks"
+ assert ranks == set(range(min(ranks), max(ranks) + 1)), (
+ "cannot skip prize leaderboard_ranks"
+ )
return self
@model_validator(mode="after")
@@ -91,9 +86,9 @@ class LeaderboardContestCreate(ContestBase):
@model_validator(mode="after")
def check_end_condition(self) -> Self:
- assert (
- not self.end_condition.target_entry_amount
- ), "target_entry_amount not valid in leaderboard contest"
+ assert not self.end_condition.target_entry_amount, (
+ "target_entry_amount not valid in leaderboard contest"
+ )
# the ends_at will get set automatically from the leaderboard_key
return self
@@ -125,12 +120,45 @@ class LeaderboardContestCreate(ContestBase):
parts | {"row_count": 0, "bpid": parts["product_id"]}
)
+ @classmethod
+ def example(cls) -> LeaderboardContestCreate:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+
+ return cls(
+ name="Prizes for top survey takers this week",
+ description="$15 1st place, $10 2nd, $5 3rd place US weekly",
+ contest_type=ContestType.LEADERBOARD,
+ prizes=[
+ ContestPrize(
+ name="$15 Cash",
+ estimated_cash_value=USDCent(15_00),
+ cash_amount=USDCent(15_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=1,
+ ),
+ ContestPrize(
+ name="$10 Cash",
+ estimated_cash_value=USDCent(10_00),
+ cash_amount=USDCent(10_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=2,
+ ),
+ ContestPrize(
+ name="$5 Cash",
+ estimated_cash_value=USDCent(5_00),
+ cash_amount=USDCent(5_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=3,
+ ),
+ ],
+ leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count",
+ )
+
class LeaderboardContest(LeaderboardContestCreate, Contest):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_leaderboard_contest,
arbitrary_types_allowed=True,
)
@@ -144,15 +172,16 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
@model_validator(mode="after")
def validate_product_lb_key(self) -> Self:
- assert (
- self.product_id == self.leaderboard_key_parts["product_id"]
- ), "leaderboard_key product_id is invalid"
+ assert self.product_id == self.leaderboard_key_parts["product_id"], (
+ "leaderboard_key product_id is invalid"
+ )
if self.country_isos:
+ assert len(self.country_isos) == 1, (
+ "Can only set 1 country_iso in a leaderboard contest"
+ )
assert (
- len(self.country_isos) == 1
- ), "Can only set 1 country_iso in a leaderboard contest"
- assert (
- list(self.country_isos)[0] == self.leaderboard_key_parts["country_iso"]
+ next(iter(self.country_isos))
+ == self.leaderboard_key_parts["country_iso"]
), "leaderboard_key country_iso must match the country_isos"
else:
self.country_isos = {self.leaderboard_key_parts["country_iso"]}
@@ -161,9 +190,9 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
@model_validator(mode="after")
def validate_tie_break(self) -> Self:
if self.tie_break_strategy == LeaderboardTieBreakStrategy.SPLIT_PRIZE_POOL:
- assert all(
- p.kind == ContestPrizeKind.CASH for p in self.prizes
- ), "All prizes must be cash due to the tie-break strategy"
+ assert all(p.kind == ContestPrizeKind.CASH for p in self.prizes), (
+ "All prizes must be cash due to the tie-break strategy"
+ )
return self
@model_validator(mode="after")
@@ -193,10 +222,12 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
return lbm
def should_end(self) -> tuple[bool, ContestEndReason | None]:
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.ends_at:
- if datetime.now(tz=timezone.utc) >= self.end_condition.ends_at:
- return True, ContestEndReason.ENDS_AT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.ends_at
+ and datetime.now(tz=UTC) >= self.end_condition.ends_at
+ ):
+ return True, ContestEndReason.ENDS_AT
return False, None
@@ -244,12 +275,46 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
)
return d
+ @classmethod
+ def example(cls) -> LeaderboardContest:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+
+ return cls(
+ name="Prizes for top survey takers this week",
+ description="$15 1st place, $10 2nd, $5 3rd place US weekly",
+ contest_type=ContestType.LEADERBOARD,
+ prizes=[
+ ContestPrize(
+ name="$15 Cash",
+ estimated_cash_value=USDCent(15_00),
+ cash_amount=USDCent(15_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=1,
+ ),
+ ContestPrize(
+ name="$10 Cash",
+ estimated_cash_value=USDCent(10_00),
+ cash_amount=USDCent(10_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=2,
+ ),
+ ContestPrize(
+ name="$5 Cash",
+ estimated_cash_value=USDCent(5_00),
+ cash_amount=USDCent(5_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=3,
+ ),
+ ],
+ leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count",
+ product_id=product_id,
+ )
+
class LeaderboardContestUserView(LeaderboardContest, ContestUserView):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_leaderboard_contest_user_view,
)
@computed_field(description="The current rank of this user in this contest")
@@ -276,7 +341,7 @@ class LeaderboardContestUserView(LeaderboardContest, ContestUserView):
if self.user_winnings:
return False, "User already won"
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
if self.leaderboard_model.period_end_utc < now:
return False, "Contest is over"
if self.leaderboard_model.period_start_utc > now:
@@ -289,3 +354,39 @@ class LeaderboardContestUserView(LeaderboardContest, ContestUserView):
return False, "contest is over"
return True, ""
+
+ @classmethod
+ def example(cls) -> LeaderboardContestUserView:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+
+ return cls(
+ name="Prizes for top survey takers this week",
+ description="$15 1st place, $10 2nd, $5 3rd place US weekly",
+ contest_type=ContestType.LEADERBOARD,
+ prizes=[
+ ContestPrize(
+ name="$15 Cash",
+ estimated_cash_value=USDCent(15_00),
+ cash_amount=USDCent(15_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=1,
+ ),
+ ContestPrize(
+ name="$10 Cash",
+ estimated_cash_value=USDCent(10_00),
+ cash_amount=USDCent(10_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=2,
+ ),
+ ContestPrize(
+ name="$5 Cash",
+ estimated_cash_value=USDCent(5_00),
+ cash_amount=USDCent(5_00),
+ kind=ContestPrizeKind.CASH,
+ leaderboard_rank=3,
+ ),
+ ],
+ leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count",
+ product_id=product_id,
+ product_user_id="test-user",
+ )
diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py
index f62be8f..db5ba2f 100644
--- a/generalresearch/models/thl/contest/milestone.py
+++ b/generalresearch/models/thl/contest/milestone.py
@@ -2,17 +2,21 @@ from __future__ import annotations
import logging
from datetime import timedelta
-from typing import Any, Literal
+from typing import Any, Literal, Self
from pydantic import (
BaseModel,
ConfigDict,
Field,
+ HttpUrl,
PositiveInt,
)
-from typing_extensions import Self
+from generalresearch.currency import USDCent
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.thl.contest import (
+ ContestPrize,
+)
from generalresearch.models.thl.contest.contest import (
Contest,
ContestBase,
@@ -23,14 +27,10 @@ from generalresearch.models.thl.contest.definitions import (
ContestEndReason,
ContestEntryTrigger,
ContestEntryType,
+ ContestPrizeKind,
ContestStatus,
ContestType,
)
-from generalresearch.models.thl.contest.examples import (
- _example_milestone,
- _example_milestone_create,
- _example_milestone_user_view,
-)
logging.basicConfig()
LOG = logging.getLogger()
@@ -104,19 +104,46 @@ class MilestoneContestCreate(ContestBase, MilestoneContestConfig):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_milestone_create,
+ # json_schema_extra=json_example_milestone_create,
)
contest_type: Literal[ContestType.MILESTONE] = Field(default=ContestType.MILESTONE)
end_condition: MilestoneContestEndCondition = Field()
+ @classmethod
+ def example(cls) -> MilestoneContestCreate:
+
+ return cls(
+ name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
+ description="Only valid for the first 50 users",
+ contest_type=ContestType.MILESTONE,
+ end_condition=MilestoneContestEndCondition(max_winners=50),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PROMOTION,
+ name="50% bonus on completes for 7 days",
+ estimated_cash_value=USDCent(0),
+ ),
+ ContestPrize(
+ kind=ContestPrizeKind.CASH,
+ name="$5.00 Bonus",
+ cash_amount=USDCent(5_00),
+ estimated_cash_value=USDCent(5_00),
+ ),
+ ],
+ entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
+ target_amount=10,
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=HttpUrl("https://www.example.com"),
+ )
+
class MilestoneContest(MilestoneContestCreate, Contest):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_milestone,
+ # json_schema_extra=json_example_milestone,
)
entry_type: Literal[ContestEntryType.COUNT] = Field(default=ContestEntryType.COUNT)
@@ -133,10 +160,12 @@ class MilestoneContest(MilestoneContestCreate, Contest):
if res:
return res, msg
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.max_winners:
- if self.win_count >= self.end_condition.max_winners:
- return True, ContestEndReason.MAX_WINNERS
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.max_winners
+ and self.win_count >= self.end_condition.max_winners
+ ):
+ return True, ContestEndReason.MAX_WINNERS
return False, None
@@ -172,12 +201,43 @@ class MilestoneContest(MilestoneContestCreate, Contest):
)
return super().model_validate_mysql(data)
+ @classmethod
+ def example(cls) -> MilestoneContest:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+ return cls(
+ name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
+ description="Only valid for the first 50 users",
+ contest_type=ContestType.MILESTONE,
+ end_condition=MilestoneContestEndCondition(max_winners=50),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PROMOTION,
+ name="50% bonus on completes for 7 days",
+ estimated_cash_value=USDCent(0),
+ ),
+ ContestPrize(
+ kind=ContestPrizeKind.CASH,
+ name="$5.00 Bonus",
+ cash_amount=USDCent(5_00),
+ estimated_cash_value=USDCent(5_00),
+ ),
+ ],
+ entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
+ target_amount=10,
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=HttpUrl("https://www.example.com"),
+ product_id=product_id,
+ uuid="747fe3b709ae460e816821dcb81aebb9",
+ created_at="2025-06-12T21:12:58.061205Z",
+ updated_at="2025-06-12T21:12:58.061205Z",
+ win_count=12,
+ )
+
class MilestoneUserView(MilestoneContest, ContestUserView):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_milestone_user_view,
)
valid_until: AwareDatetimeISO | None = Field(
@@ -190,16 +250,10 @@ class MilestoneUserView(MilestoneContest, ContestUserView):
)
def should_award(self):
- if self.status == ContestStatus.ACTIVE:
- if self.should_have_awarded():
- return True
- return False
+ return bool(self.status == ContestStatus.ACTIVE and self.should_have_awarded())
def should_have_awarded(self):
- if self.target_amount:
- if self.user_amount >= self.target_amount:
- return True
- return False
+ return bool(self.target_amount and self.user_amount >= self.target_amount)
def is_user_eligible(self, country_iso: str) -> tuple[bool, str]:
passes, msg = super().is_user_eligible(country_iso=country_iso)
@@ -223,3 +277,37 @@ class MilestoneUserView(MilestoneContest, ContestUserView):
# TODO: others in self.entry_rule ... min_completes, id_verified, etc.
return True, ""
+
+ @classmethod
+ def example(cls) -> MilestoneUserView:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+ return cls(
+ name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!",
+ description="Only valid for the first 50 users",
+ contest_type=ContestType.MILESTONE,
+ end_condition=MilestoneContestEndCondition(max_winners=50),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PROMOTION,
+ name="50% bonus on completes for 7 days",
+ estimated_cash_value=USDCent(0),
+ ),
+ ContestPrize(
+ kind=ContestPrizeKind.CASH,
+ name="$5.00 Bonus",
+ cash_amount=USDCent(5_00),
+ estimated_cash_value=USDCent(5_00),
+ ),
+ ],
+ entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
+ target_amount=10,
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=HttpUrl("https://www.example.com"),
+ product_id=product_id,
+ uuid="747fe3b709ae460e816821dcb81aebb9",
+ created_at="2025-06-12T21:12:58.061205Z",
+ updated_at="2025-06-12T21:12:58.061205Z",
+ win_count=12,
+ user_amount=8,
+ product_user_id="test-user",
+ )
diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py
index d3157e3..7e84e89 100644
--- a/generalresearch/models/thl/contest/raffle.py
+++ b/generalresearch/models/thl/contest/raffle.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import logging
import random
from collections import defaultdict
-from datetime import datetime, timezone
-from typing import Any, Literal
+from datetime import UTC, datetime
+from typing import Any, Literal, Self
from pydantic import (
ConfigDict,
@@ -14,11 +14,12 @@ from pydantic import (
model_validator,
)
from scipy.stats import hypergeom
-from typing_extensions import Self
from generalresearch.currency import USDCent
from generalresearch.models.thl.contest import (
+ ContestEndCondition,
ContestEntryRule,
+ ContestPrize,
ContestWinner,
)
from generalresearch.models.thl.contest.contest import (
@@ -26,18 +27,16 @@ from generalresearch.models.thl.contest.contest import (
ContestBase,
ContestUserView,
)
-from generalresearch.models.thl.contest.contest_entry import ContestEntry
+from generalresearch.models.thl.contest.contest_entry import (
+ ContestEntry,
+)
from generalresearch.models.thl.contest.definitions import (
ContestEndReason,
ContestEntryType,
+ ContestPrizeKind,
ContestStatus,
ContestType,
)
-from generalresearch.models.thl.contest.examples import (
- _example_raffle,
- _example_raffle_create,
- _example_raffle_user_view,
-)
logging.basicConfig()
LOG = logging.getLogger()
@@ -48,7 +47,7 @@ class RaffleContestCreate(ContestBase):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_raffle_create,
+ # json_schema_extra=json_example_raffle_create,
)
contest_type: Literal[ContestType.RAFFLE] = Field(default=ContestType.RAFFLE)
@@ -64,12 +63,36 @@ class RaffleContestCreate(ContestBase):
raise ValueError("At least one end condition must be specified")
return self
+ @classmethod
+ def example(cls) -> RaffleContestCreate:
+ return cls(
+ name="Win an iPhone",
+ description="iPhone winner will be drawn in proportion to entry "
+ "amount. Contest ends once $800 has been entered.",
+ contest_type=ContestType.RAFFLE,
+ end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PHYSICAL,
+ name="iPhone 16",
+ estimated_cash_value=USDCent(800_00),
+ )
+ ],
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=None,
+ entry_rule=ContestEntryRule(
+ max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
+ ),
+ country_isos={"us", "ca"},
+ entry_type=ContestEntryType.CASH,
+ )
+
class RaffleContest(RaffleContestCreate, Contest):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_raffle,
+ # json_schema_extra=json_example_raffle,
)
entries: list[ContestEntry] = Field(default_factory=list, exclude=True)
@@ -87,19 +110,22 @@ class RaffleContest(RaffleContestCreate, Contest):
@model_validator(mode="after")
def validate_entry_type(self):
- assert all(
- entry.entry_type == self.entry_type for entry in self.entries
- ), f"all entries must be of type {self.entry_type}"
+ assert all(entry.entry_type == self.entry_type for entry in self.entries), (
+ f"all entries must be of type {self.entry_type}"
+ )
return self
@field_validator("current_amount", mode="before")
- def coerce_current_amount(cls, v, info):
+ def coerce_current_amount(cls, v: int | USDCent, info):
if v is None:
return None
+
if info.data.get("entry_type") == ContestEntryType.CASH:
return USDCent(v)
+
elif info.data.get("entry_type") == ContestEntryType.COUNT:
return int(v)
+
return v
@model_validator(mode="after")
@@ -128,7 +154,7 @@ class RaffleContest(RaffleContestCreate, Contest):
# If there is more than 1 prize, the winning entry is subtracted
# from the user's entry count
user_amount = defaultdict(int)
- user_id_user = dict()
+ user_id_user = {}
for entry in self.entries:
user_amount[entry.user.user_id] += entry.amount
user_id_user[entry.user.user_id] = entry.user
@@ -150,10 +176,12 @@ class RaffleContest(RaffleContestCreate, Contest):
res, msg = super().should_end()
if res:
return res, msg
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.target_entry_amount:
- if self.current_amount >= self.end_condition.target_entry_amount:
- return True, ContestEndReason.TARGET_ENTRY_AMOUNT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.target_entry_amount
+ and self.current_amount >= self.end_condition.target_entry_amount
+ ):
+ return True, ContestEndReason.TARGET_ENTRY_AMOUNT
return False, None
@staticmethod
@@ -202,9 +230,7 @@ class RaffleContest(RaffleContestCreate, Contest):
c = self.end_condition
if c.target_entry_amount and self.current_amount >= c.target_entry_amount:
return True
- if c.ends_at and datetime.now(tz=timezone.utc) >= c.ends_at:
- return True
- return False
+ return bool(c.ends_at and datetime.now(tz=UTC) >= c.ends_at)
def model_dump_mysql(self) -> dict[str, Any]:
d = super().model_dump_mysql()
@@ -212,16 +238,48 @@ class RaffleContest(RaffleContestCreate, Contest):
return d
@classmethod
- def model_validate_mysql(cls, data: dict) -> Self:
+ def model_validate_mysql(cls, data: dict[str, Any]) -> Self:
data["entry_rule"] = ContestEntryRule.model_validate(data["entry_rule"])
return super().model_validate_mysql(data)
+ @classmethod
+ def example(cls) -> RaffleContest:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+ return cls(
+ name="Win an iPhone",
+ description="iPhone winner will be drawn in proportion to entry "
+ "amount. Contest ends once $800 has been entered.",
+ contest_type=ContestType.RAFFLE,
+ end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PHYSICAL,
+ name="iPhone 16",
+ estimated_cash_value=USDCent(800_00),
+ )
+ ],
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=None,
+ entry_rule=ContestEntryRule(
+ max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
+ ),
+ country_isos={"us", "ca"},
+ entry_type=ContestEntryType.CASH,
+ status=ContestStatus.ACTIVE,
+ uuid="ce3968b8e18a4b96af62007f262ed7f7",
+ created_at="2025-06-12T21:12:58.061205Z",
+ updated_at="2025-06-12T21:12:58.061205Z",
+ current_amount=4723,
+ current_participants=12,
+ product_id=product_id,
+ )
+
class RaffleUserView(RaffleContest, ContestUserView):
model_config = ConfigDict(
validate_assignment=True,
extra="forbid",
- json_schema_extra=_example_raffle_user_view,
+ # json_schema_extra=json_example_raffle_user_view,
)
user_amount: int | USDCent = Field(
@@ -279,17 +337,19 @@ class RaffleUserView(RaffleContest, ContestUserView):
return probs
def is_entry_eligible(self, entry: ContestEntry) -> tuple[bool, str]:
- if self.entry_rule.max_entry_amount_per_user:
- if (
- self.user_amount + entry.amount
- ) > self.entry_rule.max_entry_amount_per_user:
- return False, "Entry would exceed max amount per user."
-
- if self.entry_rule.max_daily_entries_per_user:
- if (
- self.user_amount_today + entry.amount
- ) > self.entry_rule.max_daily_entries_per_user:
- return False, "Entry would exceed max amount per user per day."
+ if (
+ self.entry_rule.max_entry_amount_per_user
+ and (self.user_amount + entry.amount)
+ > self.entry_rule.max_entry_amount_per_user
+ ):
+ return False, "Entry would exceed max amount per user."
+
+ if (
+ self.entry_rule.max_daily_entries_per_user
+ and (self.user_amount_today + entry.amount)
+ > self.entry_rule.max_daily_entries_per_user
+ ):
+ return False, "Entry would exceed max amount per user per day."
return True, ""
def is_user_eligible(self, country_iso: str) -> tuple[bool, str]:
@@ -297,16 +357,18 @@ class RaffleUserView(RaffleContest, ContestUserView):
if not passes:
return False, msg
- if self.entry_rule.max_entry_amount_per_user:
+ if self.entry_rule.max_entry_amount_per_user: # noqa: SIM102
# Greater or equal b/c we're asking if the user is eligible to
# enter MORE, now! If it equals, nothing is wrong, just that they
# are not eligible anymore.
if self.user_amount >= self.entry_rule.max_entry_amount_per_user:
return False, "Reached max amount per user."
- if self.entry_rule.max_daily_entries_per_user:
- if self.user_amount_today >= self.entry_rule.max_daily_entries_per_user:
- return False, "Reached max amount today."
+ if (
+ self.entry_rule.max_daily_entries_per_user
+ and self.user_amount_today >= self.entry_rule.max_daily_entries_per_user
+ ):
+ return False, "Reached max amount today."
# This would indicate something is wrong, as something else should have done this
e, _ = self.should_end()
@@ -316,3 +378,38 @@ class RaffleUserView(RaffleContest, ContestUserView):
# todo: others in self.entry_rule ... min_completes, id_verified, etc.
return True, ""
+
+ @classmethod
+ def example(cls) -> RaffleUserView:
+ product_id = "1108d053e4fa47c5b0dbdcd03a7981e7"
+ return cls(
+ name="Win an iPhone",
+ description="iPhone winner will be drawn in proportion to entry "
+ "amount. Contest ends once $800 has been entered.",
+ contest_type=ContestType.RAFFLE,
+ end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)),
+ prizes=[
+ ContestPrize(
+ kind=ContestPrizeKind.PHYSICAL,
+ name="iPhone 16",
+ estimated_cash_value=USDCent(800_00),
+ )
+ ],
+ starts_at="2025-06-12T21:12:58.061170Z",
+ terms_and_conditions=None,
+ entry_rule=ContestEntryRule(
+ max_entry_amount_per_user=10000, max_daily_entries_per_user=1000
+ ),
+ country_isos={"us", "ca"},
+ entry_type=ContestEntryType.CASH,
+ status=ContestStatus.ACTIVE,
+ uuid="ce3968b8e18a4b96af62007f262ed7f7",
+ created_at="2025-06-12T21:12:58.061205Z",
+ updated_at="2025-06-12T21:12:58.061205Z",
+ current_amount=4723,
+ current_participants=12,
+ product_id=product_id,
+ user_amount=420,
+ user_amount_today=0,
+ product_user_id="test-user",
+ )
diff --git a/generalresearch/models/thl/definitions.py b/generalresearch/models/thl/definitions.py
index 22ca4b9..c40df21 100644
--- a/generalresearch/models/thl/definitions.py
+++ b/generalresearch/models/thl/definitions.py
@@ -1,10 +1,10 @@
import copy
-from enum import Enum
+from enum import IntEnum, StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class ReservedQueryParameters(str, Enum, metaclass=ReprEnumMeta):
+class ReservedQueryParameters(StrEnum, metaclass=ReprEnumMeta):
PRODUCT_ID = "product_id"
PRODUCT_USER_ID = "bp_user_id"
BPUID = "bpuid"
@@ -48,7 +48,7 @@ class ReservedQueryParameters(str, Enum, metaclass=ReprEnumMeta):
N_BINS = "n_bins"
-class THLPaths(str, Enum, metaclass=ReprEnumMeta):
+class THLPaths(StrEnum, metaclass=ReprEnumMeta):
# Endpoints on thl-fsb
TASK_ADJUSTMENT = "f4484dbdf144451ab60cda256ce14266"
@@ -65,7 +65,7 @@ class THLPaths(str, Enum, metaclass=ReprEnumMeta):
GET_GRLIQ_JS_ATTR = "4a2954b34cc24f93be3e8b218e323b88"
-class Status(str, Enum, metaclass=ReprEnumMeta):
+class Status(StrEnum, metaclass=ReprEnumMeta):
"""
The outcome of a session or wall event. If the session is still in
progress, the status will be NULL.
@@ -85,7 +85,7 @@ class Status(str, Enum, metaclass=ReprEnumMeta):
TIMEOUT = "t"
-class WallAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class WallAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
# Task was reconciled to complete
ADJUSTED_TO_COMPLETE = "ac"
# Task was reconciled to incomplete
@@ -100,7 +100,7 @@ class WallAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
CONFIRMED_COMPLETE = "cc"
-class SessionAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class SessionAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
"""An adjusted_status is set if a session is adjusted by the marketplace
after the original return. A session can be adjusted multiple times.
This is the most recent status. If a session was originally a complete,
@@ -117,7 +117,7 @@ class SessionAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
PAYOUT_ADJUSTMENT = "pa"
-class StatusCode1(int, Enum, metaclass=ReprEnumMeta):
+class StatusCode1(IntEnum, metaclass=ReprEnumMeta):
"""
__High level status code for outcome of the session.__
This should only be NULL if the Status is ABANDON or TIMEOUT
@@ -177,7 +177,7 @@ class StatusCode1(int, Enum, metaclass=ReprEnumMeta):
SESSION_CONTINUE_QUALITY_FAIL = 19
-class SessionStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class SessionStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
__Status Detail__
This should be set if the Session.status_code_1 is SESSION_XXX_FAIL
@@ -185,7 +185,7 @@ class SessionStatusCode2(int, Enum, metaclass=ReprEnumMeta):
# Unable to parse either the bucket_id, request_id, or nudge_id from the url
ENTRY_URL_MODIFICATION = 1
- # The client's IP failed maxmind lookup, or we failed to store it for some reason
+ # The client's IP failed GRIP lookup, or we failed to store it for some reason
UNRECOGNIZED_IP = 2
# User is using an anonymous IP
USER_IS_ANONYMOUS = 3
@@ -217,7 +217,7 @@ class SessionStatusCode2(int, Enum, metaclass=ReprEnumMeta):
GRLIQ_MISSING = 13
-class WallStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class WallStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
This should be set if the Wall.status_code_1 is MARKETPLACE_FAIL
"""
@@ -289,7 +289,7 @@ WALL_ALLOWED_STATUS_CODE_1_2 = {
}
-class ReportValue(int, Enum, metaclass=ReprEnumMeta):
+class ReportValue(IntEnum, metaclass=ReprEnumMeta):
"""
The reason a user reported a task.
"""
@@ -316,7 +316,7 @@ class ReportValue(int, Enum, metaclass=ReprEnumMeta):
DIDNT_LIKE = 7
-class PayoutStatus(str, Enum, metaclass=ReprEnumMeta):
+class PayoutStatus(StrEnum, metaclass=ReprEnumMeta):
"""The max size of the db field that holds this value is 20, so please
don't add new values longer than that!
"""
diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py
index 4d4c8c2..ce4939c 100644
--- a/generalresearch/models/thl/demographics.py
+++ b/generalresearch/models/thl/demographics.py
@@ -3,14 +3,13 @@ from __future__ import annotations
import copy
from collections import Counter, defaultdict
from dataclasses import dataclass
-from enum import Enum
+from enum import Enum, StrEnum
from typing import TYPE_CHECKING, Any, Literal
import numpy as np
-from generalresearch.models.thl.locales import CountryISO
-
if TYPE_CHECKING:
+ from generalresearch.models.thl.locales import CountryISO
from generalresearch.models.thl.survey import MarketplaceTask
@@ -35,7 +34,7 @@ class DemographicTarget:
}
-class Gender(str, Enum):
+class Gender(StrEnum):
"""
The respondent's gender
"""
@@ -76,7 +75,7 @@ class AgeGroup(Enum):
return self.label
-def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list:
+def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list[dict[str, Any]]:
"""
Measurement: marketplace_survey_demographics
tags: source (marketplace)
@@ -86,7 +85,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list:
"""
source = {opp.source for opp in opps}
assert len(source) == 1
- source = list(source)[0]
+ source = next(iter(source))
survey_cpi = defaultdict(list)
target_open = defaultdict(int)
for opp in opps:
@@ -100,7 +99,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list:
survey_counter = {k: len(v) for k, v in survey_cpi.items()}
survey_counter = {k: {"count": v} for k, v in survey_counter.items() if v}
- grp_stats = dict()
+ grp_stats = {}
for grp, costs in survey_cpi.items():
stats = {
"cost_min": np.min(costs),
@@ -155,7 +154,7 @@ def calculate_used_question_metrics(
"""
source = {opp.source for opp in opps}
assert len(source) == 1
- source = list(source)[0]
+ source = next(iter(source))
country_q_counter = defaultdict(Counter)
for opp in opps:
for q in opp.used_question_ids:
diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py
index b72ecf6..f012cbf 100644
--- a/generalresearch/models/thl/finance.py
+++ b/generalresearch/models/thl/finance.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import random
-from datetime import timezone
+from datetime import UTC
from typing import TYPE_CHECKING
from uuid import uuid4
@@ -16,20 +16,21 @@ from pydantic import (
model_validator,
)
from pydantic.json_schema import SkipJsonSchema
-from generalresearch.config import is_debug
+from generalresearch.config import is_debug
from generalresearch.currency import USDCent
from generalresearch.decorators import LOG
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
from generalresearch.models.thl.definitions import SessionAdjustedStatus
-from generalresearch.pg_helper import PostgresConfig
payout_example = random.randint(150, 750 * 100)
adjustment_example = random.randint(-1_000, 50 * 100)
if TYPE_CHECKING:
+
from generalresearch.managers.thl.product import ProductManager
- from generalresearch.models.thl.ledger import AccountType, Direction, LedgerAccount
+ from generalresearch.models.thl.ledger import LedgerAccount
+ from generalresearch.models.thl.product import Product
class AdjustmentType(BaseModel):
@@ -125,10 +126,10 @@ class POPFinancial(BaseModel):
Direction,
)
- assert all([a.account_type == AccountType.BP_WALLET for a in accounts])
- assert all([a.normal_balance == Direction.CREDIT for a in accounts])
+ assert all(a.account_type == AccountType.BP_WALLET for a in accounts)
+ assert all(a.normal_balance == Direction.CREDIT for a in accounts)
if not is_debug():
- assert all([a.currency == "USD" for a in accounts])
+ assert all(a.currency == "USD" for a in accounts)
if input_data.empty:
return []
@@ -150,7 +151,7 @@ class POPFinancial(BaseModel):
index: int # Not useful, just a RangeIndex
row: pd.DataFrame
- row["time_idx"] = row.time_idx.to_pydatetime().replace(tzinfo=timezone.utc)
+ row["time_idx"] = row.time_idx.to_pydatetime().replace(tzinfo=UTC)
instance = ProductBalances.from_pandas(row)
res.append(
@@ -324,8 +325,6 @@ class ProductBalances(BaseModel):
)
@property
def payout_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.payout).to_usd_str()
@computed_field(
@@ -384,8 +383,6 @@ class ProductBalances(BaseModel):
)
@property
def payment_usd_str(self):
- from generalresearch.currency import USDCent
-
return USDCent(self.payment).to_usd_str()
@computed_field(
@@ -427,8 +424,6 @@ class ProductBalances(BaseModel):
)
@property
def retainer_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.retainer).to_usd_str()
@computed_field(
@@ -460,8 +455,6 @@ class ProductBalances(BaseModel):
)
@property
def available_balance_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.available_balance).to_usd_str()
@computed_field(
@@ -477,8 +470,6 @@ class ProductBalances(BaseModel):
)
@property
def recoup(self) -> USDCent:
- from generalresearch.currency import USDCent
-
if self.balance >= 0:
return USDCent(0)
@@ -517,7 +508,7 @@ class ProductBalances(BaseModel):
if isinstance(input_data, pd.Series):
return ProductBalances.model_validate(input_data.to_dict())
- elif isinstance(input_data, pd.DataFrame):
+ else:
assert isinstance(input_data.index, pd.DatetimeIndex), "Invalid input data"
# The pop merge is grouped by 1min intervals. Therefore, if we take
@@ -530,9 +521,6 @@ class ProductBalances(BaseModel):
pb.last_event = pq_last_event_close.to_pydatetime()
return pb
- else:
- raise NotImplementedError("Can't handle this input")
-
def __str__(self) -> str:
return (
f"Product: {self.product_id or '—'}\n"
@@ -558,7 +546,7 @@ class BusinessBalances(BaseModel):
they all explicitly are set
"""
- if any([pb.product_id is None for pb in v]):
+ if any(pb.product_id is None for pb in v):
raise ValueError("'product_id' must be set for BusinessBalance children.")
return v
@@ -581,8 +569,6 @@ class BusinessBalances(BaseModel):
)
@property
def payout_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.payout).to_usd_str()
@computed_field(
@@ -684,8 +670,6 @@ class BusinessBalances(BaseModel):
)
@property
def payment_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.payment).to_usd_str()
@computed_field(
@@ -733,8 +717,6 @@ class BusinessBalances(BaseModel):
)
@property
def retainer_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.retainer).to_usd_str()
@computed_field(
@@ -766,8 +748,6 @@ class BusinessBalances(BaseModel):
)
@property
def available_balance_usd_str(self) -> str:
- from generalresearch.currency import USDCent
-
return USDCent(self.available_balance).to_usd_str()
# --- Properties: account related ---
@@ -802,8 +782,6 @@ class BusinessBalances(BaseModel):
"""Returns the sum of this Business' recouped amount from any
children Products.
"""
- from generalresearch.currency import USDCent
-
return USDCent(sum([i.recoup for i in self.product_balances]))
@computed_field(
@@ -835,7 +813,7 @@ class BusinessBalances(BaseModel):
def from_pandas(
input_data: pd.DataFrame,
accounts: list[LedgerAccount],
- thl_pg_config: PostgresConfig,
+ product_manager: ProductManager,
) -> BusinessBalances:
LOG.debug(f"BusinessBalances.from_pandas(input_data={input_data.shape})")
@@ -846,16 +824,14 @@ class BusinessBalances(BaseModel):
AccountType,
Direction,
)
- from generalresearch.models.thl.product import Product
- from generalresearch.managers.thl.product import ProductManager
# Validate the input accounts
assert len(accounts) > 0, "Must provide accounts"
- assert all([a.account_type == AccountType.BP_WALLET for a in accounts])
- assert all([a.normal_balance == Direction.CREDIT for a in accounts])
+ assert all(a.account_type == AccountType.BP_WALLET for a in accounts)
+ assert all(a.normal_balance == Direction.CREDIT for a in accounts)
if not is_debug():
- assert all([a.currency == "USD" for a in accounts])
+ assert all(a.currency == "USD" for a in accounts)
# Validate the input dataframe
assert input_data.index.name == "account_id"
@@ -873,8 +849,7 @@ class BusinessBalances(BaseModel):
# Sort the ProductBalances so that they're always in a consistent
# sorted order.
- pm = ProductManager(pg_config=thl_pg_config)
- products: list[Product] = pm.get_by_uuids(
+ products: list[Product] = product_manager.get_by_uuids(
product_uuids=[pb.product_id for pb in product_balances]
)
sorted_products_uuids = [
diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py
index 0d254e2..44dd408 100644
--- a/generalresearch/models/thl/ipinfo.py
+++ b/generalresearch/models/thl/ipinfo.py
@@ -1,10 +1,11 @@
from __future__ import annotations
import ipaddress
-from datetime import datetime, timezone
-from typing import Any, Literal
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any, Literal, Self
from faker import Faker
+from grip_client.enums import AccessType
from pydantic import (
BaseModel,
ConfigDict,
@@ -13,15 +14,15 @@ from pydantic import (
PrivateAttr,
field_validator,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
IPvAnyAddressStr,
)
-from generalresearch.models.thl.maxmind.definitions import UserType
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ipinfo import IPGeonameManager
fake = Faker()
@@ -95,7 +96,7 @@ class IPGeoname(BaseModel):
is_in_european_union: bool | None = Field(default=None)
updated: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
)
@field_validator(
@@ -119,7 +120,7 @@ class IPGeoname(BaseModel):
@classmethod
def from_mysql(cls, d: dict[str, Any]) -> Self:
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
return cls.model_validate(d)
@@ -138,8 +139,7 @@ class IPInformation(BaseModel):
registered_country_iso: CountryISOLike | None = Field(
default=None,
- description="The ISO code of the country where the IP address is "
- "registered.",
+ description="The ISO code of the country where the IP address is registered.",
examples=[fake.country_code().lower()],
)
is_anonymous: bool | None = Field(
@@ -160,7 +160,7 @@ class IPInformation(BaseModel):
domain: str | None = Field(default=None, max_length=255)
isp: str | None = Field(
default=None,
- description="The Internet Service Provider associated with the " "IP address.",
+ description="The Internet Service Provider associated with the IP address.",
examples=["Comcast"],
)
@@ -174,11 +174,11 @@ class IPInformation(BaseModel):
default=None,
description="A score indicating the likelihood that the IP address is static.",
)
- user_type: UserType | None = Field(
+ user_type: AccessType | None = Field(
default=None,
description="The type of user associated with the IP address "
"(e.g., 'residential', 'business').",
- examples=[UserType.SCHOOL],
+ examples=[AccessType.RESIDENTIAL],
)
postal_code: str | None = Field(
default=None,
@@ -205,7 +205,7 @@ class IPInformation(BaseModel):
)
updated: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
)
_geoname: IPGeoname | None = PrivateAttr(default=None)
@@ -219,8 +219,8 @@ class IPInformation(BaseModel):
@property
def basic(self) -> bool:
- # This could be almost any field, but we're checking here if maxmind
- # insights was run on this record. If not, then most of the optional
+ # This could be almost any field, but we're checking here if GRIP
+ # was run on this record. If not, then most of the optional
# fields will be None
return self.is_anonymous is None
@@ -236,26 +236,25 @@ class IPInformation(BaseModel):
# --- prefetch_* ---
def prefetch_geoname(
self,
- pg_config: PostgresConfig,
+ ip_gm: IPGeonameManager,
) -> None:
if self.geoname_id is None:
raise ValueError("Must provide geoname_id")
- from generalresearch.managers.thl.ipinfo import IPGeonameManager
-
- ip_gm = IPGeonameManager(pg_config=pg_config)
+ # from generalresearch.managers.thl.ipinfo import IPGeonameManager
+ # ip_gm = IPGeonameManager(pg_config=pg_config)
self._geoname = ip_gm.get_by_id(geoname_id=self.geoname_id)
# --- ORM ---
def model_dump_mysql(self):
- d = self.model_dump(mode="json", exclude={"geoname"})
+ d = self.model_dump(mode="json")
d["updated"] = self.updated
return d
@classmethod
- def from_mysql(cls, d: dict) -> Self:
- d["updated"] = d["updated"].replace(tzinfo=timezone.utc)
+ def from_mysql(cls, d: dict[str, Any]) -> Self:
+ d["updated"] = d["updated"].replace(tzinfo=UTC)
return cls.model_validate(d)
diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py
index d73c7cf..df60a7f 100644
--- a/generalresearch/models/thl/leaderboard.py
+++ b/generalresearch/models/thl/leaderboard.py
@@ -2,8 +2,8 @@ from __future__ import annotations
import logging
import math
-from datetime import datetime, timedelta, timezone
-from enum import Enum
+from datetime import UTC, datetime, timedelta
+from enum import StrEnum
from typing import Literal
from uuid import UUID, uuid3
from zoneinfo import ZoneInfo
@@ -27,7 +27,7 @@ from generalresearch.utils.enum import ReprEnumMeta
logger = logging.getLogger()
-class LeaderboardCode(str, Enum, metaclass=ReprEnumMeta):
+class LeaderboardCode(StrEnum, metaclass=ReprEnumMeta):
"""
The type of leaderboard. What the "values" represent.
"""
@@ -40,7 +40,7 @@ class LeaderboardCode(str, Enum, metaclass=ReprEnumMeta):
SUM_PAYOUTS = "sum_user_payout"
-class LeaderboardFrequency(str, Enum, metaclass=ReprEnumMeta):
+class LeaderboardFrequency(StrEnum, metaclass=ReprEnumMeta):
"""
The time period range for the leaderboard.
"""
@@ -115,9 +115,10 @@ class Leaderboard(BaseModel):
examples=[LeaderboardFrequency.DAILY],
)
- timezone_name: str = Field(
+ timezone_name: str | None = Field(
description="The timezone for the requested country",
examples=["America/New_York"],
+ default=None,
)
sort_order: Literal["ascending", "descending"] = Field(default="descending")
@@ -146,7 +147,7 @@ class Leaderboard(BaseModel):
# exclude=True,
)
- period_end_local: AwareDatetime = Field(
+ period_end_local: AwareDatetime | None = Field(
description="The end of the time period covered by this board in local time, tz-aware",
examples=[
datetime(
@@ -160,6 +161,7 @@ class Leaderboard(BaseModel):
tzinfo=ZoneInfo("America/New_York"),
)
],
+ default=None,
# exclude=True,
)
@@ -176,13 +178,13 @@ class Leaderboard(BaseModel):
def period_start_utc(self) -> datetime:
# The start of the time period covered by this board in UTC, tz-aware
# e.g. datetime(2024, 7, 12, 4, 0, 0, 0, tzinfo=timezone.utc)
- return self.period_start_local.astimezone(timezone.utc)
+ return self.period_start_local.astimezone(UTC)
@property
def period_end_utc(self) -> datetime:
# The end of the time period covered by this board in UTC, tz-aware
# e.g. datetime(2024, 7, 13, 3, 59, 59, 999999, tzinfo=timezone.utc)
- return self.period_end_local.astimezone(timezone.utc)
+ return self.period_end_local.astimezone(UTC)
@computed_field(
description="(unix timestamp) The start time of the time range this leaderboard covers.",
diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py
index d30f78a..613d16e 100644
--- a/generalresearch/models/thl/ledger.py
+++ b/generalresearch/models/thl/ledger.py
@@ -1,8 +1,8 @@
from __future__ import annotations
-from datetime import datetime, timezone
-from enum import Enum
-from typing import Annotated, Any, Literal, Union
+from datetime import UTC, datetime
+from enum import IntEnum, StrEnum
+from typing import Annotated, Any, Literal, Self
from uuid import uuid4
from pydantic import (
@@ -15,7 +15,6 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import (
AwareDatetimeISO,
@@ -23,12 +22,6 @@ from generalresearch.models.custom_types import (
UUIDStr,
check_valid_uuid,
)
-from generalresearch.models.thl.ledger_example import (
- _example_user_tx_adjustment,
- _example_user_tx_bonus,
- _example_user_tx_complete,
- _example_user_tx_payout,
-)
from generalresearch.models.thl.pagination import Page
from generalresearch.models.thl.payout_format import (
PayoutFormatType,
@@ -37,7 +30,54 @@ from generalresearch.models.thl.payout_format import (
from generalresearch.utils.enum import ReprEnumMeta
-class Direction(int, Enum, metaclass=ReprEnumMeta):
+def _example_user_tx_payout(schema: dict[str, Any]) -> None:
+
+ schema["example"] = UserLedgerTransactionUserPayout(
+ product_id=uuid4().hex,
+ payout_id=uuid4().hex,
+ amount=-5,
+ description="HIT Reward",
+ payout_format="${payout/100:.2f}",
+ created=datetime.now(tz=UTC),
+ ).model_dump(mode="json")
+
+
+def _example_user_tx_bonus(schema: dict[str, Any]) -> None:
+
+ schema["example"] = UserLedgerTransactionUserBonus(
+ product_id=uuid4().hex,
+ amount=100,
+ description="Compensation Bonus",
+ payout_format="${payout/100:.2f}",
+ created=datetime.now(tz=UTC),
+ ).model_dump(mode="json")
+
+
+def _example_user_tx_complete(schema: dict[str, Any]) -> None:
+
+ schema["example"] = UserLedgerTransactionTaskComplete(
+ product_id=uuid4().hex,
+ amount=38,
+ description="Task Complete",
+ payout_format="${payout/100:.2f}",
+ created=datetime.now(tz=UTC),
+ tsid=uuid4().hex,
+ ).model_dump(mode="json")
+
+
+def _example_user_tx_adjustment(schema: dict[str, Any]) -> None:
+
+ schema["example"] = UserLedgerTransactionTaskAdjustment(
+ product_id=uuid4().hex,
+ amount=-38,
+ description="Task Adjustment",
+ payout_format="${payout/100:.2f}",
+ created=datetime.now(tz=UTC),
+ tsid=uuid4().hex,
+ ).model_dump(mode="json")
+
+
+class Direction(IntEnum, metaclass=ReprEnumMeta):
"""Entries on the debit side will increase debit normal accounts, while
entries on the credit side will decrease them. Conversely, entries on
the credit side will increase credit normal accounts, while entries on
@@ -51,13 +91,13 @@ class Direction(int, Enum, metaclass=ReprEnumMeta):
DEBIT = 1
-class OrderBy(str, Enum, metaclass=ReprEnumMeta):
+class OrderBy(StrEnum, metaclass=ReprEnumMeta):
ASC = "ASC"
DESC = "DESC"
-class AccountType(str, Enum, metaclass=ReprEnumMeta):
+class AccountType(StrEnum, metaclass=ReprEnumMeta):
# Revenue from BP payment commission
BP_COMMISSION = "bp_commission"
# BP wallets (owed balance)
@@ -86,7 +126,7 @@ class AccountType(str, Enum, metaclass=ReprEnumMeta):
WA_CREDIT_LINE = "wa_credit_line"
-class TransactionMetadataColumns(str, Enum):
+class TransactionMetadataColumns(StrEnum):
BONUS = "bonus_id"
# Note: EVENT & EVENT2 represent the same concept. I accidentally made
# this inconsistent.
@@ -105,7 +145,7 @@ class TransactionMetadataColumns(str, Enum):
CONTEST = "contest"
-class TransactionType(str, Enum):
+class TransactionType(StrEnum):
"""These are used in the Ledger to annotate the type of transaction (in
metadata: tx_type)
"""
@@ -258,7 +298,7 @@ class LedgerTransaction(BaseModel):
id: int | None = Field(default=None)
created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When the Transaction (TX) was created into the database."
"This does not represent the exact time for any action"
"which may be responsible for this Transaction (TX), and "
@@ -288,9 +328,7 @@ class LedgerTransaction(BaseModel):
"""Created should not be in the future. This will mess up
LedgerAccountStatement / groupby rollups.
"""
- assert (
- datetime.now(tz=timezone.utc) > created
- ), "created cannot be in the future"
+ assert datetime.now(tz=UTC) > created, "created cannot be in the future"
return created
@field_validator("entries", mode="after")
@@ -302,13 +340,13 @@ class LedgerTransaction(BaseModel):
"""
if entries:
assert len(entries) >= 2, "ledger transaction must have 2 or more entries"
- assert (
- sum(x.amount * x.direction for x in entries) == 0
- ), "ledger entries must balance"
+ assert sum(x.amount * x.direction for x in entries) == 0, (
+ "ledger entries must balance"
+ )
return entries
- def model_dump_mysql(self, *args, **kwargs) -> dict[str, Any]:
- d = self.model_dump(mode="json", *args, **kwargs)
+ def model_dump_mysql(self, **kwargs) -> dict[str, Any]:
+ d = self.model_dump(mode="json", **kwargs)
if "created" in d:
d["created"] = self.created.replace(tzinfo=None)
return d
@@ -316,7 +354,7 @@ class LedgerTransaction(BaseModel):
def to_user_tx(
self, user_account: LedgerAccount, product_id: str, payout_format: str
):
- from generalresearch.models.thl.wallet import PayoutType
+ from generalresearch.models.thl.wallet.definitions import PayoutType
d = self.model_dump(include={"created"})
d["tx_type"] = self.metadata.get("tx_type")
@@ -401,7 +439,7 @@ class UserLedgerTransaction(BaseModel):
# It is optional b/c we'll calculate this from the query
balance_after: int | None = Field(default=None)
- def create_url(self, product_id: str):
+ def create_url(self, product_id: str) -> str | None:
raise NotImplementedError()
@computed_field(
@@ -439,7 +477,7 @@ class UserLedgerTransactionUserPayout(UserLedgerTransaction):
examples=["a3848e0a53d64f68a74ced5f61b6eb68"],
)
- def create_url(self, product_id: str):
+ def create_url(self, product_id: str) -> str | None:
return f"https://fsb.generalresearch.com/{product_id}/cashout/{self.payout_id}/"
@model_validator(mode="after")
@@ -467,7 +505,7 @@ class UserLedgerTransactionUserBonus(UserLedgerTransaction):
default="Compensation Bonus",
)
- def create_url(self, product_id: str):
+ def create_url(self, product_id: str) -> str | None:
return None
@model_validator(mode="after")
@@ -505,7 +543,7 @@ class UserLedgerTransactionTaskComplete(UserLedgerTransaction):
examples=["a3848e0a53d64f68a74ced5f61b6eb68"],
)
- def create_url(self, product_id: str):
+ def create_url(self, product_id: str) -> str | None:
return f"https://fsb.generalresearch.com/{product_id}/status/{self.tsid}/"
@model_validator(mode="after")
@@ -536,17 +574,15 @@ class UserLedgerTransactionTaskAdjustment(UserLedgerTransaction):
examples=["a3848e0a53d64f68a74ced5f61b6eb68"],
)
- def create_url(self, product_id: str):
+ def create_url(self, product_id: str) -> str | None:
return f"https://fsb.generalresearch.com/{product_id}/status/{self.tsid}/"
UserLedgerTransactionType = Annotated[
- Union[
- UserLedgerTransactionUserPayout,
- UserLedgerTransactionUserBonus,
- UserLedgerTransactionTaskAdjustment,
- UserLedgerTransactionTaskComplete,
- ],
+ UserLedgerTransactionUserPayout
+ | UserLedgerTransactionUserBonus
+ | UserLedgerTransactionTaskAdjustment
+ | UserLedgerTransactionTaskComplete,
Field(discriminator="tx_type"),
]
diff --git a/generalresearch/models/thl/ledger_example.py b/generalresearch/models/thl/ledger_example.py
deleted file mode 100644
index 767be85..0000000
--- a/generalresearch/models/thl/ledger_example.py
+++ /dev/null
@@ -1,64 +0,0 @@
-from __future__ import annotations
-
-from datetime import datetime, timezone
-from typing import Any
-from uuid import uuid4
-
-
-def _example_user_tx_payout(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.ledger import (
- UserLedgerTransactionUserPayout,
- )
-
- schema["example"] = UserLedgerTransactionUserPayout(
- product_id=uuid4().hex,
- payout_id=uuid4().hex,
- amount=-5,
- description="HIT Reward",
- payout_format="${payout/100:.2f}",
- created=datetime.now(tz=timezone.utc),
- ).model_dump(mode="json")
-
-
-def _example_user_tx_bonus(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.ledger import (
- UserLedgerTransactionUserBonus,
- )
-
- schema["example"] = UserLedgerTransactionUserBonus(
- product_id=uuid4().hex,
- amount=100,
- description="Compensation Bonus",
- payout_format="${payout/100:.2f}",
- created=datetime.now(tz=timezone.utc),
- ).model_dump(mode="json")
-
-
-def _example_user_tx_complete(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.ledger import (
- UserLedgerTransactionTaskComplete,
- )
-
- schema["example"] = UserLedgerTransactionTaskComplete(
- product_id=uuid4().hex,
- amount=38,
- description="Task Complete",
- payout_format="${payout/100:.2f}",
- created=datetime.now(tz=timezone.utc),
- tsid=uuid4().hex,
- ).model_dump(mode="json")
-
-
-def _example_user_tx_adjustment(schema: dict[str, Any]) -> None:
- from generalresearch.models.thl.ledger import (
- UserLedgerTransactionTaskAdjustment,
- )
-
- schema["example"] = UserLedgerTransactionTaskAdjustment(
- product_id=uuid4().hex,
- amount=-38,
- description="Task Adjustment",
- payout_format="${payout/100:.2f}",
- created=datetime.now(tz=timezone.utc),
- tsid=uuid4().hex,
- ).model_dump(mode="json")
diff --git a/generalresearch/models/thl/maxmind/__init__.py b/generalresearch/models/thl/maxmind/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/generalresearch/models/thl/maxmind/__init__.py
+++ /dev/null
diff --git a/generalresearch/models/thl/maxmind/definitions.py b/generalresearch/models/thl/maxmind/definitions.py
deleted file mode 100644
index 01431c7..0000000
--- a/generalresearch/models/thl/maxmind/definitions.py
+++ /dev/null
@@ -1,22 +0,0 @@
-from enum import Enum
-
-from generalresearch.utils.enum import ReprEnumMeta
-
-
-class UserType(Enum, metaclass=ReprEnumMeta):
- # https://support.maxmind.com/hc/en-us/articles/4408430082971-IP-Trait-Risk-Data#h_01FN6V8JMQMWZGWNPPAW77ZPY4
- BUSINESS = "business"
- CAFE = "cafe"
- CELLULAR = "cellular"
- COLLEGE = "college"
- CDN = "content_delivery_network"
- CPN = "consumer_privacy_network"
- GOVERNMENT = "government"
- HOSTING = "hosting"
- LIBRARY = "library"
- MILITARY = "military"
- RESIDENTIAL = "residential"
- ROUTER = "router"
- SCHOOL = "school"
- SEARCH_ENGINE = "search_engine_spider"
- TRAVELER = "traveler"
diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py
index d2d7d36..599cc1d 100644
--- a/generalresearch/models/thl/offerwall/__init__.py
+++ b/generalresearch/models/thl/offerwall/__init__.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import hashlib
import json
from decimal import Decimal
-from enum import Enum
-from typing import Any, Literal
+from enum import StrEnum
+from typing import Any, Literal, Self
from pydantic import (
BaseModel,
@@ -13,10 +13,9 @@ from pydantic import (
computed_field,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.models import Source
from generalresearch.models.custom_types import IPvAnyAddressStr
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.locales import (
CountryISO,
LanguageISO,
@@ -32,7 +31,7 @@ from generalresearch.models.thl.product import (
from generalresearch.models.thl.user import User
-class OfferWallType(str, Enum):
+class OfferWallType(StrEnum):
"""
The specific offerwall type
"""
@@ -57,7 +56,7 @@ class OfferWallType(str, Enum):
STARWALL = "b59a2d2b"
-class OfferWallTypeClass(str, Enum):
+class OfferWallTypeClass(StrEnum):
"""
A higher level "class" to organize similar offerwall types.
For e.g. STARWALL_PLUS_BLOCK, STARWALL_PLUS, STARWALL all use the same
@@ -268,7 +267,7 @@ class OfferWallRequest(BaseModel):
# We need this so thl-core can refresh an offerwall in order to continue
# a session
d = self.model_dump(mode="json")
- kwargs = dict()
+ kwargs = {}
keys = [
"n_bins",
"min_bin_size",
diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py
index 3a867b6..ead48c3 100644
--- a/generalresearch/models/thl/offerwall/base.py
+++ b/generalresearch/models/thl/offerwall/base.py
@@ -4,7 +4,7 @@ import statistics
from datetime import timedelta
from decimal import Decimal
from string import Formatter
-from typing import Any
+from typing import TYPE_CHECKING, Annotated, Any, Self
from uuid import uuid4
import numpy as np
@@ -18,10 +18,9 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Annotated, Self
-from generalresearch.models import Source
from generalresearch.models.custom_types import HttpsUrl, UUIDStr
+from generalresearch.models.definitions import Source
from generalresearch.models.legacy.bucket import (
Bucket as LegacyBucket,
)
@@ -45,7 +44,9 @@ from generalresearch.models.thl.offerwall.bucket import (
)
from generalresearch.models.thl.profiling.upk_question import UpkQuestion
from generalresearch.models.thl.soft_pair import SoftPairResultType
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.user import User
class MergeTableFeatures(BaseModel):
@@ -272,9 +273,9 @@ class TaskResult(BaseModel):
if fname
]
)
- assert all(
- x in {"domain", "mid"} for x in fmt_str
- ), "unrecognized format variable"
+ assert all(x in {"domain", "mid"} for x in fmt_str), (
+ "unrecognized format variable"
+ )
else:
assert self.entry_link is None, f"entry link not allowed for {self.source}"
return self
@@ -399,8 +400,10 @@ class OfferwallBucket(BaseModel):
)
uri: HttpsUrl | None = Field(
examples=[
- "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
- "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ (
+ "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i"
+ "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e"
+ )
],
description="The URL to send a respondent into. Must not edit this URL in any way",
default=None,
diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py
index c36568e..abb213e 100644
--- a/generalresearch/models/thl/offerwall/cache.py
+++ b/generalresearch/models/thl/offerwall/cache.py
@@ -1,12 +1,12 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import Any
from pydantic import BaseModel, Field
-from generalresearch.models import Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.offerwall import OfferWallRequest
from generalresearch.models.thl.offerwall.base import (
OfferwallBase,
@@ -26,9 +26,7 @@ class GetOfferWallCache(BaseModel):
request_id: str = Field()
offerwall: OfferwallBase = Field()
all_sids: list[str] = Field()
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(UTC))
latest_ip_info: dict[str, Any] = Field(
description="So we can easily check if user's IP info has changed"
)
@@ -51,9 +49,7 @@ class SessionInfoCache(BaseModel):
# will get pruned as tasks are attempted
tasks: list[ScoredTaskResult] = Field()
- started: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# The count of attempts per marketplace
mp_retry_count: dict[Source, int] = Field(default_factory=dict)
diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py
index 1a9d534..8759cd4 100644
--- a/generalresearch/models/thl/payout.py
+++ b/generalresearch/models/thl/payout.py
@@ -1,28 +1,32 @@
from __future__ import annotations
import json
-from datetime import datetime, timezone
-from typing import Collection
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Self
from uuid import uuid4
from pydantic import (
BaseModel,
+ ConfigDict,
Field,
PositiveInt,
computed_field,
field_validator,
+ model_validator,
)
-from typing_extensions import Self
+from pydantic.json_schema import SkipJsonSchema
from generalresearch.currency import USDCent
-from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.custom_types import (
+ AwareDatetimeISO,
+ UUIDStr,
+ UUIDStrCoerce,
+)
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.ledger import OrderBy
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailOrderData,
)
-from generalresearch.redis_helper import RedisConfig
+from generalresearch.models.thl.wallet.definitions import PayoutType
class PayoutEvent(BaseModel):
@@ -39,29 +43,27 @@ class PayoutEvent(BaseModel):
multiple BrokerageProductPayoutEvents.
"""
- uuid: UUIDStr = Field(
+ uuid: UUIDStrCoerce = Field(
title="Payout Event Unique Identifier",
default_factory=lambda: uuid4().hex,
examples=["9453cd076713426cb68d05591c7145aa"],
)
- debit_account_uuid: UUIDStr = Field(
+ debit_account_uuid: UUIDStrCoerce | None = Field(
description="The LedgerAccount.uuid that money is being requested from. "
"Thie User or Brokerage Product is retrievable through the "
"LedgerAccount.reference_uuid",
examples=["18298cb1583846fbb06e4747b5310693"],
)
- cashout_method_uuid: UUIDStr = Field(
+ cashout_method_uuid: UUIDStrCoerce | None = Field(
description="References a row in the account_cashoutmethod table. This "
"is the cashout method that was used to request this "
"payout. (A cashout is the same thing as a payout)",
examples=["a6dc1fc1bf934557b952f253dee12813"],
)
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# In the smallest unit of the currency being transacted. For USD, this
# is cents.
@@ -84,8 +86,8 @@ class PayoutEvent(BaseModel):
description=PayoutType.as_openapi(), examples=[PayoutType.ACH]
)
- request_data: dict = Field(
- default_factory=dict,
+ request_data: dict | None = Field(
+ default=None,
description="Stores payout-type-specific information that is used to "
"request this payout from the external provider.",
)
@@ -144,16 +146,15 @@ class PayoutEvent(BaseModel):
# --- ORM ---
- def model_dump_mysql(self, *args, **kwargs) -> dict:
- d = self.model_dump(mode="json", *args, **kwargs)
-
- if "created" in d:
- d["created"] = self.created.replace(tzinfo=None)
-
- if d.get("request_data") is not None:
- d["request_data"] = json.dumps(self.request_data)
+ def model_dump_postgres(self) -> dict:
+ d = self.model_dump(mode="json", exclude={"request_data", "order_data"})
- if d.get("order_data") is not None:
+ d["request_data"] = (
+ json.dumps(self.request_data) if self.request_data is not None else None
+ )
+ if self.order_data is None:
+ d["order_data"] = None
+ else:
if isinstance(self.order_data, dict):
d["order_data"] = json.dumps(self.order_data)
else:
@@ -196,7 +197,7 @@ class BrokerageProductPayoutEvent(PayoutEvent):
- created: When the Brokerage Product was paid out
"""
- product_id: UUIDStr = Field(
+ product_id: UUIDStrCoerce = Field(
description="The Brokerage Product that was paid out",
examples=["1108d053e4fa47c5b0dbdcd03a7981e7"],
)
@@ -220,136 +221,136 @@ class BrokerageProductPayoutEvent(PayoutEvent):
def amount_usd_str(self) -> str:
return self.amount_usd.to_usd_str()
- # --- ORM ---
- @classmethod
- def from_payout_event(
- cls,
- pe: PayoutEvent,
- account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
- redis_config: RedisConfig | None = None,
- ) -> Self:
- # TODO!: prevent re-assignment, rework this...
-
- if account_product_mapping is None:
- rc = redis_config.create_redis_client()
- account_product_mapping: dict = rc.hgetall(name="pem:account_to_product")
- assert isinstance(account_product_mapping, dict)
- assert pe.uuid in account_product_mapping.keys()
-
- d = pe.model_dump()
- d["product_id"] = account_product_mapping[pe.debit_account_uuid]
- return cls.model_validate(d)
+class BusinessPayoutEventCreate(BaseModel):
+ """A single payout event to a supplier Business."""
- @classmethod
- def from_payout_events(
- cls,
- payout_events: Collection[PayoutEvent],
- order_by=OrderBy,
- account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
- redis_config: RedisConfig | None = None,
- ) -> list[Self]:
- # TODO!: prevent re-assignment, rework this...
-
- if account_product_mapping is None:
- rc = redis_config.create_redis_client()
- account_product_mapping: dict = rc.hgetall(name="pem:account_to_product")
- assert isinstance(account_product_mapping, dict)
-
- res = []
- for pe in payout_events:
- res.append(
- cls.from_payout_event(
- pe=pe, account_product_mapping=account_product_mapping
- )
- )
+ model_config = ConfigDict(validate_assignment=True)
- match order_by:
- case OrderBy.ASC:
- sorted_list = sorted(res, key=lambda x: x.created, reverse=False)
- case OrderBy.DESC:
- sorted_list = sorted(res, key=lambda x: x.created, reverse=True)
- case _:
- raise ValueError("Invalid order provided..")
+ # Used for holding a *unique*, external, payout-type-specific identifier.
+ ext_ref_id: str = Field(title="Unique external reference ID")
- return sorted_list
+ business_id: UUIDStr = Field(
+ description="The Business receiving this supplier payout.",
+ examples=[uuid4().hex],
+ )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
-class BusinessPayoutEvent(BaseModel):
- """A single ACH or Wire event to a Business Bank Account"""
+ # In the smallest unit of the currency being transacted. For USD, this
+ # is cents.
+ amount: PositiveInt = Field(
+ lt=2**63 - 1,
+ strict=True,
+ title="Amount",
+ description="The amount issued to the supplier.",
+ examples=[1_982_343],
+ )
- bp_payouts: list[BrokerageProductPayoutEvent] = Field(
- description="Here is the list of Brokerage Product Payouts that"
- "this Business Payout includes.",
- min_length=1,
+ status: PayoutStatus = Field(
+ default=PayoutStatus.PENDING,
+ description=PayoutStatus.as_openapi(),
+ examples=[PayoutStatus.COMPLETE],
)
- @computed_field(
- title="Amount",
- description="The amount issued to the Bank Account",
- examples=[19_823_43],
- return_type=USDCent,
+ payout_type: PayoutType = Field(
+ description=PayoutType.as_openapi(), examples=[PayoutType.ACH]
)
- @property
- def amount(self) -> USDCent:
- return USDCent(sum([p.amount for p in self.bp_payouts]))
- @computed_field(
- title="Amount USD Str",
- description="The amount issued to the Bank Account as a USD string",
- examples=["$19,823.43"],
- return_type=str,
+ request_data: dict | None = Field(
+ default=None,
+ description="Stores payout-type-specific information that is used to "
+ "request this payout from the external provider.",
)
- @property
- def amount_usd_str(self) -> str:
- return self.amount.to_usd_str()
- @computed_field(
- title="Created",
- description="This is equal to the created time of the first"
- "Brokerage Product Payout Event.",
- return_type=AwareDatetimeISO,
+ order_data: dict | None = Field(
+ default=None,
+ description="Stores payout-type-specific order information that is "
+ "returned from the external payout provider.",
)
- @property
- def created(self) -> AwareDatetimeISO:
- return self.bp_payouts[0].created
- @computed_field(
- title="Line Items",
- description="The number of sub-payments",
- return_type=PositiveInt,
+ bp_payouts: list[BrokerageProductPayoutEvent] = Field(
+ description="The list of Brokerage Product Payouts that this Business Payout includes",
+ min_length=1,
)
- @property
- def line_items(self):
- return len(self.bp_payouts)
@computed_field(
- title="External Reference ID",
- description="ACH Transaction ID",
- return_type=str | None,
+ title="Amount USD Str",
+ description="The amount issued to the supplier as a USD string",
+ examples=["$19,823.43"],
+ return_type=str,
)
@property
- def ext_ref_id(self):
- return self.bp_payouts[0].ext_ref_id
+ def amount_usd_str(self) -> str:
+ return USDCent(self.amount).to_usd_str()
# --- Validators ---
+ @field_validator("payout_type", mode="before")
+ @classmethod
+ def normalize_payout_type(cls, v):
+ if isinstance(v, str):
+ try:
+ return PayoutType[v.upper()]
+ except KeyError:
+ raise ValueError(f"Invalid payout_type: {v}")
+ return v
+
@field_validator("bp_payouts", mode="before")
@classmethod
- def normalize_enum(cls, v):
+ def validate_bp_payouts_type(cls, v):
"""This can be a list of Instances or Python Dictionaries depending
on how it's initialized.
"""
+ if v is None:
+ return v
+
assert isinstance(v, list)
+ return v
- def get_field(obj, field):
- if isinstance(obj, dict):
- return obj.get(field)
- return getattr(obj, field, None)
+ @model_validator(mode="after")
+ def validate_bp_payouts(self) -> Self:
+ bp_payout_amount = sum([p.amount for p in self.bp_payouts])
+ if bp_payout_amount != self.amount:
+ raise ValueError(
+ "BusinessPayoutEvent.amount must equal the sum of "
+ f"bp_payouts amounts ({self.amount=} {bp_payout_amount=})"
+ )
- assert all(
- get_field(i, "ext_ref_id") == get_field(v[0], "ext_ref_id") for i in v
- ), "Not all group values are the same"
+ invalid_payout_types = [
+ p.payout_type for p in self.bp_payouts if p.payout_type != self.payout_type
+ ]
+ if invalid_payout_types:
+ raise ValueError(
+ "All BrokerageProductPayoutEvent.payout_type values must equal "
+ f"BusinessPayoutEvent.payout_type ({self.payout_type})"
+ )
- return v
+ invalid_ext_ids = [
+ p.ext_ref_id for p in self.bp_payouts if p.ext_ref_id != self.ext_ref_id
+ ]
+ if invalid_ext_ids:
+ raise ValueError(
+ "All BrokerageProductPayoutEvent.ext_ref_id values must equal "
+ f"BusinessPayoutEvent.ext_ref_id ({self.ext_ref_id})"
+ )
+
+ return self
+
+ def model_dump_postgres(self):
+ d = self.model_dump(
+ mode="json",
+ exclude={"bp_payouts"},
+ )
+ d["request_data"] = (
+ json.dumps(self.request_data) if self.request_data is not None else None
+ )
+ d["order_data"] = (
+ json.dumps(self.order_data) if self.order_data is not None else None
+ )
+ return d
+
+
+class BusinessPayoutEvent(BusinessPayoutEventCreate):
+ id: SkipJsonSchema[PositiveInt] = Field(exclude=True)
diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py
index 9f22ace..d29c9de 100644
--- a/generalresearch/models/thl/payout_format.py
+++ b/generalresearch/models/thl/payout_format.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import decimal
import re
+from typing import Annotated
from pydantic import AfterValidator, Field
-from typing_extensions import Annotated
# Matches only digits, parenthesis, + , -, *, / and the string payout.
xform_format_re = re.compile(pattern=r"^[\d()+\-*/.]*payout[\d()+\-*/.]*$")
@@ -45,7 +45,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str:
try:
xform, formatstr = inside.split(":")
- except ValueError as e:
+ except ValueError:
raise ValueError(
"Payout format string must contain ':' to distinguish between transformations and formatting."
)
@@ -61,17 +61,17 @@ def format_payout_format(payout_format: str, payout_int: int) -> str:
payout = decimal.Decimal(eval(xform, {"payout": payout_int}))
- except NameError as e:
+ except NameError:
raise ValueError("Payout format string must contain 'payout' variable.")
- except ZeroDivisionError as e:
+ except ZeroDivisionError:
raise ValueError("Cannot divide by zero.")
- except TypeError as e:
+ except TypeError:
# "{payout()*1:}" - TypeError: 'int' object is not callable
raise ValueError("Invalid type reference.")
- except Exception as e:
- raise ValueError(f"Invalid payout transformation")
+ except Exception: # noqa
+ raise ValueError("Invalid payout transformation")
formatstr = f"{{:{formatstr}}}"
diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py
index a179142..7441318 100644
--- a/generalresearch/models/thl/product.py
+++ b/generalresearch/models/thl/product.py
@@ -6,14 +6,15 @@ import json
import math
import warnings
from collections import defaultdict
+from collections.abc import Callable
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from functools import cached_property, partial
from typing import (
TYPE_CHECKING,
Any,
- Callable,
Literal,
+ Self,
)
from urllib.parse import parse_qs, urlencode, urlsplit, urlunsplit
from uuid import uuid4
@@ -34,18 +35,23 @@ from pydantic import (
model_validator,
)
from pydantic.json_schema import SkipJsonSchema
-from typing_extensions import Self
from generalresearch.currency import USDCent
from generalresearch.decorators import LOG
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
HttpsUrlStr,
UUIDStr,
)
-from generalresearch.models.thl.ledger import LedgerAccount
+from generalresearch.models.definitions import Source
+from generalresearch.models.thl.finance import (
+ POPFinancial,
+ ProductBalances,
+)
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+)
from generalresearch.models.thl.payout_format import (
PayoutFormatType,
format_payout_format,
@@ -57,7 +63,7 @@ from generalresearch.models.thl.payout_format import (
examples as payout_format_examples,
)
from generalresearch.models.thl.supplier_tag import SupplierTag
-from generalresearch.models.thl.wallet import PayoutType
+from generalresearch.models.thl.wallet.definitions import PayoutType
from generalresearch.models.utils import decimal_to_usd_cents
from generalresearch.redis_helper import RedisConfig
@@ -70,14 +76,7 @@ if TYPE_CHECKING:
from generalresearch.managers.thl.payout import (
BrokerageProductPayoutEventManager,
)
- from generalresearch.models.thl.finance import (
- POPFinancial,
- ProductBalances,
- )
- from generalresearch.models.thl.payout import (
- BrokerageProductPayoutEvent,
- )
- from generalresearch.models.thl.user import User
+ from generalresearch.models.thl.ledger import LedgerAccount
# fmt: off
@@ -446,16 +445,16 @@ class UserWalletConfig(BaseModel):
@field_serializer("supported_payout_types", when_used="json")
def serialize_supported_payout_types_in_order(
self, supported_payout_types: set[PayoutType]
- ) -> set[PayoutType]:
- return set(sorted(supported_payout_types))
+ ) -> list[PayoutType]:
+ return sorted(supported_payout_types)
@field_validator("min_cashout", "failed_attempt_credit", mode="after")
@classmethod
def check_payout_decimal_places(cls, v: Decimal) -> Decimal:
if v is not None:
- assert (
- v.as_tuple().exponent >= -2
- ), "Must have 2 or fewer decimal places ('XXX.YY')"
+ assert v.as_tuple().exponent >= -2, (
+ "Must have 2 or fewer decimal places ('XXX.YY')"
+ )
# explicitly make sure it is 2 decimal places, after checking that it is
# already 2 or less.
v = v.quantize(Decimal("0.00"))
@@ -465,9 +464,9 @@ class UserWalletConfig(BaseModel):
def check_enabled(self):
if self.enabled is False:
assert self.amt is False, "amt can't be set if enabled is False"
- assert (
- self.min_cashout is None
- ), "min_cashout can't be set if enabled is False"
+ assert self.min_cashout is None, (
+ "min_cashout can't be set if enabled is False"
+ )
else:
if self.min_cashout is None:
self.min_cashout = Decimal("0.01")
@@ -501,9 +500,9 @@ class PayoutTransformationPercentArgs(BaseModel):
@classmethod
def check_payout_decimal_places(cls, v: Decimal) -> Decimal:
if v is not None:
- assert (
- v.as_tuple().exponent >= -2
- ), "Must have 2 or fewer decimal places ('XXX.YY')"
+ assert v.as_tuple().exponent >= -2, (
+ "Must have 2 or fewer decimal places ('XXX.YY')"
+ )
# explicitly make sure it is 2 decimal places, after checking that it is
# already 2 or less.
v = v.quantize(Decimal("0.00"))
@@ -562,8 +561,8 @@ class PayoutTransformation(BaseModel):
def payout_transformation_percent(
self,
payout: Decimal,
- pct: Decimal = 1,
- min_payout: Decimal | None = 0,
+ pct: Decimal = Decimal(1),
+ min_payout: Decimal | None = None,
max_payout: Decimal | None = None,
) -> Decimal:
"""Payout transformation for user displayed values"""
@@ -571,14 +570,14 @@ class PayoutTransformation(BaseModel):
min_payout = Decimal(0)
pct = Decimal(pct)
- payout = Decimal(payout)
+ _payout = Decimal(payout)
min_payout = Decimal(min_payout)
max_payout = Decimal(max_payout) if max_payout else None
- payout: Decimal = payout * pct
- payout: Decimal = max([payout, min_payout])
- payout: Decimal = min([payout, max_payout]) if max_payout else payout
- return payout
+ _payout = _payout * pct
+ _payout = max(_payout, min_payout)
+ _payout = min(_payout, max_payout) if max_payout is not None else _payout
+ return _payout
def payout_transformation_amt(
self, payout: Decimal, user_wallet_balance: Decimal | None = None
@@ -588,22 +587,22 @@ class PayoutTransformation(BaseModel):
# (display, adjustment) so ignore the 7-cent rounding.
if user_wallet_balance is None:
return self.payout_transformation_percent(payout=payout, pct=Decimal(".95"))
- payout = Decimal(payout)
+ _payout = Decimal(payout)
- payout: Decimal = payout * Decimal("0.95")
- new_balance = payout + user_wallet_balance
+ _payout: Decimal = _payout * Decimal("0.95")
+ new_balance = _payout + user_wallet_balance
# If the new_balance is <0, we aren't paying anything, so use the
# full amount
if new_balance < 0:
- return payout
+ return _payout
amt = (5 * math.floor((int(new_balance * 100) - 2) / 5)) + 2
rounded_new_balance = Decimal(amt / 100).quantize(Decimal("0.00"))
- payout = rounded_new_balance - user_wallet_balance
- if payout < Decimal(0):
+ _payout = rounded_new_balance - user_wallet_balance
+ if _payout < Decimal(0):
return Decimal(0)
- return payout
+ return _payout
class SourceConfig(BaseModel):
@@ -646,13 +645,13 @@ class SourceConfig(BaseModel):
)
-class Scope(str, Enum):
+class Scope(StrEnum):
GLOBAL = "global"
TEAM = "team"
PRODUCT = "product"
-class IntegrationMode(str, Enum):
+class IntegrationMode(StrEnum):
# We handle integration, get paid
PLATFORM = "platform"
# "external" credentials, we do not get paid for this activity
@@ -685,9 +684,9 @@ class SupplyConfig(BaseModel):
if c.scope == Scope.TEAM
for team_id in c.team_ids
]
- assert len(team_names) == len(
- set(team_names)
- ), "Can only have one TEAM policy per Source per Team"
+ assert len(team_names) == len(set(team_names)), (
+ "Can only have one TEAM policy per Source per Team"
+ )
return self
@model_validator(mode="after")
@@ -698,9 +697,9 @@ class SupplyConfig(BaseModel):
if c.scope == Scope.PRODUCT
for product_id in c.product_ids
]
- assert len(bp_names) == len(
- set(bp_names)
- ), "Can only have one PRODUCT policy per Source per BP"
+ assert len(bp_names) == len(set(bp_names)), (
+ "Can only have one PRODUCT policy per Source per BP"
+ )
return self
@property
@@ -750,8 +749,8 @@ class SupplyConfig(BaseModel):
Use global config.
"""
d = self.global_scoped_policies_dict.copy()
- d.update(self.team_scoped_policies_dict.get(team_id, dict()))
- d.update(self.product_scoped_policies_dict.get(product_id, dict()))
+ d.update(self.team_scoped_policies_dict.get(team_id, {}))
+ d.update(self.product_scoped_policies_dict.get(product_id, {}))
return d
def get_config_for_product(self, product: Product) -> MergedSupplyConfig:
@@ -770,7 +769,7 @@ class SupplyConfig(BaseModel):
supply_policy=policy_dict[source],
source_config=sources_dict[source],
)
- for source in policy_dict.keys()
+ for source in policy_dict
]
)
@@ -959,9 +958,7 @@ class Product(BaseModel, validate_assignment=True):
# Initialization is deferred until unless it's called
# (see .prebuild_***())
- balance: ProductBalances | None = Field(
- default=None, description="Product Balance"
- )
+ balance: ProductBalances | None = Field(default=None, description="Product Balance")
payouts_total_str: str | None = Field(default=None)
payouts_total: USDCent | None = Field(default=None)
@@ -978,7 +975,7 @@ class Product(BaseModel, validate_assignment=True):
@field_validator("harmonizer_domain", mode="before")
def harmonizer_domain_https(cls, s: str | None):
# in the db, this has no scheme. accept both with a default of https://
- if s is not None and not (s.startswith("https://") or s.startswith("http://")):
+ if s is not None and not (s.startswith(("https://", "http://"))):
s = f"https://{s}"
return s
@@ -992,15 +989,15 @@ class Product(BaseModel, validate_assignment=True):
def harmonizer_domain_only(cls, s: str):
# maks sure there is no path
url_split = urlsplit(s)
- assert (
- url_split.path == "/"
- ), f"harmonizer_domain should be a schema+domain only: {url_split.path}"
- assert (
- url_split.query == ""
- ), f"harmonizer_domain should be a schema+domain only: {url_split.query}"
- assert (
- url_split.fragment == ""
- ), f"harmonizer_domain should be a schema+domain only: {url_split.fragment}"
+ assert url_split.path == "/", (
+ f"harmonizer_domain should be a schema+domain only: {url_split.path}"
+ )
+ assert url_split.query == "", (
+ f"harmonizer_domain should be a schema+domain only: {url_split.query}"
+ )
+ assert url_split.fragment == "", (
+ f"harmonizer_domain should be a schema+domain only: {url_split.fragment}"
+ )
return s
@field_validator("redirect_url", mode="after")
@@ -1021,10 +1018,12 @@ class Product(BaseModel, validate_assignment=True):
@property
def business_uuid(self) -> UUIDStr:
+ assert self.business_id
return self.business_id
@property
def team_uuid(self) -> UUIDStr:
+ assert self.team_id
return self.team_id
@property
@@ -1108,7 +1107,6 @@ class Product(BaseModel, validate_assignment=True):
from generalresearch.incite.schemas.mergers.pop_ledger import (
numerical_col_names,
)
- from generalresearch.models.thl.ledger import LedgerAccount
account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=self)
assert self.id == account.reference_uuid
@@ -1225,7 +1223,6 @@ class Product(BaseModel, validate_assignment=True):
from generalresearch.models.thl.ledger import OrderBy
self.payouts = bp_pem.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm,
product_uuids=[self.uuid],
order_by=OrderBy.DESC,
)
@@ -1378,9 +1375,7 @@ class Product(BaseModel, validate_assignment=True):
if self.payout_config.payout_transformation is None:
return lambda x: x
else:
- return (
- self.payout_config.payout_transformation.get_payout_transformation_func()
- )
+ return self.payout_config.payout_transformation.get_payout_transformation_func()
def calculate_user_payment(
self, bp_payout: Decimal, user_wallet_balance: Decimal | None = None
@@ -1392,7 +1387,7 @@ class Product(BaseModel, validate_assignment=True):
if self.payout_config.payout_transformation is None:
return None
payout_xform_func = self.get_payout_transformation_func()
- kwargs = dict()
+ kwargs = {}
if "user_wallet_balance" in inspect.signature(payout_xform_func).parameters:
kwargs["user_wallet_balance"] = user_wallet_balance
user_payout: Decimal = payout_xform_func(bp_payout, **kwargs)
@@ -1419,6 +1414,7 @@ class Product(BaseModel, validate_assignment=True):
def model_dump_mysql(self, *args, **kwargs) -> dict[str, Any]:
d = self.model_dump(mode="json", *args, **kwargs)
+ assert self.created
if "created" in d:
d["created"] = self.created.replace(tzinfo=None)
diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py
index 027aa4c..38ce47e 100644
--- a/generalresearch/models/thl/profiling/marketplace.py
+++ b/generalresearch/models/thl/profiling/marketplace.py
@@ -1,19 +1,19 @@
from __future__ import annotations
from abc import ABC, abstractmethod
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from functools import cached_property
from typing import Any
from pydantic import BaseModel, ConfigDict, Field, PositiveInt, computed_field
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
LanguageISOLike,
UUIDStr,
)
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.locales import CountryISO, LanguageISO
@@ -82,10 +82,9 @@ class MarketplaceQuestion(BaseModel, ABC):
# question has more than 6.
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k == "options":
- if v and len(v) > 6:
- v = v[:3] + ["..."] + v[-3:]
- repr_args[n] = ("options", v)
+ if k == "options" and v and len(v) > 6:
+ v = v[:3] + ["..."] + v[-3:]
+ repr_args[n] = ("options", v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -111,9 +110,7 @@ class MarketplaceUserQuestionAnswer(BaseModel):
# This may be a pipe-separated string if the question_type is multi. Regex
# means any chars except capital letters
option_id: str = Field(pattern=r"^[^A-Z]*$")
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
country_iso: CountryISO = Field(frozen=True)
language_iso: LanguageISO = Field(frozen=True)
diff --git a/generalresearch/models/thl/profiling/other_option.py b/generalresearch/models/thl/profiling/other_option.py
index ae58e90..2f3cac9 100644
--- a/generalresearch/models/thl/profiling/other_option.py
+++ b/generalresearch/models/thl/profiling/other_option.py
@@ -40,7 +40,7 @@ texts_in = {
}
-def option_is_catch_all(c: "UpkQuestionChoice") -> bool:
+def option_is_catch_all(c: UpkQuestionChoice) -> bool:
"""
Exclusive not specifically in the sense that it is a multi-select question
and if this option is selected no others can be selected. But also in the
@@ -51,6 +51,4 @@ def option_is_catch_all(c: "UpkQuestionChoice") -> bool:
return True
if c.text.lower() in texts_exact:
return True
- if any(t in c.text.lower() for t in texts_in):
- return True
- return False
+ return bool(any(t in c.text.lower() for t in texts_in))
diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py
index 9e78692..e46e00a 100644
--- a/generalresearch/models/thl/profiling/upk_property.py
+++ b/generalresearch/models/thl/profiling/upk_property.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from uuid import uuid4
@@ -11,7 +11,7 @@ from generalresearch.models.thl.category import Category
from generalresearch.utils.enum import ReprEnumMeta
-class PropertyType(str, Enum, metaclass=ReprEnumMeta):
+class PropertyType(StrEnum, metaclass=ReprEnumMeta):
# UserProfileKnowledge Item
UPK_ITEM = "i"
# UserProfileKnowledge Numerical
@@ -25,7 +25,7 @@ class PropertyType(str, Enum, metaclass=ReprEnumMeta):
# UPK_DATE = "d"
-class Cardinality(str, Enum, metaclass=ReprEnumMeta):
+class Cardinality(StrEnum, metaclass=ReprEnumMeta):
# Zero or More
ZERO_OR_MORE = "*"
# Zero or One
diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py
index 307bc33..9c7383a 100644
--- a/generalresearch/models/thl/profiling/upk_question.py
+++ b/generalresearch/models/thl/profiling/upk_question.py
@@ -3,9 +3,9 @@ from __future__ import annotations
import hashlib
import json
import re
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
-from typing import Any, List, Literal, Union
+from typing import Annotated, Any, Literal
from pydantic import (
BaseModel,
@@ -16,10 +16,9 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Annotated
-from generalresearch.models import Source
from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.category import Category
@@ -99,7 +98,7 @@ class UpkQuestionChoiceOut(UpkQuestionChoice):
# importance: Optional[UPKImportance] = Field(default=None, exclude=True)
-class UpkQuestionType(str, Enum):
+class UpkQuestionType(StrEnum):
# The question has options that the user must select from. A MC question
# can be e.g. Selector.SINGLE_ANSWER or Selector.MULTIPLE_ANSWER to
# indicate only 1 or more than 1 option can be selected respectively.
@@ -112,7 +111,7 @@ class UpkQuestionType(str, Enum):
HIDDEN = "HIDDEN"
-class UpkQuestionSelector(str, Enum):
+class UpkQuestionSelector(StrEnum):
pass
@@ -215,11 +214,9 @@ SelectorType = (
| UpkQuestionSelectorHIDDEN
)
Configuration = Annotated[
- Union[
- UpkQuestionConfigurationMC,
- UpkQuestionConfigurationTE,
- UpkQuestionConfigurationSLIDER,
- ],
+ UpkQuestionConfigurationMC
+ | UpkQuestionConfigurationTE
+ | UpkQuestionConfigurationSLIDER,
Field(discriminator="type"),
]
@@ -370,7 +367,7 @@ class UpkQuestion(BaseModel):
self.choices is None
), f"No `choices` are allowed for type `{self.type}`"
else:
- assert self.choices is not None, f"`choices` must be set"
+ assert self.choices is not None, "`choices` must be set"
return self
@model_validator(mode="after")
@@ -433,7 +430,7 @@ class UpkQuestion(BaseModel):
@field_validator("choices")
@classmethod
- def order_choices(cls, choices: List):
+ def order_choices(cls, choices: list):
if choices:
choices.sort(key=lambda x: x.order)
return choices
@@ -478,10 +475,9 @@ class UpkQuestion(BaseModel):
# Almost nothing has >1k options, besides location stuff (cities,
# etc.) which should get harmonized. When presenting them, we'll
# filter down options to at most 50.
- if self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000):
- return False
-
- return True
+ return not (
+ self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000)
+ )
@property
def md5sum(self):
@@ -537,7 +533,7 @@ class UpkQuestion(BaseModel):
), "Multiple of the same answer submitted"
if self.type == UpkQuestionType.MULTIPLE_CHOICE:
assert len(answer) >= 1, "MC question with no selected answers"
- choice_codes = set(x.id for x in self.choices)
+ choice_codes = {x.id for x in self.choices}
if self.selector == UpkQuestionSelectorMC.SINGLE_ANSWER:
assert (
len(answer) == 1
@@ -566,9 +562,7 @@ class UpkQuestion(BaseModel):
assert len(answer) == 1, "Only one answer allowed"
answer = answer[0]
assert len(answer) > 0, "Must provide answer"
- max_length = (
- self.configuration.max_length if self.configuration else 0 or 100000
- )
+ max_length = self.configuration.max_length if self.configuration else 100000
assert len(answer) <= max_length, "Answer longer than allowed"
if self.validation and self.validation.patterns:
for pattern in self.validation.patterns:
diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py
index 0024e68..f25baa7 100644
--- a/generalresearch/models/thl/profiling/upk_question_answer.py
+++ b/generalresearch/models/thl/profiling/upk_question_answer.py
@@ -1,7 +1,7 @@
from __future__ import annotations
-from datetime import datetime, timezone
-from typing import Any
+from datetime import UTC, datetime
+from typing import Any, Self
from uuid import uuid4
from pydantic import (
@@ -12,14 +12,13 @@ from pydantic import (
computed_field,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.models import MAX_INT32
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
UUIDStr,
)
+from generalresearch.models.definitions import MAX_INT32
from generalresearch.models.thl.profiling.upk_property import (
Cardinality,
PropertyType,
@@ -60,9 +59,7 @@ class UpkQuestionAnswer(BaseModel):
# ISO 3166-1 alpha-2 (two-letter codes, lowercase)
country_iso: CountryISOLike = Field()
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# If the property is PropertyType.UPK_ITEM, it should have an item (and no value).
# If the property is UPK_NUMERICAL or UPK_TEXT, it'll have a value (and no item).
diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py
index 32af704..46733d4 100644
--- a/generalresearch/models/thl/profiling/user_info.py
+++ b/generalresearch/models/thl/profiling/user_info.py
@@ -3,8 +3,8 @@ from __future__ import annotations
from pydantic import BaseModel, ConfigDict, Field
from pydantic.json_schema import SkipJsonSchema
-from generalresearch.models import Source
from generalresearch.models.custom_types import AwareDatetimeISO
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.profiling.user_question_answer import (
MarketplaceResearchProfileQuestion,
)
diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py
index 8248623..342b65e 100644
--- a/generalresearch/models/thl/profiling/user_question_answer.py
+++ b/generalresearch/models/thl/profiling/user_question_answer.py
@@ -1,8 +1,9 @@
from __future__ import annotations
import json
-from datetime import datetime, timedelta, timezone
-from typing import Any, Iterator, Literal
+from collections.abc import Iterator
+from datetime import UTC, datetime, timedelta
+from typing import Any, Literal
from pydantic import (
BaseModel,
@@ -12,25 +13,20 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.grpc import timestamp_to_datetime
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.locales import CountryISO, LanguageISO
from generalresearch.models.thl.profiling.upk_question import UpkQuestion
class UserQuestionAnswer(BaseModel):
-
model_config = ConfigDict(validate_assignment=True)
user_id: PositiveInt | None = Field(lt=MAX_INT32, default=None)
question_id: UUIDStr = Field()
answer: tuple[str, ...] = Field()
- timestamp: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ timestamp: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
country_iso: CountryISO | Literal["xx"] = Field()
language_iso: LanguageISO | Literal["xxx"] = Field()
@@ -41,20 +37,24 @@ class UserQuestionAnswer(BaseModel):
calc_answers: dict[str, tuple[str, ...]] | None = Field(default=None)
@field_validator("calc_answers")
- def sorted_calc_answers(cls, calc_answers) -> dict[str, tuple[str, ...]] | None:
+ def sorted_calc_answers(
+ cls, calc_answers: dict[str, tuple[str, ...]] | None
+ ) -> dict[str, tuple[str, ...]] | None:
if calc_answers is None:
return None
return {k: tuple(sorted(v)) for k, v in calc_answers.items()}
@field_validator("calc_answers")
- def validate_keys(cls, calc_answers) -> dict[str, tuple[str, ...]] | None:
+ def validate_keys(
+ cls, calc_answers: dict[str, tuple[str, ...]] | None
+ ) -> dict[str, tuple[str, ...]] | None:
if calc_answers is None:
return None
- assert all(
- ":" in k for k in calc_answers.keys()
- ), "calc_answers expects the keys to be in format source:question_code"
+ assert all(":" in k for k in calc_answers), (
+ "calc_answers expects the keys to be in format source:question_code"
+ )
return calc_answers
def model_dump_mysql(self, session_id: str | None = None) -> dict[str, Any]:
@@ -68,6 +68,7 @@ class UserQuestionAnswer(BaseModel):
return d
def get_mrpqs(self) -> Iterator[MarketplaceResearchProfileQuestion]:
+ assert self.calc_answers
for k, v in self.calc_answers.items():
source, question_code = k.split(":", 1)
yield MarketplaceResearchProfileQuestion(
@@ -92,12 +93,12 @@ class UserQuestionAnswer(BaseModel):
"""
try:
assert question.id == self.question_id, "mismatched question id"
- assert (
- question.country_iso == self.country_iso
- ), "country_iso doesn't match question's country"
- assert (
- question.language_iso == self.language_iso
- ), "language_iso doesn't match question's language"
+ assert question.country_iso == self.country_iso, (
+ "country_iso doesn't match question's country"
+ )
+ assert question.language_iso == self.language_iso, (
+ "language_iso doesn't match question's language"
+ )
question._validate_question_answer(self.answer)
except AssertionError as e:
return False, str(e)
@@ -105,22 +106,7 @@ class UserQuestionAnswer(BaseModel):
return True, ""
def is_stale(self) -> bool:
- return self.timestamp < datetime.now(tz=timezone.utc) - timedelta(days=30)
-
- @classmethod
- def from_grpc(cls, msg, default_timestamp: datetime) -> Self:
- """
- Handles correctly issues with grpc timestamps
- :param msg: "thl.protos.generalresearch_pb2.ProfilingQuestionAnswer"
- """
- assert default_timestamp.tzinfo is not None, "must use tz-aware timestamps"
- timestamp = timestamp_to_datetime(msg.timestamp)
- timestamp = default_timestamp if timestamp < datetime(2000, 1, 1) else timestamp
- return cls(
- question_id=msg.question_id,
- answer=tuple(msg.answer),
- timestamp=timestamp,
- )
+ return self.timestamp < datetime.now(tz=UTC) - timedelta(days=30)
# We can't set a redis list to [] vs None. We'll push this dummy answer into
@@ -129,11 +115,11 @@ class UserQuestionAnswer(BaseModel):
DUMMY_UQA = UserQuestionAnswer(
question_id="f118edd01cf1476ba7200a175fb4351d",
answer=("0",),
- timestamp=datetime(2020, 1, 1, tzinfo=timezone.utc),
+ timestamp=datetime(2020, 1, 1, tzinfo=UTC),
country_iso="xx",
language_iso="xxx",
property_code="dummy",
- calc_answers=dict(),
+ calc_answers={},
)
@@ -152,9 +138,9 @@ class MarketplaceResearchProfileQuestion(BaseModel):
@model_validator(mode="after")
def validate_keys(self):
- assert (
- ":" not in self.question_code
- ), "question_code expected to not be in curie format"
+ assert ":" not in self.question_code, (
+ "question_code expected to not be in curie format"
+ )
return self
@property
diff --git a/generalresearch/models/thl/report_task.py b/generalresearch/models/thl/report_task.py
index d29599d..d816420 100644
--- a/generalresearch/models/thl/report_task.py
+++ b/generalresearch/models/thl/report_task.py
@@ -3,11 +3,13 @@ from __future__ import annotations
import random
from collections import defaultdict
from collections.abc import Collection
+from typing import TYPE_CHECKING
from pydantic import BaseModel, ConfigDict, Field
from generalresearch.models.thl.definitions import ReportValue
-from generalresearch.models.thl.user import BPUIDStr
+from generalresearch.models.thl.user_identifiers import BPUIDStr
+
# If a report is made with multiple values, we'll take the one with the
# highest priority
@@ -28,7 +30,7 @@ def prioritize_report_values(
return None
report_values = list(set(report_values))
random.shuffle(report_values)
- return sorted(report_values, key=lambda x: REPORT_PRIORITY[x])[-1]
+ return max(report_values, key=lambda x: REPORT_PRIORITY[x])
class ReportTask(BaseModel):
diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py
index b5e4401..a9fe442 100644
--- a/generalresearch/models/thl/session.py
+++ b/generalresearch/models/thl/session.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import json
import logging
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import TYPE_CHECKING, Annotated, Any
+from typing import TYPE_CHECKING, Annotated, Any, Self
from uuid import uuid4
from pydantic import (
@@ -17,21 +17,15 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
-from generalresearch.models import DeviceType, Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
EnumNameSerializer,
IPvAnyAddressStr,
UUIDStr,
)
+from generalresearch.models.definitions import DeviceType, Source
from generalresearch.models.legacy.bucket import Bucket
-from generalresearch.models.thl import (
- Product,
- decimal_to_int_cents,
- int_cents_to_decimal,
-)
from generalresearch.models.thl.definitions import (
WALL_ALLOWED_STATUS_CODE_1_2,
WALL_ALLOWED_STATUS_STATUS_CODE,
@@ -43,7 +37,12 @@ from generalresearch.models.thl.definitions import (
WallAdjustedStatus,
WallStatusCode2,
)
+from generalresearch.models.thl.product import Product
from generalresearch.models.thl.user import User
+from generalresearch.models.thl.utils import (
+ decimal_to_int_cents,
+ int_cents_to_decimal,
+)
if TYPE_CHECKING:
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
@@ -69,9 +68,7 @@ class WallBase(BaseModel):
buyer_id: str | None = Field(default=None, max_length=32)
req_survey_id: str = Field(max_length=32)
req_cpi: Decimal = Field(decimal_places=5, lt=1000, ge=0)
- started: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# These get set on creation, or updated when the wall event is finished. So
# they shouldn't really ever be NULL, but you don't have to pass them in
@@ -127,9 +124,9 @@ class WallBase(BaseModel):
@classmethod
def check_cpi_decimal_places(cls, v: Decimal) -> Decimal:
if v is not None:
- assert (
- v.as_tuple().exponent >= -5
- ), "Must have 5 or fewer decimal places ('XXX.YYYYY')"
+ assert v.as_tuple().exponent >= -5, (
+ "Must have 5 or fewer decimal places ('XXX.YYYYY')"
+ )
return v
@model_validator(mode="before")
@@ -158,29 +155,27 @@ class WallBase(BaseModel):
@model_validator(mode="after")
def check_timestamps(self):
- assert self.started <= datetime.now(
- tz=timezone.utc
- ), "Started must not be in the future"
+ assert self.started <= datetime.now(tz=UTC), "Started must not be in the future"
if self.finished:
assert self.finished > self.started, "Finished must be after started"
- assert self.finished - self.started <= timedelta(
- minutes=90
- ), "Maximum wall event time is 90 min"
+ assert self.finished - self.started <= timedelta(minutes=90), (
+ "Maximum wall event time is 90 min"
+ )
return self
@model_validator(mode="after")
def check_ext_statuses(self):
if self.ext_status_code_3 is not None:
- assert (
- self.ext_status_code_1 is not None
- ), "Set ext_status_code_1 before ext_status_code_3"
- assert (
- self.ext_status_code_2 is not None
- ), "Set ext_status_code_2 before ext_status_code_3"
+ assert self.ext_status_code_1 is not None, (
+ "Set ext_status_code_1 before ext_status_code_3"
+ )
+ assert self.ext_status_code_2 is not None, (
+ "Set ext_status_code_2 before ext_status_code_3"
+ )
if self.ext_status_code_2 is not None:
- assert (
- self.ext_status_code_1 is not None
- ), "Set ext_status_code_1 before ext_status_code_2"
+ assert self.ext_status_code_1 is not None, (
+ "Set ext_status_code_1 before ext_status_code_2"
+ )
return self
@model_validator(mode="after")
@@ -188,27 +183,27 @@ class WallBase(BaseModel):
if self.status in {Status.COMPLETE, Status.FAIL}:
assert self.finished is not None, "finished should be set"
if self.status == Status.COMPLETE:
- assert (
- self.status_code_1 == StatusCode1.COMPLETE
- ), "status_code_1 should be COMPLETE"
+ assert self.status_code_1 == StatusCode1.COMPLETE, (
+ "status_code_1 should be COMPLETE"
+ )
return self
@model_validator(mode="after")
def check_status_status_code_agreement(self) -> Self:
if self.status_code_1:
options = WALL_ALLOWED_STATUS_STATUS_CODE.get(self.status, {})
- assert (
- self.status_code_1 in options
- ), f"If status is {self.status.value}, status_code_1 should be in {options}"
+ assert self.status_code_1 in options, (
+ f"If status is {self.status.value}, status_code_1 should be in {options}"
+ )
return self
@model_validator(mode="after")
def check_status_code1_2_agreement(self) -> Self:
if self.status_code_2:
options = WALL_ALLOWED_STATUS_CODE_1_2.get(self.status_code_1, {})
- assert (
- self.status_code_2 in options
- ), f"If status_code_1 is {self.status_code_1.value}, status_code_2 should be in {options}"
+ assert self.status_code_2 in options, (
+ f"If status_code_1 is {self.status_code_1.value}, status_code_2 should be in {options}"
+ )
return self
# --- Methods ---
@@ -239,10 +234,7 @@ class WallBase(BaseModel):
return self.is_visible() and self.status == Status.COMPLETE
def allow_session(self) -> bool:
- if self.status == Status.COMPLETE:
- return False
-
- return True
+ return self.status != Status.COMPLETE
def update(self, **kwargs) -> None:
"""
@@ -278,7 +270,7 @@ class WallBase(BaseModel):
# This is just used in tests at the moment. This needs to be adjusted.
if finished is None:
- finished = datetime.now(tz=timezone.utc)
+ finished = datetime.now(tz=UTC)
self.update(
status=status,
@@ -304,16 +296,16 @@ class WallBase(BaseModel):
finished: datetime | None = None,
) -> None:
# This should be called by the wall manager in order to actually update db
- from generalresearch import wall_status_codes
+ from generalresearch.wall_status_codes import annotate_status_code
- status, status_code_1, status_code_2 = wall_status_codes.annotate_status_code(
+ status, status_code_1, status_code_2 = annotate_status_code(
self.source,
ext_status_code_1,
ext_status_code_2,
ext_status_code_3,
)
if finished is None:
- finished = datetime.now(tz=timezone.utc)
+ finished = datetime.now(tz=UTC)
self.update(
status=status,
status_code_1=status_code_1,
@@ -325,18 +317,18 @@ class WallBase(BaseModel):
)
def is_soft_fail(self) -> bool:
- from generalresearch import wall_status_codes
+ from generalresearch.wall_status_codes import is_soft_fail
assert self.status is not None, "status should not be None"
assert self.status_code_1 is not None, "status_code_1 should not be None"
- return wall_status_codes.is_soft_fail(self)
+ return is_soft_fail(self)
def stop_marketplace_session(self) -> bool:
- from generalresearch import wall_status_codes
+ from generalresearch.wall_status_codes import stop_marketplace_session
assert self.status is not None, "status should not be None"
assert self.status_code_1 is not None, "status_code_1 should not be None"
- return wall_status_codes.stop_marketplace_session(self)
+ return stop_marketplace_session(self)
def get_status_after_adjustment(self) -> Status:
if self.adjusted_status in {
@@ -357,10 +349,13 @@ class WallBase(BaseModel):
WallAdjustedStatus.CPI_ADJUSTMENT,
}:
return self.adjusted_cpi
+
elif self.adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL:
return Decimal(0)
+
elif self.status == Status.COMPLETE:
return self.cpi
+
else:
return Decimal(0)
@@ -386,7 +381,7 @@ class WallBase(BaseModel):
TODO: Transition this over to use the ReportTask pydantic model.
"""
report_timestamp = (
- report_timestamp if report_timestamp else datetime.now(tz=timezone.utc)
+ report_timestamp if report_timestamp else datetime.now(tz=UTC)
)
if self.status is None and self.finished is None:
self.status = Status.ABANDON
@@ -419,15 +414,15 @@ class Wall(WallBase):
@model_validator(mode="after")
def check_adjusted_null(self) -> Self:
if self.adjusted_status is not None or self.adjusted_cpi is not None:
- assert (
- self.adjusted_cpi is not None
- ), "Set adjusted_cpi if the wall has been adjusted"
- assert (
- self.adjusted_status is not None
- ), "Set adjusted_status if the wall has been adjusted"
- assert (
- self.adjusted_timestamp is not None
- ), "Set adjusted_timestamp if the wall has been adjusted"
+ assert self.adjusted_cpi is not None, (
+ "Set adjusted_cpi if the wall has been adjusted"
+ )
+ assert self.adjusted_status is not None, (
+ "Set adjusted_status if the wall has been adjusted"
+ )
+ assert self.adjusted_timestamp is not None, (
+ "Set adjusted_timestamp if the wall has been adjusted"
+ )
return self
@model_validator(mode="after")
@@ -439,9 +434,9 @@ class Wall(WallBase):
# --- Properties ---
- @computed_field
+ @computed_field()
@property
- def elapsed(self) -> timedelta:
+ def elapsed(self) -> timedelta | None:
return self.finished - self.started if self.finished else None
def to_json(self) -> str:
@@ -450,9 +445,9 @@ class Wall(WallBase):
d = self.model_dump(mode="json", exclude={"elapsed"})
return json.dumps(d)
- def model_dump_mysql(self, *args, **kwargs) -> dict:
+ def model_dump_mysql(self) -> dict[str, Any]:
# Generate a dictionary representation of the model, with special handling for datetimes
- d = self.model_dump(mode="json", exclude={"elapsed"}, *args, **kwargs)
+ d = self.model_dump(mode="json", exclude={"elapsed"})
d["started"] = self.started.replace(tzinfo=None)
if self.finished:
d["finished"] = self.finished.replace(tzinfo=None)
@@ -462,7 +457,6 @@ class Wall(WallBase):
class WallOut(WallBase):
-
# These get serialized to the enum name instead of the int value (for ease in UI)
status_code_1: Annotated[StatusCode1, EnumNameSerializer] | None = Field(
default=None,
@@ -504,13 +498,13 @@ class WallOut(WallBase):
)
# Serialize user_cpi to an int
- @field_serializer("user_cpi", return_type=int)
- def serialize_user_cpi(self, v: Decimal, _info):
+ @field_serializer("user_cpi", return_type=int | None)
+ def serialize_user_cpi(self, v: Decimal | None) -> int | None:
return decimal_to_int_cents(v)
# If user_cpi is an int, put it back to a decimal
@field_validator("user_cpi", mode="before")
- def deserialize_user_cpi(cls, v):
+ def deserialize_user_cpi(cls, v: Decimal | None) -> Decimal | None:
if isinstance(v, int):
return int_cents_to_decimal(v)
return v
@@ -518,11 +512,11 @@ class WallOut(WallBase):
# noinspection PyNestedDecorators
@field_validator("user_cpi", mode="after")
@classmethod
- def check_cpi_decimal_places(cls, v: Decimal) -> Decimal:
+ def check_cpi_decimal_places(cls, v: Decimal | None) -> Decimal | None:
if v is not None:
- assert (
- v.as_tuple().exponent >= -5
- ), "Must have 5 or fewer decimal places ('XXX.YYYYY')"
+ assert v.as_tuple().exponent >= -5, (
+ "Must have 5 or fewer decimal places ('XXX.YYYYY')"
+ )
return v
@field_validator("status_code_1", mode="before")
@@ -587,9 +581,7 @@ class Session(BaseModel):
id: int | None = None
uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex)
user: User
- started: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ started: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# This is the "bucket" the user clicked on to start this session. We only
# store the 4 fields: loi_min, loi_max, user_payout_min, user_payout_max
@@ -677,9 +669,9 @@ class Session(BaseModel):
@classmethod
def check_payout_decimal_places(cls, v: Decimal) -> Decimal:
if v is not None:
- assert (
- v.as_tuple().exponent >= -2
- ), "Must have 2 or fewer decimal places ('XXX.YY')"
+ assert v.as_tuple().exponent >= -2, (
+ "Must have 2 or fewer decimal places ('XXX.YY')"
+ )
# explicitly make sure it is 2 decimal places, after checking that it is already 2 or less.
v = v.quantize(Decimal("0.00"))
return v
@@ -701,17 +693,21 @@ class Session(BaseModel):
StatusCode1.PS_FAIL,
StatusCode1.PS_QUALITY,
StatusCode1.PS_BLOCKED,
- }, f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ }, (
+ f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ )
elif self.status in {Status.TIMEOUT, Status.ABANDON}:
assert self.status_code_1 in {
StatusCode1.PS_ABANDON,
StatusCode1.GRS_ABANDON,
StatusCode1.BUYER_ABANDON,
- }, f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ }, (
+ f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ )
elif self.status == Status.COMPLETE:
- assert (
- self.status_code_1 == StatusCode1.COMPLETE
- ), f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ assert self.status_code_1 == StatusCode1.COMPLETE, (
+ f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}"
+ )
else:
assert self.status_code_1 is None, (
f"status_code_1 {self.status_code_1.name} invalid for status "
@@ -734,9 +730,9 @@ class Session(BaseModel):
@model_validator(mode="after")
def check_payout_when_complete(self):
if self.status == Status.COMPLETE:
- assert (
- self.payout is not None
- ), "there should be a payout if the session is marked complete"
+ assert self.payout is not None, (
+ "there should be a payout if the session is marked complete"
+ )
return self
# @model_validator(mode='after')
@@ -762,19 +758,19 @@ class Session(BaseModel):
@model_validator(mode="after")
def check_adjusted(self):
if self.adjusted_status is not None or self.adjusted_payout is not None:
- assert (
- self.adjusted_payout is not None
- ), "Set adjusted_payout if the session has been adjusted"
- assert (
- self.adjusted_status is not None
- ), "Set adjusted_status if the session has been adjusted"
- assert (
- self.adjusted_timestamp is not None
- ), "Set adjusted_timestamp if the session has been adjusted"
+ assert self.adjusted_payout is not None, (
+ "Set adjusted_payout if the session has been adjusted"
+ )
+ assert self.adjusted_status is not None, (
+ "Set adjusted_status if the session has been adjusted"
+ )
+ assert self.adjusted_timestamp is not None, (
+ "Set adjusted_timestamp if the session has been adjusted"
+ )
if self.adjusted_user_payout is not None:
- assert (
- self.adjusted_payout is not None
- ), "Set adjusted_payout if adjusted_user_payout is set"
+ assert self.adjusted_payout is not None, (
+ "Set adjusted_payout if adjusted_user_payout is set"
+ )
# NOTE: the other way around is NOT required!
# (the adjusted_user_payout / user_payout can be null)
return self
@@ -787,9 +783,9 @@ class Session(BaseModel):
"the adjusted_status should be null"
)
if self.adjusted_status == SessionAdjustedStatus.ADJUSTED_TO_FAIL:
- assert (
- self.status == Status.COMPLETE
- ), "Session.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL"
+ assert self.status == Status.COMPLETE, (
+ "Session.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL"
+ )
return self
# --- Properties ---
@@ -828,7 +824,7 @@ class Session(BaseModel):
The status of the BP's user_wallet_config.failed_attempt_credit_enabled does
not matter here.
"""
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
min_session_length = timedelta(minutes=1)
if self.status is None:
@@ -850,13 +846,12 @@ class Session(BaseModel):
)
def model_dump_mysql(
- self, *args, **kwargs
+ self, **kwargs
) -> dict[str, str | int | datetime | float | None]:
# Generate a dictionary representation of the model, with special
# handling for datetimes, and nested models such as User & Bucket
-
- d = self.model_dump(mode="json", *args, **kwargs)
+ d = self.model_dump(mode="json", **kwargs)
d["started"] = self.started.replace(tzinfo=None)
if self.finished:
@@ -907,7 +902,7 @@ class Session(BaseModel):
if (
last_wall.status is None
and self.status is None
- and datetime.now(tz=timezone.utc)
+ and datetime.now(tz=UTC)
> self.started + timedelta(seconds=task_timeout_seconds)
):
last_wall.status = Status.TIMEOUT
@@ -988,7 +983,7 @@ class Session(BaseModel):
self, max_session_len: timedelta, max_session_hard_retry: int
) -> bool:
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
last_wall = self.get_last_visible_wall()
if last_wall and last_wall.status == Status.COMPLETE:
@@ -1002,10 +997,7 @@ class Session(BaseModel):
return True
# Hard limit of 40 wall events per session
- if len(self.wall_events) >= 40:
- return True
-
- return False
+ return len(self.wall_events) >= 40
def determine_payments(
self,
@@ -1018,6 +1010,7 @@ class Session(BaseModel):
)
product = self.user.product
+ assert product
# Handle brokerage product payouts
bp_pay: Decimal = product.determine_bp_payment(thl_net)
commission_amount: Decimal = thl_net - bp_pay
@@ -1124,9 +1117,9 @@ class Session(BaseModel):
return False
if self.status == Status.COMPLETE:
- assert (
- self.adjusted_status != SessionAdjustedStatus.ADJUSTED_TO_COMPLETE
- ), "Can't have complete adj to complete"
+ assert self.adjusted_status != SessionAdjustedStatus.ADJUSTED_TO_COMPLETE, (
+ "Can't have complete adj to complete"
+ )
if self.adjusted_status in {
None,
SessionAdjustedStatus.PAYOUT_ADJUSTMENT,
@@ -1270,19 +1263,19 @@ def check_adjusted_status_consistent(
assert adjusted_cpi == cpi, "adjusted_cpi should be equal to the original cpi"
elif adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL:
- assert (
- status == Status.COMPLETE
- ), "Wall.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL"
- assert (
- adjusted_cpi == 0
- ), "adjusted_cpi should be 0 if adjusted_status is ADJUSTED_TO_FAIL"
+ assert status == Status.COMPLETE, (
+ "Wall.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL"
+ )
+ assert adjusted_cpi == 0, (
+ "adjusted_cpi should be 0 if adjusted_status is ADJUSTED_TO_FAIL"
+ )
elif adjusted_status == WallAdjustedStatus.CPI_ADJUSTMENT:
# the original status is allowed to be anything
# the adjusted cpi should be something different
- assert (
- adjusted_cpi != 0 and adjusted_cpi != cpi
- ), "If CPI_ADJUSTMENT, the adjusted_cpi should be different from the original cpi or 0"
+ assert adjusted_cpi != 0 and adjusted_cpi != cpi, (
+ "If CPI_ADJUSTMENT, the adjusted_cpi should be different from the original cpi or 0"
+ )
elif adjusted_status is None:
assert adjusted_cpi is None, "incompatible adjusted values"
@@ -1343,21 +1336,21 @@ def _check_adjusted_status_wall_consistent(
# status / adjusted_status agreement
if status == Status.COMPLETE:
- assert (
- new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_COMPLETE
- ), "adjusted status can't be ADJUSTED_TO_COMPLETE if the status is already COMPLETE"
+ assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_COMPLETE, (
+ "adjusted status can't be ADJUSTED_TO_COMPLETE if the status is already COMPLETE"
+ )
elif status == Status.FAIL:
- assert (
- new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL
- ), "adjusted status can't be ADJUSTED_TO_FAIL if the status is already FAIL"
+ assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL, (
+ "adjusted status can't be ADJUSTED_TO_FAIL if the status is already FAIL"
+ )
else:
# status is None/timeout/abandon, which we treat as a fail anyway
- assert (
- new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL
- ), "attempt is already a failure"
+ assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL, (
+ "attempt is already a failure"
+ )
# adjusted_status / new_adjusted_status agreement
if new_adjusted_status == WallAdjustedStatus.CPI_ADJUSTMENT:
- assert (
- new_adjusted_cpi != adjusted_cpi
- ), f"adjusted_cpi is already {adjusted_cpi}"
+ assert new_adjusted_cpi != adjusted_cpi, (
+ f"adjusted_cpi is already {adjusted_cpi}"
+ )
diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py
index 6ff1165..c0bf2dd 100644
--- a/generalresearch/models/thl/soft_pair.py
+++ b/generalresearch/models/thl/soft_pair.py
@@ -2,11 +2,14 @@ from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
+from typing import TYPE_CHECKING
-from generalresearch.models import Source
-from generalresearch.models.thl.survey.condition import (
- MarketplaceCondition,
-)
+if TYPE_CHECKING:
+ from generalresearch.models.definitions import Source
+ from generalresearch.models.dynata.survey import DynataCondition
+ from generalresearch.models.thl.survey.condition import (
+ MarketplaceCondition,
+ )
class SoftPairResultType(int, Enum):
@@ -34,7 +37,7 @@ class SoftPairResult:
pair_type: SoftPairResultType
source: Source
survey_id: str
- conditions: set[MarketplaceCondition] | None = None
+ conditions: set[MarketplaceCondition | DynataCondition] | None = None
@property
def survey_sid(self) -> str:
@@ -49,7 +52,7 @@ class SoftPairResult:
return (
self.survey_id
+ ":"
- + ";".join(sorted(set([c.question_id for c in self.conditions])))
+ + ";".join(sorted({c.question_id for c in self.conditions}))
)
else:
return None
diff --git a/generalresearch/models/thl/supplier_tag.py b/generalresearch/models/thl/supplier_tag.py
index ad84c9b..739b895 100644
--- a/generalresearch/models/thl/supplier_tag.py
+++ b/generalresearch/models/thl/supplier_tag.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class SupplierTag(str, Enum):
+class SupplierTag(StrEnum):
"""Available tags which can be used to annotate supplier traffic
Note: should not include commas!
diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py
index b6ac740..7749bdb 100644
--- a/generalresearch/models/thl/survey/__init__.py
+++ b/generalresearch/models/thl/survey/__init__.py
@@ -3,12 +3,11 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from decimal import Decimal
from itertools import product
-from typing import Type
from more_itertools import flatten
from pydantic import BaseModel, Field
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.demographics import (
AgeGroup,
DemographicTarget,
@@ -108,11 +107,10 @@ class MarketplaceTask(BaseModel, ABC):
@property
@abstractmethod
- def condition_model(self) -> Type[MarketplaceCondition]:
+ def condition_model(self) -> type[MarketplaceCondition]:
"""
The Condition Model for this survey class
"""
- pass
@property
@abstractmethod
@@ -120,7 +118,6 @@ class MarketplaceTask(BaseModel, ABC):
"""
The age question ID
"""
- pass
@property
@abstractmethod
@@ -130,7 +127,6 @@ class MarketplaceTask(BaseModel, ABC):
"""
Mapping of generic Gender to the marketplace condition for that gender
"""
- pass
@property
def marketplace_age_groups(
diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py
index 6d4d7a1..91102b4 100644
--- a/generalresearch/models/thl/survey/buyer.py
+++ b/generalresearch/models/thl/survey/buyer.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from math import log
from typing import Annotated
@@ -16,12 +16,12 @@ from pydantic import (
)
from scipy.stats import beta as beta_dist
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
UUIDStr,
)
+from generalresearch.models.definitions import Source
class Buyer(BaseModel):
@@ -46,7 +46,7 @@ class Buyer(BaseModel):
)
label: str | None = Field(default=None, max_length=255)
created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When this entry was made, or when the buyer was first seen",
)
@@ -177,9 +177,10 @@ class BuyerCountryStat(BaseModel):
)
# ---- Scoring ----
- score: float = Field(
+ score: float | None = Field(
description="Composite score calculated from all of the individual features",
examples=[-5.329389837486194],
+ default=None,
)
@model_validator(mode="after")
diff --git a/generalresearch/models/thl/survey/condition.py b/generalresearch/models/thl/survey/condition.py
index 927b7e1..90cf27b 100644
--- a/generalresearch/models/thl/survey/condition.py
+++ b/generalresearch/models/thl/survey/condition.py
@@ -4,7 +4,7 @@ import hashlib
from abc import ABC
from enum import Enum
from functools import cached_property
-from typing import Any
+from typing import Annotated, Any, Self
from pydantic import (
BaseModel,
@@ -16,9 +16,8 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Annotated, Self
-from generalresearch.models import LogicalOperator
+from generalresearch.models.definitions import LogicalOperator
MarketplaceConditionHash = Annotated[
str, StringConstraints(min_length=7, max_length=7, pattern=r"^[a-f0-9]+$")
@@ -249,7 +248,7 @@ class MarketplaceCondition(BaseModel, ABC):
return d
@staticmethod
- def is_numeric_including_inf(s) -> bool:
+ def is_numeric_including_inf(s: Any) -> bool:
try:
float(s)
return True
@@ -264,10 +263,9 @@ class MarketplaceCondition(BaseModel, ABC):
# Fancy repr that only shows the first and last 3 values if there are more than 6.
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k == "values":
- if v and len(v) > 6:
- v = v[:3] + ["…"] + v[-3:]
- repr_args[n] = ("values", v)
+ if k == "values" and v and len(v) > 6:
+ v = v[:3] + ["…"] + v[-3:]
+ repr_args[n] = ("values", v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -335,5 +333,5 @@ class MarketplaceCondition(BaseModel, ABC):
except ValueError:
return None
values = self.values_ranges
- passes = any([start <= x <= end for start, end in values for x in answer])
+ passes = any(start <= x <= end for start, end in values for x in answer)
return not passes if self.negate else passes
diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py
index 3794c00..9e3c03d 100644
--- a/generalresearch/models/thl/survey/model.py
+++ b/generalresearch/models/thl/survey/model.py
@@ -1,8 +1,8 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
-from typing import Any
+from typing import Annotated, Any
from pydantic import (
BaseModel,
@@ -15,10 +15,8 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Annotated
from generalresearch.managers.thl.buyer import Buyer
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
@@ -26,6 +24,7 @@ from generalresearch.models.custom_types import (
PropertyCode,
SurveyKey,
)
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.category import Category
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.pagination import Page
@@ -71,12 +70,8 @@ class Survey(BaseModel):
min_length=1, max_length=128, default=None, examples=["124"]
)
- created_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
- updated_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
+ updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
is_live: bool = Field(default=True)
is_recontact: bool = Field(default=False)
@@ -101,12 +96,12 @@ class Survey(BaseModel):
@model_validator(mode="after")
def category_strengths(self):
if any(s.strength is not None for s in self.categories):
- assert all(
- s.strength is not None for s in self.categories
- ), "If any category strength is not None, all should be set"
- assert (
- abs(sum(s.strength for s in self.categories) - 1) <= 0.01
- ), "Strengths should some to 1"
+ assert all(s.strength is not None for s in self.categories), (
+ "If any category strength is not None, all should be set"
+ )
+ assert abs(sum(s.strength for s in self.categories) - 1) <= 0.01, (
+ "Strengths should some to 1"
+ )
return self
def model_dump_sql(self):
@@ -188,9 +183,7 @@ class SurveyStat(BaseModel):
# ---- Metadata ----
- updated_at: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ updated_at: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
@property
def natural_key(self) -> str:
diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py
index 0f544b8..b8ec697 100644
--- a/generalresearch/models/thl/survey/penalty.py
+++ b/generalresearch/models/thl/survey/penalty.py
@@ -1,16 +1,16 @@
from __future__ import annotations
import abc
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
-from generalresearch.models import Source
from generalresearch.models.custom_types import (
AwareDatetimeISO,
UUIDStr,
)
+from generalresearch.models.definitions import Source
class SurveyPenalty(BaseModel, abc.ABC):
@@ -28,9 +28,7 @@ class SurveyPenalty(BaseModel, abc.ABC):
penalty: float = Field(ge=0, le=1)
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
@property
def sid(self):
diff --git a/generalresearch/models/thl/survey/task_collection.py b/generalresearch/models/thl/survey/task_collection.py
index b80166f..d8db0d1 100644
--- a/generalresearch/models/thl/survey/task_collection.py
+++ b/generalresearch/models/thl/survey/task_collection.py
@@ -38,7 +38,8 @@ class TaskCollection(BaseModel):
except pa.errors.SchemaErrors as exc:
idx = exc.failure_cases["index"]
if len(idx) >= len(df) * 0.10:
- raise exc
+ raise
+
logger.info(f"{self.__repr_name__()}:handle_df:{json.dumps(exc.message)}")
df.drop(index=list(idx), inplace=True)
# we need to redo the validation after removing failing rows!
diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py
index 89a3873..404f6ea 100644
--- a/generalresearch/models/thl/task_adjustment.py
+++ b/generalresearch/models/thl/task_adjustment.py
@@ -1,13 +1,13 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from uuid import uuid4
from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.definitions import (
WallAdjustedStatus,
)
@@ -26,11 +26,11 @@ class TaskAdjustmentEvent(BaseModel):
uuid: UUIDStr = Field(default_factory=lambda: uuid4().hex)
created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When this event was created in the db",
)
alerted: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
description="When we were notified about this change",
)
diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py
index ee713a5..8fbb76d 100644
--- a/generalresearch/models/thl/task_status.py
+++ b/generalresearch/models/thl/task_status.py
@@ -1,7 +1,7 @@
from __future__ import annotations
from datetime import datetime
-from typing import Annotated, Any, Literal
+from typing import TYPE_CHECKING, Annotated, Any, Literal
from pydantic import (
BaseModel,
@@ -12,14 +12,12 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import (
AwareDatetimeISO,
EnumNameSerializer,
UUIDStr,
)
-from generalresearch.models.thl import decimal_to_int_cents
from generalresearch.models.thl.definitions import (
SessionAdjustedStatus,
SessionStatusCode2,
@@ -31,11 +29,13 @@ from generalresearch.models.thl.payout_format import (
PayoutFormatOptionalField,
PayoutFormatType,
)
-from generalresearch.models.thl.product import (
- PayoutTransformation,
- Product,
-)
-from generalresearch.models.thl.session import Session, WallOut
+from generalresearch.models.thl.product import PayoutTransformation
+from generalresearch.models.thl.session import WallOut
+from generalresearch.models.thl.utils import decimal_to_int_cents
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
# API uses the ints, b/c this is what the grpc returned originally ...
STATUS_MAP = {
@@ -64,8 +64,7 @@ class TaskStatusResponse(BaseModel):
product_user_id: str = Field(
min_length=3,
max_length=128,
- description="A unique identifier for each user, which is set by the "
- "Supplier",
+ description="A unique identifier for each user, which is set by the Supplier",
examples=["app-user-9329ebd"],
)
@@ -172,12 +171,12 @@ class TaskStatusResponse(BaseModel):
# Serialize enum → int
@field_serializer("status", return_type=int)
- def serialize_status(self, v: Status | None, _info):
+ def serialize_status(self, v: Status | None):
return STATUS_MAP[v]
# Accept int OR string for input, but internally store a Status enum
@field_validator("status", mode="before")
- def deserialize_status(cls, v):
+ def deserialize_status(cls, v: Status | None):
# int → enum
if isinstance(v, int):
return REVERSE_STATUS_MAP[v]
@@ -225,11 +224,12 @@ class TaskStatusResponse(BaseModel):
return v or 0
@field_validator("kwargs", mode="after")
- def sanitize_kwargs(cls, v: dict | None) -> dict | None:
+ def sanitize_kwargs(cls, v: dict[str, Any] | None) -> dict[str, Any] | None:
if v and "clicked_timestamp" in v:
try:
- clicked_timestamp = datetime.strptime(
- v["clicked_timestamp"], "%Y-%m-%d %H:%M:%S.%f"
+ clicked_timestamp = datetime.strptime( # noqa
+ date_string=v["clicked_timestamp"],
+ format="%Y-%m-%d %H:%M:%S.%f",
)
v["clicked_timestamp"] = (
clicked_timestamp.isoformat(timespec="microseconds") + "Z"
@@ -239,7 +239,7 @@ class TaskStatusResponse(BaseModel):
return v
@model_validator(mode="before")
- def transform_user_payout(cls, d):
+ def transform_user_payout(cls, d: dict[str, Any]):
# If the user_payout is None and there is a payout_format, make the user_payout 0
if d.get("user_payout") is None and d.get("payout_format"):
d["user_payout"] = 0
@@ -252,7 +252,7 @@ class TaskStatusResponse(BaseModel):
return self.product_user_id
@classmethod
- def from_session(cls, session: Session, product: Product) -> Self:
+ def from_session(cls, session: Session, product: Product) -> TaskStatusResponse:
user_payout_string = None
if session.user_payout is not None:
@@ -275,7 +275,7 @@ class TaskStatusResponse(BaseModel):
user_payout_string=user_payout_string,
product_id=session.user.product_id,
product_user_id=session.user.product_user_id,
- kwargs=session.url_metadata or dict(),
+ kwargs=session.url_metadata or {},
status_code_1=session.status_code_1,
status_code_2=session.status_code_2,
adjusted_status=session.adjusted_status,
diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py
index 08e59a0..59a14db 100644
--- a/generalresearch/models/thl/user.py
+++ b/generalresearch/models/thl/user.py
@@ -2,42 +2,42 @@ from __future__ import annotations
import json
import logging
-import re
-from datetime import datetime, timezone
-from typing import TYPE_CHECKING, Annotated
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING, Any, Self
from uuid import UUID, uuid4
from pydantic import (
- AfterValidator,
AwareDatetime,
BaseModel,
ConfigDict,
Field,
PositiveInt,
- StringConstraints,
field_validator,
model_validator,
)
from sentry_sdk import set_tag, set_user
-from typing_extensions import Self
-from generalresearch.models import MAX_INT32
-from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.custom_types import (
+ AwareDatetimeISO,
+ UUIDStr,
+)
+from generalresearch.models.definitions import MAX_INT32
from generalresearch.models.thl.ipinfo import GeoIPInformation
from generalresearch.models.thl.ledger import LedgerTransaction
from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user_identifiers import BPUIDStr
+from generalresearch.models.thl.user_ref import UserRef
from generalresearch.models.thl.userhealth import AuditLog
-from generalresearch.pg_helper import PostgresConfig
if TYPE_CHECKING:
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
)
from generalresearch.managers.thl.userhealth import AuditLogManager
+ from generalresearch.pg_helper import PostgresConfig
-logger = logging.getLogger()
-BPUID_ALLOWED = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[\]^_{|}~"
+logger = logging.getLogger()
class User(BaseModel):
@@ -102,42 +102,28 @@ class User(BaseModel):
)
# --- Validation ---
- @field_validator("product_user_id")
- def check_product_user_id(cls, v: str) -> str:
- if v is not None:
- if " " in v:
- raise ValueError("String cannot contain spaces")
- if "\\" in v:
- raise ValueError("String cannot contain backslash")
- if "/" in v:
- raise ValueError("String cannot contain slash")
- # I think the * on the regex messes up value matches that are
- # the same length as the
- rex = re.fullmatch("[" + BPUID_ALLOWED + "]*", v)
- if not bool(rex):
- raise ValueError("String is not valid regex")
- return v
# noinspection PyNestedDecoratorsk
@field_validator("created", "last_seen")
@classmethod
- def check_not_in_future(cls, v: AwareDatetime) -> AwareDatetime:
+ def check_not_in_future(cls, v: AwareDatetime | None) -> AwareDatetime | None:
if v is not None:
try:
- assert v < datetime.now(tz=timezone.utc)
- except Exception:
+ assert v < datetime.now(tz=UTC)
+ except AssertionError:
raise ValueError("Input is in the future")
return v
# noinspection PyNestedDecorators
@field_validator("created", "last_seen")
@classmethod
- def check_after_anno_domini(cls, v: AwareDatetime) -> AwareDatetime:
+ def check_after_anno_domini(cls, v: AwareDatetime | None) -> AwareDatetime | None:
if v is not None:
try:
- assert v > datetime(year=2016, month=7, day=13, tzinfo=timezone.utc)
- except Exception:
+ assert v > datetime(year=2016, month=7, day=13, tzinfo=UTC)
+ except AssertionError:
raise ValueError("Input is before Anno Domini")
+
return v
@model_validator(mode="after")
@@ -167,7 +153,7 @@ class User(BaseModel):
)
@classmethod
- def is_valid_ubp(cls, *, product_id, product_user_id) -> bool:
+ def is_valid_ubp(cls, *, product_id: str, product_user_id: str) -> bool:
# Attempt to create common_struct solely for validation purposes,
# using the product_id and product_user_id
try:
@@ -177,7 +163,7 @@ class User(BaseModel):
product_id=product_id,
product_user_id=product_user_id,
)
- except Exception as e:
+ except ValueError as e:
logger.info(e)
return False
else:
@@ -185,7 +171,7 @@ class User(BaseModel):
# --- Methods ---
@staticmethod
- def check_bpuid_is_not_bpid(product_id, product_user_id):
+ def check_bpuid_is_not_bpid(product_id: str | None, product_user_id: str | None):
"""Unfortunately users were already created failing this constraint,
so only check for new users!
"""
@@ -197,7 +183,7 @@ class User(BaseModel):
raise ValueError("product_user_id must not equal the product_id")
return True
- def to_dict(self) -> dict:
+ def to_dict(self) -> dict[str, Any]:
return self.model_dump(mode="python", exclude={"product"})
def to_json(self) -> str:
@@ -205,6 +191,16 @@ class User(BaseModel):
d["user_id"] = self.user_id
return json.dumps(d)
+ def to_user_ref(self) -> UserRef:
+ assert self.user_id is not None
+ assert self.product_id is not None
+ assert self.product_user_id is not None
+ return UserRef(
+ user_id=self.user_id,
+ product_id=self.product_id,
+ product_user_id=self.product_user_id,
+ )
+
def set_sentry_user(self):
# https://docs.sentry.io/platforms/python/enriching-events/identify-user/
set_user(
@@ -248,7 +244,7 @@ class User(BaseModel):
# # Delete from db.thl-marketplaces
# We need DELETE credentials for all these...
- # from generalresearch.models import Source
+ # from generalresearch.models.definitions import Source
# mp_db_table = {
# Source.SPECTRUM: "`thl-spectrum`.`spectrum_marketresearchprofilequestion`",
# Source.INNOVATE: "`thl-innovate`.`innovate_marketresearchprofilequestion`",
@@ -290,11 +286,11 @@ class User(BaseModel):
# --- Prebuild ---
@classmethod
- def from_db(cls, res) -> Self:
+ def from_db(cls, res: dict[str, Any]) -> Self:
if res["created"]:
- res["created"] = res["created"].replace(tzinfo=timezone.utc)
+ res["created"] = res["created"].replace(tzinfo=UTC)
if res["last_seen"]:
- res["last_seen"] = res["last_seen"].replace(tzinfo=timezone.utc)
+ res["last_seen"] = res["last_seen"].replace(tzinfo=UTC)
res["product_id"] = UUID(res["product_id"]).hex
res["uuid"] = UUID(res["uuid"]).hex
return cls(
@@ -308,10 +304,4 @@ class User(BaseModel):
)
-# Used in other places where the bpuid is part of a model that's used in
-# the API (separate from a User)
-BPUIDStr = Annotated[
- str,
- StringConstraints(min_length=3, max_length=128),
- AfterValidator(User.check_product_user_id),
-]
+User.model_rebuild()
diff --git a/generalresearch/models/thl/user_identifiers.py b/generalresearch/models/thl/user_identifiers.py
new file mode 100644
index 0000000..57a2970
--- /dev/null
+++ b/generalresearch/models/thl/user_identifiers.py
@@ -0,0 +1,33 @@
+import re
+from typing import Annotated
+
+from pydantic import (
+ AfterValidator,
+ StringConstraints,
+)
+
+BPUID_ALLOWED = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[\]^_{|}~"
+
+
+def validate_product_user_id(v: str) -> str:
+ if " " in v:
+ raise ValueError("String cannot contain spaces")
+ if "\\" in v:
+ raise ValueError("String cannot contain backslash")
+ if "/" in v:
+ raise ValueError("String cannot contain slash")
+ # I think the * on the regex messes up value matches that are
+ # the same length as the
+ rex = re.fullmatch("[" + BPUID_ALLOWED + "]*", v)
+ if not bool(rex):
+ raise ValueError("String is not valid regex")
+ return v
+
+
+# Used in other places where the bpuid is part of a model that's used in
+# the API (separate from a User)
+BPUIDStr = Annotated[
+ str,
+ StringConstraints(min_length=3, max_length=128),
+ AfterValidator(validate_product_user_id),
+]
diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py
index 5892d41..a7eadf4 100644
--- a/generalresearch/models/thl/user_iphistory.py
+++ b/generalresearch/models/thl/user_iphistory.py
@@ -1,7 +1,8 @@
from __future__ import annotations
import ipaddress
-from datetime import datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING, Self
from faker import Faker
from pydantic import (
@@ -11,21 +12,20 @@ from pydantic import (
PositiveInt,
field_validator,
)
-from typing_extensions import Self
from generalresearch.models.custom_types import (
AwareDatetimeISO,
CountryISOLike,
IPvAnyAddressStr,
)
-from generalresearch.models.thl.ipinfo import (
- GeoIPInformation,
- normalize_ip,
-)
-from generalresearch.models.thl.maxmind.definitions import UserType
+from generalresearch.models.thl.ipinfo import GeoIPInformation, normalize_ip
from generalresearch.models.thl.user import User
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+
+if TYPE_CHECKING:
+ from grip_client.enums import AccessType
+
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
fake = Faker()
@@ -53,7 +53,11 @@ class UserIPRecord(BaseModel):
)
@property
- def user_type(self) -> UserType | None:
+ def user_type(self) -> AccessType | None:
+ return self.information.user_type if self.information else None
+
+ @property
+ def access_type(self) -> AccessType | None:
return self.information.user_type if self.information else None
@property
@@ -113,7 +117,7 @@ class IPRecord(BaseModel):
# --- ORM ---
@classmethod
def from_mysql(cls, d: dict) -> Self:
- created = d["created"].replace(tzinfo=timezone.utc)
+ created = d["created"].replace(tzinfo=UTC)
d["created"] = created
d["forwarded_ip_records"] = []
@@ -169,7 +173,7 @@ class UserIPHistory(BaseModel):
def ips_timestamp(cls, ips):
if ips is None:
return None
- cutoff = datetime.now(tz=timezone.utc) - timedelta(days=28)
+ cutoff = datetime.now(tz=UTC) - timedelta(days=28)
return sorted(
[x for x in ips if x.created > cutoff],
key=lambda x: x.created,
@@ -204,8 +208,6 @@ class UserIPHistory(BaseModel):
if res.get(x.ip):
x.information = res[x.ip]
- return None
-
def collapse_ip_records(self):
"""
- Records where sequential ipv6 addresses are in the same /64 block,
diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py
index 13e8af1..9df7b23 100644
--- a/generalresearch/models/thl/user_profile.py
+++ b/generalresearch/models/thl/user_profile.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import hashlib
-from typing import Annotated, Any
+from typing import Annotated, Any, Self
from pydantic import (
BaseModel,
@@ -12,10 +12,9 @@ from pydantic import (
computed_field,
)
from pydantic.json_schema import SkipJsonSchema
-from typing_extensions import Self
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.user import User
from generalresearch.models.thl.user_streak import UserStreak
diff --git a/generalresearch/models/thl/user_quality_event.py b/generalresearch/models/thl/user_quality_event.py
index d6ebddc..ab82999 100644
--- a/generalresearch/models/thl/user_quality_event.py
+++ b/generalresearch/models/thl/user_quality_event.py
@@ -1,16 +1,16 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from typing import Literal
from pydantic import BaseModel, Field, PositiveInt
-from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
+from generalresearch.models.definitions import MAX_INT32, Source
from generalresearch.models.thl.definitions import WallAdjustedStatus
-from generalresearch.models.thl.user import BPUIDStr
+from generalresearch.models.thl.user_identifiers import BPUIDStr
from generalresearch.utils.enum import ReprEnumMeta
"""
@@ -18,7 +18,7 @@ Typically used internally. These affect a user's quality standing.
"""
-class QualityEventType(str, Enum, metaclass=ReprEnumMeta):
+class QualityEventType(StrEnum, metaclass=ReprEnumMeta):
"""
Currently, the grpc call SendUserQualityEvents handles both the
recons/task adj, access control, and "security/hash failure" events.
@@ -61,9 +61,7 @@ class TaskAdjustmentEvent(BaseModel):
mid: UUIDStr = Field()
source: Source = Field()
status: WallAdjustedStatus = Field()
- alert_time: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ alert_time: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
quality_event_type: Literal[QualityEventType.task_adjustment] = Field(
default=QualityEventType.task_adjustment
)
diff --git a/generalresearch/models/thl/user_ref.py b/generalresearch/models/thl/user_ref.py
new file mode 100644
index 0000000..9622abe
--- /dev/null
+++ b/generalresearch/models/thl/user_ref.py
@@ -0,0 +1,17 @@
+from pydantic import BaseModel, PositiveInt
+
+from generalresearch.models.custom_types import UUIDStr
+from generalresearch.models.thl.user_identifiers import BPUIDStr
+
+
+class UserRef(BaseModel):
+ """
+ Use in place of the full User model in places where we want to
+ associate something with a User, but can't use the full User
+ model due to cyclic import issues.
+ As a side-effect, this also avoids the type|None ruff issues.
+ """
+
+ user_id: PositiveInt
+ product_id: UUIDStr
+ product_user_id: BPUIDStr
diff --git a/generalresearch/models/thl/user_streak.py b/generalresearch/models/thl/user_streak.py
index 278809b..9105b27 100644
--- a/generalresearch/models/thl/user_streak.py
+++ b/generalresearch/models/thl/user_streak.py
@@ -1,7 +1,8 @@
from __future__ import annotations
from datetime import date, datetime, timedelta
-from enum import Enum
+from enum import StrEnum
+from zoneinfo import ZoneInfo
import pandas as pd
from pydantic import (
@@ -15,14 +16,13 @@ from pydantic import (
model_validator,
)
from pydantic.json_schema import SkipJsonSchema
-from zoneinfo import ZoneInfo
from generalresearch.managers.leaderboard import country_timezone
-from generalresearch.models import MAX_INT32
+from generalresearch.models.definitions import MAX_INT32
from generalresearch.models.thl.locales import CountryISO
-class StreakPeriod(str, Enum):
+class StreakPeriod(StrEnum):
# Midnight to midnight in the tz associated with the user's country
DAY = "day"
# Sunday midnight - sunday midnight
@@ -31,7 +31,7 @@ class StreakPeriod(str, Enum):
MONTH = "month"
-class StreakFulfillment(str, Enum):
+class StreakFulfillment(StrEnum):
"""
What has to happen for a user to fulfill a period for a streak
"""
@@ -42,7 +42,7 @@ class StreakFulfillment(str, Enum):
COMPLETE = "complete"
-class StreakState(str, Enum):
+class StreakState(StrEnum):
# The activity for today was completed!
ACTIVE = "active"
# They had activity yesterday, but not today, and can still continue today
diff --git a/generalresearch/models/thl/userhealth.py b/generalresearch/models/thl/userhealth.py
index e556dc8..8535dfa 100644
--- a/generalresearch/models/thl/userhealth.py
+++ b/generalresearch/models/thl/userhealth.py
@@ -1,11 +1,10 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from enum import Enum
-from typing import Dict, Optional
+from typing import Any, Self
from pydantic import BaseModel, Field, NonNegativeFloat, PositiveInt
-from typing_extensions import Self
from generalresearch.models.custom_types import AwareDatetimeISO
@@ -26,12 +25,12 @@ class AuditLog(BaseModel):
are related to a User
"""
- id: Optional[PositiveInt] = Field(default=None)
+ id: PositiveInt | None = Field(default=None)
user_id: PositiveInt = Field()
created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc),
- examples=[datetime.now(tz=timezone.utc)],
+ default_factory=lambda: datetime.now(tz=UTC),
+ examples=[datetime.now(tz=UTC)],
description="When did this event occur",
)
@@ -51,14 +50,14 @@ class AuditLog(BaseModel):
# e.g. "upk-audit", "ip-audit", "entrance-limit"
event_type: str = Field(max_length=64, examples=["entrance-limit"])
- event_msg: Optional[str] = Field(
+ event_msg: str | None = Field(
default=None,
min_length=3,
max_length=256,
description="The event message. Could be displayed on user's page",
)
- event_value: Optional[NonNegativeFloat] = Field(
+ event_value: NonNegativeFloat | None = Field(
default=None,
description="Optionally store a numeric value associated with this "
"event. For e.g. if we recalculate the user's normalized "
@@ -68,12 +67,12 @@ class AuditLog(BaseModel):
examples=[0.42],
)
- def model_dump_mysql(self, **kwargs) -> Dict:
+ def model_dump_mysql(self, **kwargs) -> dict:
d = self.model_dump(mode="json", **kwargs)
d["created"] = self.created.replace(tzinfo=None)
return d
@classmethod
- def from_mysql(cls, d: Dict) -> Self:
- d["created"] = d["created"].replace(tzinfo=timezone.utc)
- return AuditLog.model_validate(d)
+ def from_mysql(cls, d: dict[str, Any]) -> Self:
+ d["created"] = d["created"].replace(tzinfo=UTC)
+ return cls.model_validate(d)
diff --git a/generalresearch/models/thl/utils.py b/generalresearch/models/thl/utils.py
new file mode 100644
index 0000000..3e14065
--- /dev/null
+++ b/generalresearch/models/thl/utils.py
@@ -0,0 +1,11 @@
+from decimal import Decimal
+
+
+def decimal_to_int_cents(usd: Decimal | None) -> int | None:
+ return round(usd * 100) if usd is not None else None
+
+
+def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None:
+ if value is None:
+ return None
+ return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals)
diff --git a/generalresearch/models/thl/wallet/__init__.py b/generalresearch/models/thl/wallet/__init__.py
index 1403ddf..e69de29 100644
--- a/generalresearch/models/thl/wallet/__init__.py
+++ b/generalresearch/models/thl/wallet/__init__.py
@@ -1,87 +0,0 @@
-from enum import Enum
-
-from generalresearch.utils.enum import ReprEnumMeta
-
-
-class PayoutType(str, Enum, metaclass=ReprEnumMeta):
- """
- The method in which the requested payout is delivered.
- """
-
- # The max size of the db field that holds this value is 14, so please
- # don't add new values longer than that!
-
- # User is paid out to their personal PayPal email address
- PAYPAL = "PAYPAL"
- # User is paid out via a Tango Gift Card
- TANGO = "TANGO"
- # DWOLLA
- DWOLLA = "DWOLLA"
- # A payment is made to a bank account using ACH
- ACH = "ACH"
- # A payment is made to a bank account using ACH
- WIRE = "WIRE"
- # A payment is made in cash and mailed to the user.
- CASH_IN_MAIL = "CASH_IN_MAIL"
- # A payment is made as a prize with some monetary value
- PRIZE = "PRIZE"
-
- # This is used to designate either AMT_BONUS or AMT_HIT
- AMT = "AMT"
- # Amazon Mechanical Turk as a Bonus
- AMT_BONUS = "AMT_BONUS"
- # Amazon Mechanical Turk for a HIT
- AMT_HIT = "AMT_ASSIGNMENT"
- AMT_ASSIGNMENT = "AMT_ASSIGNMENT"
-
-
-class Currency(str, Enum):
- # United States Dollar
- USD = "USD"
- # Canadian Dollar
- CAD = "CAD"
- # British Pound Sterling
- GBP = "GBP"
- # Euro
- EUR = "EUR"
- # Indian Rupee
- INR = "INR"
- # Australian Dollar
- AUD = "AUD"
- # Polish Zloty
- PLN = "PLN"
- # Swedish Krona
- SEK = "SEK"
- # Singapore Dollar
- SGD = "SGD"
- # Mexican Peso
- MXN = "MXN"
-
-
-CURRENCY_FORMATTER = {
- "USD": lambda x: f"${x / 100:,.2f}",
- "CAD": lambda x: f"${x / 100:,.2f} CAD",
- "GBP": lambda x: f"{x / 100:,.2f} £",
- "EUR": lambda x: f"€{x / 100:,.2f}",
- "INR": lambda x: f"₹{x / 100:,.2f}",
- "AUD": lambda x: f"${x / 100:,.2f} AUD",
- "PLN": lambda x: f"{x / 100:,.2f} zł",
- "SEK": lambda x: f"{x / 100:,.2f} kr",
- "SGD": lambda x: f"${x / 100:,.2f} SGD",
- "MXN": lambda x: f"${x / 100:,.2f} MXN",
-}
-
-# The max value user can redeem in one go in foreign currencies. should be < $250
-# in order to avoid exchange rate issues
-CURRENCY_MAX_VALUE = {
- "USD": 250,
- "CAD": 200,
- "GBP": 100,
- "EUR": 100,
- "INR": 10000,
- "AUD": 200,
- "PLN": 500,
- "SEK": 1000,
- "SGD": 200,
- "MXN": 4000,
-}
diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py
index 59cf721..6ad0bf2 100644
--- a/generalresearch/models/thl/wallet/cashout_method.py
+++ b/generalresearch/models/thl/wallet/cashout_method.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import hashlib
import logging
-from datetime import datetime, timezone
-from enum import Enum
-from typing import Any, Literal
+from datetime import UTC, datetime
+from enum import StrEnum
+from typing import Any, Literal, Self
from pydantic import (
BaseModel,
@@ -16,7 +16,6 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing_extensions import Self
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import (
@@ -27,8 +26,9 @@ from generalresearch.models.custom_types import (
from generalresearch.models.legacy.api_status import StatusResponse
from generalresearch.models.thl.definitions import PayoutStatus
from generalresearch.models.thl.locales import CountryISO
-from generalresearch.models.thl.user import BPUIDStr, User
-from generalresearch.models.thl.wallet import Currency, PayoutType
+from generalresearch.models.thl.user_identifiers import BPUIDStr
+from generalresearch.models.thl.user_ref import UserRef
+from generalresearch.models.thl.wallet.definitions import Currency, PayoutType
from generalresearch.utils.enum import ReprEnumMeta
logger = logging.getLogger()
@@ -131,34 +131,31 @@ class CashoutMethodBase(BaseModel):
f"Invalid amount requested: ${amount / 100:.2f}. Must be between"
f" ${int(self.min_value) / 100:.2f} and ${int(self.max_value) / 100:.2f}"
)
- if self.type == PayoutType.CASH_IN_MAIL:
- if amount % 500 != 0:
- raise ValueError("Amount must be in increments of $5.00")
+ if self.type == PayoutType.CASH_IN_MAIL and amount % 500 != 0:
+ raise ValueError("Amount must be in increments of $5.00")
return True
class CashoutMethod(CashoutMethodBase):
- user: User | None = Field(
+ user: UserRef | None = Field(
default=None,
description="If set, this cashout method is custom for this user. For example"
"a user may have a paypal cashout method with their paypal"
"email associated.",
)
- last_updated: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ last_updated: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
is_live: bool = Field(default=True)
@model_validator(mode="after")
def validate_user(self) -> Self:
if self.type in {PayoutType.PAYPAL, PayoutType.CASH_IN_MAIL}:
- assert (
- self.user is not None
- ), "user_id must be set for this cashout method type"
+ assert self.user is not None, (
+ "user_id must be set for this cashout method type"
+ )
else:
- assert (
- self.user is None
- ), "user_id must NOT be set for this cashout method type"
+ assert self.user is None, (
+ "user_id must NOT be set for this cashout method type"
+ )
return self
@@ -249,7 +246,7 @@ class CashoutMethodsResponse(StatusResponse):
cashout_methods: list[CashoutMethodOut] = Field()
-class DeliveryStatus(str, Enum):
+class DeliveryStatus(StrEnum):
PENDING = "Pending"
SHIPPED = "Shipped"
IN_TRANSIT = "In Transit"
@@ -261,14 +258,14 @@ class DeliveryStatus(str, Enum):
LOST = "Lost"
-class ShippingCarrier(str, Enum):
+class ShippingCarrier(StrEnum):
USPS = "USPS"
FEDEX = "FedEx"
UPS = "UPS"
DHL = "DHL"
-class ShippingMethod(str, Enum):
+class ShippingMethod(StrEnum):
STANDARD = "Standard"
EXPRESS = "Express"
TWO_DAY = "Two-Day"
@@ -306,8 +303,7 @@ class CashMailOrderData(BaseModel):
default=None,
min_length=1,
max_length=50,
- description="Current status of delivery, e.g., pending, in "
- "transit, delivered",
+ description="Current status of delivery, e.g., pending, in transit, delivered",
)
last_updated: AwareDatetimeISO | None = Field(
default=None,
@@ -394,7 +390,7 @@ example_foreign_value = {
}
-class RedemptionCurrency(str, Enum, metaclass=ReprEnumMeta):
+class RedemptionCurrency(StrEnum, metaclass=ReprEnumMeta):
"""
Supported Currencies for Foreign Redemptions
"""
diff --git a/generalresearch/models/thl/wallet/definitions.py b/generalresearch/models/thl/wallet/definitions.py
new file mode 100644
index 0000000..2d1eb8d
--- /dev/null
+++ b/generalresearch/models/thl/wallet/definitions.py
@@ -0,0 +1,87 @@
+from enum import StrEnum
+
+from generalresearch.utils.enum import ReprEnumMeta
+
+
+class PayoutType(StrEnum, metaclass=ReprEnumMeta):
+ """
+ The method in which the requested payout is delivered.
+ """
+
+ # The max size of the db field that holds this value is 14, so please
+ # don't add new values longer than that!
+
+ # User is paid out to their personal PayPal email address
+ PAYPAL = "PAYPAL"
+ # User is paid out via a Tango Gift Card
+ TANGO = "TANGO"
+ # DWOLLA
+ DWOLLA = "DWOLLA"
+ # A payment is made to a bank account using ACH
+ ACH = "ACH"
+ # A payment is made to a bank account using ACH
+ WIRE = "WIRE"
+ # A payment is made in cash and mailed to the user.
+ CASH_IN_MAIL = "CASH_IN_MAIL"
+ # A payment is made as a prize with some monetary value
+ PRIZE = "PRIZE"
+
+ # This is used to designate either AMT_BONUS or AMT_HIT
+ AMT = "AMT"
+ # Amazon Mechanical Turk as a Bonus
+ AMT_BONUS = "AMT_BONUS"
+ # Amazon Mechanical Turk for a HIT
+ AMT_HIT = "AMT_ASSIGNMENT"
+ AMT_ASSIGNMENT = "AMT_ASSIGNMENT"
+
+
+class Currency(StrEnum):
+ # United States Dollar
+ USD = "USD"
+ # Canadian Dollar
+ CAD = "CAD"
+ # British Pound Sterling
+ GBP = "GBP"
+ # Euro
+ EUR = "EUR"
+ # Indian Rupee
+ INR = "INR"
+ # Australian Dollar
+ AUD = "AUD"
+ # Polish Zloty
+ PLN = "PLN"
+ # Swedish Krona
+ SEK = "SEK"
+ # Singapore Dollar
+ SGD = "SGD"
+ # Mexican Peso
+ MXN = "MXN"
+
+
+CURRENCY_FORMATTER = {
+ "USD": lambda x: f"${x / 100:,.2f}",
+ "CAD": lambda x: f"${x / 100:,.2f} CAD",
+ "GBP": lambda x: f"{x / 100:,.2f} £",
+ "EUR": lambda x: f"€{x / 100:,.2f}",
+ "INR": lambda x: f"₹{x / 100:,.2f}",
+ "AUD": lambda x: f"${x / 100:,.2f} AUD",
+ "PLN": lambda x: f"{x / 100:,.2f} zł",
+ "SEK": lambda x: f"{x / 100:,.2f} kr",
+ "SGD": lambda x: f"${x / 100:,.2f} SGD",
+ "MXN": lambda x: f"${x / 100:,.2f} MXN",
+}
+
+# The max value user can redeem in one go in foreign currencies. should be < $250
+# in order to avoid exchange rate issues
+CURRENCY_MAX_VALUE = {
+ "USD": 250,
+ "CAD": 200,
+ "GBP": 100,
+ "EUR": 100,
+ "INR": 10000,
+ "AUD": 200,
+ "PLN": 500,
+ "SEK": 1000,
+ "SGD": 200,
+ "MXN": 4000,
+}
diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py
index 8c78bef..1fc0f77 100644
--- a/generalresearch/models/thl/wallet/payout.py
+++ b/generalresearch/models/thl/wallet/payout.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import json
from collections.abc import Collection
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import Any
from uuid import uuid4
@@ -17,10 +17,10 @@ from pydantic import (
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailOrderData,
)
+from generalresearch.models.thl.wallet.definitions import PayoutType
class PayoutEvent(BaseModel, validate_assignment=True):
@@ -50,9 +50,7 @@ class PayoutEvent(BaseModel, validate_assignment=True):
# populated from the db and so does not need to be set (there is no
# `description` field in event_payout)
description: str | None = Field(default=None)
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
- )
+ created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
# In the smallest unit of the currency being transacted. For USD, this
# is cents.
@@ -131,13 +129,16 @@ class PayoutEvent(BaseModel, validate_assignment=True):
else:
raise ValueError("this shouldn't happen")
- def model_dump_mysql(self, *args, **kwargs) -> dict[str, Any]:
- d = self.model_dump(mode="json", *args, **kwargs)
+ def model_dump_mysql(self) -> dict[str, Any]:
+ d = self.model_dump(mode="json")
+
if "created" in d:
d["created"] = self.created.replace(tzinfo=None)
if d.get("request_data") is not None:
d["request_data"] = json.dumps(self.request_data)
if d.get("order_data") is not None:
+ assert self.order_data
+
if isinstance(self.order_data, dict):
d["order_data"] = json.dumps(self.order_data)
else:
@@ -159,7 +160,7 @@ class BPPayoutEvent(BaseModel):
created: AwareDatetimeISO = Field(
description="When the Brokerage Product was paid out",
- default_factory=lambda: datetime.now(tz=timezone.utc),
+ default_factory=lambda: datetime.now(tz=UTC),
)
amount: USDCent = Field(
diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py
index b9a7d79..b1ac556 100644
--- a/generalresearch/pg_helper.py
+++ b/generalresearch/pg_helper.py
@@ -1,12 +1,12 @@
from __future__ import annotations
-from datetime import timezone
+from datetime import UTC
import psycopg
-from psycopg.adapt import Buffer
+from psycopg.abc import Query
from psycopg.rows import RowFactory, dict_row
from psycopg.types.datetime import TimestampLoader
-from psycopg.types.net import Address, InetLoader, Interface
+from psycopg.types.net import InetLoader
from psycopg.types.string import TextLoader
from psycopg.types.uuid import UUIDLoader
from pydantic import PostgresDsn
@@ -24,7 +24,7 @@ class UTCTimestampLoader(TimestampLoader):
if dt is None:
return None
assert dt.tzinfo is None, "expected naive dt"
- return dt.replace(tzinfo=timezone.utc)
+ return dt.replace(tzinfo=UTC)
class BPCharLoader(TextLoader):
@@ -75,7 +75,9 @@ class PostgresConfig:
self.row_factory = row_factory
@property
- def db(self):
+ def db(self) -> str:
+ assert self.dsn
+ assert self.dsn.path
return self.dsn.path[1:]
def make_connection(self) -> psycopg.Connection:
@@ -99,21 +101,18 @@ class PostgresConfig:
conn.adapters.register_loader("inet", InetHostLoader)
return conn
- def execute_sql_query(self, query, params=None):
+ def execute_sql_query(self, query: Query, params=None):
# This is only intended for SELECT queries
- assert "SELECT" in query.upper(), "Supports SELECTs only"
+ assert "SELECT" in str(query).upper(), "Supports SELECTs only"
- with self.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query=query, params=params)
- return c.fetchall()
+ with self.make_connection() as conn, conn.cursor() as c:
+ c.execute(query=query, params=params)
+ return c.fetchall()
- def execute_write(self, query, params=None) -> int:
+ def execute_write(self, query: Query, params=None) -> int:
cmd = query.lstrip().upper()
- assert (
- cmd.startswith("INSERT")
- or cmd.startswith("UPDATE")
- or cmd.startswith("DELETE")
+ assert cmd.startswith(
+ ("INSERT", "UPDATE", "DELETE")
), "Supports INSERT/UPDATE only"
with self.make_connection() as conn:
diff --git a/generalresearch/schemas/survey_stats.py b/generalresearch/schemas/survey_stats.py
index 0d509e1..dd592d4 100644
--- a/generalresearch/schemas/survey_stats.py
+++ b/generalresearch/schemas/survey_stats.py
@@ -2,7 +2,7 @@ import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
from generalresearch.locales import Localelator
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
COUNTRY_ISOS = Localelator().get_all_countries()
kosovo = "xk"
@@ -70,7 +70,7 @@ UnitInterval = Column(
SID_CHECKS = [
Check.str_length(min_value=3, max_value=67),
- Check.str_matches("^[a-z]{1,2}\:[A-Za-z0-9]+"),
+ Check.str_matches(r"^[a-z]{1,2}\:[A-Za-z0-9]+"),
Check(
lambda x: len(set(x.str.split(":").str[0])) == 1,
error="the sources must all be the same",
diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py
index 2c526f3..ef53c3b 100644
--- a/generalresearch/sql_helper.py
+++ b/generalresearch/sql_helper.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
-from typing import Any, Optional
+from typing import Any
from uuid import UUID
from pydantic import MariaDBDsn, MySQLDsn, PostgresDsn
@@ -14,6 +14,9 @@ ListOrTupleOfListOrTuple = (
DataBaseDsn = MySQLDsn | MariaDBDsn | PostgresDsn | None
+logging.basicConfig()
+logger = logging.getLogger(__name__)
+
class MultipleObjectsReturned(Exception):
pass
@@ -39,14 +42,7 @@ class SqlConnector:
# I'm intentionally doing a match case here so that we'll make sure
# we can NOT use this on old versions of python 😈
- if "mysql" in self.dsn.scheme:
- import pymysql as engine_module
-
- self.engine_module = engine_module
- self.cursor_class = engine_module.cursors.DictCursor
- self.quote_char = "`"
-
- elif "maria" in self.dsn.scheme:
+ if "mysql" in self.dsn.scheme or "maria" in self.dsn.scheme:
import pymysql as engine_module
self.engine_module = engine_module
@@ -130,7 +126,7 @@ def decode_uuids(row: dict[str, Any]) -> dict[str, Any]:
class SqlHelper(SqlConnector):
- def __init__(self, dsn: Optional[DataBaseDsn] = None, **kwargs):
+ def __init__(self, dsn: DataBaseDsn | None = None, **kwargs):
super().__init__(dsn, **kwargs)
def execute_sql_query(
@@ -138,7 +134,7 @@ class SqlHelper(SqlConnector):
) -> list[dict[str, Any]]:
for param in params if params else []:
if isinstance(param, (tuple, list, set)) and len(param) == 0:
- logging.warning("param is empty. not executing query")
+ logger.warning("param is empty. not executing query")
return []
connection = self.make_connection()
c = connection.cursor()
@@ -173,7 +169,7 @@ class SqlHelper(SqlConnector):
:param cursor: If cursor is passed, the insert is NOT committed!
:param ignore_existing: adds 'ON CONFLICT DO NOTHING' to SQL statement.
"""
- assert len(set([len(x) for x in values_to_insert])) == 1
+ assert len({len(x) for x in values_to_insert}) == 1
if cursor is None:
connection = self.make_connection()
c = connection.cursor()
@@ -199,8 +195,6 @@ class SqlHelper(SqlConnector):
if cursor is None:
c.connection.commit()
- return None
-
def bulk_update(
self,
table_name: str,
@@ -211,7 +205,7 @@ class SqlHelper(SqlConnector):
if len(values_to_insert) == 0:
return
- assert len(set([len(x) for x in values_to_insert])) == 1
+ assert len({len(x) for x in values_to_insert}) == 1
if cursor is None:
connection = self.make_connection()
c = connection.cursor()
@@ -233,7 +227,7 @@ class SqlHelper(SqlConnector):
if cursor is None:
c.connection.commit()
- return None
+ return
def get_or_create(
self,
@@ -253,7 +247,7 @@ class SqlHelper(SqlConnector):
lookup_fns = ",".join(
["`" + x + "`" for x in set(lookup_dict.keys()) | {primary_key}]
)
- lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict.keys()])
+ lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict])
table_name_str = self._quote(table_name)
query = f"SELECT {lookup_fns} FROM {table_name_str} WHERE {lookup_vals} LIMIT 2"
if cursor is None:
@@ -281,7 +275,7 @@ class SqlHelper(SqlConnector):
cursor=None,
commit=True,
primary_key=None,
- ) -> Optional[int]:
+ ) -> int | None:
"""
Create the item in table `table_name`.
In postgresql, `primary_key` needs to be given in order to return the
@@ -293,7 +287,7 @@ class SqlHelper(SqlConnector):
else:
c = cursor
field_names = ",".join(map(self._quote, create_dict))
- vals = ",".join([f"%({fn})s" for fn in create_dict.keys()])
+ vals = ",".join([f"%({fn})s" for fn in create_dict])
table_name_str = self._quote(table_name)
query = f"INSERT INTO {table_name_str} ({field_names}) VALUES ({vals})"
c.execute(query, create_dict)
@@ -324,7 +318,7 @@ class SqlHelper(SqlConnector):
field_names = ["`" + x + "`" for x in field_names]
field_name_str = ",".join(field_names)
if filter_d:
- lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d.keys()])
+ lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d])
lookup_str = f" WHERE {lookup_vals}"
else:
lookup_str = ""
@@ -351,5 +345,3 @@ class SqlHelper(SqlConnector):
c.execute(query)
if cursor is None:
c.connection.commit()
-
- return None
diff --git a/generalresearch/thl_django/app/manage.py b/generalresearch/thl_django/app/manage.py
index 33f2367..dabd5b3 100644
--- a/generalresearch/thl_django/app/manage.py
+++ b/generalresearch/thl_django/app/manage.py
@@ -1,11 +1,7 @@
#!/usr/bin/env python
-import os
import sys
if __name__ == "__main__":
- os.environ.setdefault(
- "DJANGO_SETTINGS_MODULE", "generalresearch.thl_django.app.settings"
- )
from django.core.management import execute_from_command_line
execute_from_command_line(sys.argv)
diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py
new file mode 100644
index 0000000..4168513
--- /dev/null
+++ b/generalresearch/thl_django/app/test_settings.py
@@ -0,0 +1,17 @@
+DATABASES = {
+ "default": {
+ "ENGINE": "django.db.backends.postgresql",
+ "NAME": 'unittest-2026-09-03-38dbfc',
+ "USER": 'jenkins',
+ "PASSWORD": '123456789',
+ "HOST": 'unittest-postgresql.fmt2.grl.internal',
+ "PORT": 5432,
+ }
+}
+INSTALLED_APPS = ['django.contrib.postgres', 'django.contrib.contenttypes', 'generalresearch.thl_django']
+DEFAULT_AUTO_FIELD = "django.db.models.BigAutoField"
+LANGUAGE_CODE = "en-us"
+TIME_ZONE = "UTC"
+USE_I18N = True
+USE_L10N = True
+USE_TZ = True
diff --git a/generalresearch/thl_django/apps.py b/generalresearch/thl_django/apps.py
index 2813947..bd87110 100644
--- a/generalresearch/thl_django/apps.py
+++ b/generalresearch/thl_django/apps.py
@@ -6,11 +6,11 @@ class THLSchemaConfig(AppConfig):
label = "thl_django"
def ready(self):
- from .accounting import models # noqa: F401 # pycharm: keep
- from .common import models # noqa: F401 # pycharm: keep
- from .contest import models # noqa: F401 # pycharm: keep
- from .event import models # noqa: F401 # pycharm: keep
- from .marketplace import models # noqa: F401 # pycharm: keep
- from .network import models # noqa: F401 # pycharm: keep
- from .userhealth import models # noqa: F401 # pycharm: keep
+ from .accounting import models # pycharm: keep
+ from .common import models # pycharm: keep
+ from .contest import models # pycharm: keep
+ from .event import models # pycharm: keep
+ from .marketplace import models # pycharm: keep
+ from .network import models # pycharm: keep
+ from .userhealth import models # pycharm: keep
from .userprofile import models # noqa: F401 # pycharm: keep
diff --git a/generalresearch/thl_django/event/models.py b/generalresearch/thl_django/event/models.py
index 51a8e2f..6d92b20 100644
--- a/generalresearch/thl_django/event/models.py
+++ b/generalresearch/thl_django/event/models.py
@@ -51,7 +51,11 @@ class SupplierPayout(models.Model):
to the Business.
"""
- uuid = models.UUIDField(default=uuid.uuid4, primary_key=True)
+ id = models.BigAutoField(primary_key=True)
+
+ # Used for holding a unique, external, payouttype-specific identifier.
+ # For ACH, this is the ACH transaction id.
+ ext_ref_id = models.CharField(max_length=64, unique=True)
# The Business receiving this payout
business_id = models.UUIDField(null=True)
@@ -64,10 +68,6 @@ class SupplierPayout(models.Model):
# generalresearch/models/thl/payout.py:PayoutStatus
status = models.CharField(max_length=20, null=True)
- # Used for holding an external, payouttype-specific identifier.
- # For ACH, this is the ACH transaction id.
- ext_ref_id = models.CharField(max_length=64, null=True)
-
# The allowed values for `payout_type` are defined in generalresearch:
# generalresearch/models/thl/payout.py:PayoutType
payout_type = models.CharField(max_length=14)
@@ -86,7 +86,6 @@ class SupplierPayout(models.Model):
indexes = [
models.Index(fields=["created"]),
models.Index(fields=["business_id"]),
- models.Index(fields=["ext_ref_id"]),
]
@@ -133,7 +132,6 @@ class Payout(models.Model):
# payout transaction that this product-level split belongs to.
supplier_payout = models.ForeignKey(
SupplierPayout,
- db_column="supplier_payout_uuid",
null=True,
on_delete=models.DO_NOTHING,
)
@@ -145,5 +143,4 @@ class Payout(models.Model):
models.Index(fields=["created"]),
models.Index(fields=["debit_account_uuid"]),
models.Index(fields=["ext_ref_id"]),
- models.Index(fields=["supplier_payout"]),
]
diff --git a/generalresearch/thl_django/fields.py b/generalresearch/thl_django/fields.py
index 5e40ef0..251faa5 100644
--- a/generalresearch/thl_django/fields.py
+++ b/generalresearch/thl_django/fields.py
@@ -1,6 +1,7 @@
-from django.db import models
import ipaddress
+from django.db import models
+
class CIDRField(models.Field):
description = "PostgreSQL CIDR network"
diff --git a/generalresearch/thl_django/migrations/0001_initial.py b/generalresearch/thl_django/migrations/0001_initial.py
index ecae35a..cf147e3 100644
--- a/generalresearch/thl_django/migrations/0001_initial.py
+++ b/generalresearch/thl_django/migrations/0001_initial.py
@@ -1,7 +1,8 @@
# Generated by Django 6.0 on 2025-12-26 20:53
-import django.db.models.deletion
import uuid
+
+import django.db.models.deletion
from django.db import migrations, models
diff --git a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py
index 211c48a..f767afc 100644
--- a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py
+++ b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py
@@ -1,7 +1,7 @@
# Generated by Django 6.0 on 2025-12-28 16:49
-from django.db import migrations, models
from django.contrib.postgres.operations import AddIndexConcurrently
+from django.db import migrations, models
class Migration(migrations.Migration):
diff --git a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py
index ecaf0a9..dcf9ef2 100644
--- a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py
+++ b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py
@@ -1,10 +1,10 @@
# Generated by Django 6.0 on 2025-12-29 21:22
-from django.db import migrations, models
from django.contrib.postgres.operations import (
AddIndexConcurrently,
RemoveIndexConcurrently,
)
+from django.db import migrations, models
class Migration(migrations.Migration):
diff --git a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py
index e2492ab..64338c8 100644
--- a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py
+++ b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py
@@ -1,7 +1,7 @@
# Generated by Django 6.0 on 2026-01-02 17:38
-from django.db import migrations
from django.contrib.postgres.operations import RemoveIndexConcurrently
+from django.db import migrations
class Migration(migrations.Migration):
diff --git a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py
index e19a353..e8ac2c2 100644
--- a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py
+++ b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py
@@ -1,13 +1,14 @@
# Generated by Django 6.0 on 2026-03-15 20:17
+import uuid
+
import django.contrib.postgres.indexes
import django.db.models.deletion
import django.utils.timezone
from django.contrib.postgres.operations import CreateExtension
+from django.db import migrations, models
import generalresearch.thl_django.fields
-import uuid
-from django.db import migrations, models
class Migration(migrations.Migration):
diff --git a/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py b/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py
new file mode 100644
index 0000000..5c3319c
--- /dev/null
+++ b/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py
@@ -0,0 +1,51 @@
+# Generated by Django 6.1 on 2026-09-01 23:15
+
+import django.db.models.deletion
+from django.db import migrations, models
+
+
+class Migration(migrations.Migration):
+
+ dependencies = [
+ (
+ "thl_django",
+ "0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more",
+ ),
+ ]
+
+ operations = [
+ migrations.CreateModel(
+ name="SupplierPayout",
+ fields=[
+ ("id", models.BigAutoField(primary_key=True, serialize=False)),
+ ("ext_ref_id", models.CharField(max_length=64, unique=True)),
+ ("business_id", models.UUIDField(null=True)),
+ ("created", models.DateTimeField(auto_now_add=True)),
+ ("amount", models.BigIntegerField()),
+ ("status", models.CharField(max_length=20, null=True)),
+ ("payout_type", models.CharField(max_length=14)),
+ ("request_data", models.JSONField(null=True)),
+ ("order_data", models.JSONField(null=True)),
+ ],
+ options={
+ "db_table": "supplier_payout",
+ "indexes": [
+ models.Index(
+ fields=["created"], name="supplier_pa_created_336236_idx"
+ ),
+ models.Index(
+ fields=["business_id"], name="supplier_pa_busines_2c7a4e_idx"
+ ),
+ ],
+ },
+ ),
+ migrations.AddField(
+ model_name="payout",
+ name="supplier_payout",
+ field=models.ForeignKey(
+ null=True,
+ on_delete=django.db.models.deletion.DO_NOTHING,
+ to="thl_django.supplierpayout",
+ ),
+ ),
+ ]
diff --git a/generalresearch/thl_django/network/models.py b/generalresearch/thl_django/network/models.py
index 167af02..733c0ab 100644
--- a/generalresearch/thl_django/network/models.py
+++ b/generalresearch/thl_django/network/models.py
@@ -1,12 +1,11 @@
from uuid import uuid4
-from django.utils import timezone
-from django.contrib.postgres.indexes import GistIndex, GinIndex
+from django.contrib.postgres.indexes import GinIndex, GistIndex
from django.db import models
+from django.utils import timezone
from generalresearch.thl_django.fields import CIDRField
-
#######
# ** Signals **
# ToolRun
diff --git a/generalresearch/utils/aggregation.py b/generalresearch/utils/aggregation.py
index 4023dc9..ef3a52d 100644
--- a/generalresearch/utils/aggregation.py
+++ b/generalresearch/utils/aggregation.py
@@ -1,8 +1,8 @@
from collections import defaultdict
-from typing import Any, Dict, List
+from typing import Any
-def group_by_year(records: List[Dict], datetime_field: str) -> Dict[int, List[Any]]:
+def group_by_year(records: list[dict], datetime_field: str) -> dict[int, list[Any]]:
"""Memory efficient - processes records one at a time"""
by_year = defaultdict(list)
diff --git a/generalresearch/utils/enum.py b/generalresearch/utils/enum.py
index 14a31de..b4620b3 100644
--- a/generalresearch/utils/enum.py
+++ b/generalresearch/utils/enum.py
@@ -3,7 +3,6 @@ from __future__ import annotations
import inspect
import re
from enum import EnumMeta
-from typing import Dict
class ReprEnumMeta(EnumMeta):
@@ -20,7 +19,7 @@ class ReprEnumMeta(EnumMeta):
[f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self]
)
else:
- return f"\nAllowed values: \n" + "\n".join(
+ return "\nAllowed values: \n" + "\n".join(
[f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self]
)
@@ -36,12 +35,12 @@ class ReprEnumMeta(EnumMeta):
[f" - __{e.name}__: {descriptions[e.name]}" for e in self]
)
else:
- return f"\nAllowed values: \n" + "\n".join(
+ return "\nAllowed values: \n" + "\n".join(
[f" - __{e.name}__: {descriptions[e.name]}" for e in self]
)
-def get_enum_comments(enum_class) -> Dict:
+def get_enum_comments(enum_class) -> dict:
source = inspect.getsource(enum_class)
# Regular expression to match multi-line comments and enum values
pattern = re.compile(r"((?:\s*#.*?\n)+)\s*(\w+)\s*=")
diff --git a/generalresearch/utils/grpc_logger.py b/generalresearch/utils/grpc_logger.py
index 59f7471..8f2f454 100644
--- a/generalresearch/utils/grpc_logger.py
+++ b/generalresearch/utils/grpc_logger.py
@@ -33,9 +33,11 @@ try:
response = handler_func(request, context)
code = context.code() or grpc.StatusCode.OK
return response
- except Exception as e:
+
+ except Exception:
code = context.code() or grpc.StatusCode.INTERNAL
- raise e
+ raise
+
finally:
duration_ms = int((time.time() - start_time) * 1000)
peer = context.peer() or "unknown"
diff --git a/generalresearch/wall_status_codes/__init__.py b/generalresearch/wall_status_codes/__init__.py
index 3a0abb8..cca1a19 100644
--- a/generalresearch/wall_status_codes/__init__.py
+++ b/generalresearch/wall_status_codes/__init__.py
@@ -1,8 +1,7 @@
-from typing import Optional, Tuple
+from typing import TYPE_CHECKING
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import Status, StatusCode1
-from generalresearch.models.thl.session import Wall
from generalresearch.wall_status_codes import (
cint,
dynata,
@@ -18,13 +17,16 @@ from generalresearch.wall_status_codes import (
spectrum,
)
+if TYPE_CHECKING:
+ from generalresearch.models.thl.session import Wall
+
def annotate_status_code(
source: Source,
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, Optional[StatusCode1], Optional[str]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1 | None, str | None]:
"""
:params ext_status_code_1: marketplace-dependent code
:params ext_status_code_2: marketplace-dependent code
diff --git a/generalresearch/wall_status_codes/cint.py b/generalresearch/wall_status_codes/cint.py
index 8042cd2..13028dc 100644
--- a/generalresearch/wall_status_codes/cint.py
+++ b/generalresearch/wall_status_codes/cint.py
@@ -1,14 +1,16 @@
-from typing import Any, Optional, Tuple
+from typing import TYPE_CHECKING, Any
-from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.wall_status_codes import lucid
+if TYPE_CHECKING:
+ from generalresearch.models.thl.definitions import Status, StatusCode1
+
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
return lucid.annotate_status_code(
ext_status_code_1=ext_status_code_1,
ext_status_code_2=ext_status_code_2,
diff --git a/generalresearch/wall_status_codes/dynata.py b/generalresearch/wall_status_codes/dynata.py
index 958e06f..5d008c1 100644
--- a/generalresearch/wall_status_codes/dynata.py
+++ b/generalresearch/wall_status_codes/dynata.py
@@ -4,11 +4,11 @@ checked by Greg 2023-10-10
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_name: Dict[str, str] = {
+status_codes_name: dict[str, str] = {
"0.0": "Unknown",
"0.1": "Missing Language",
"0.2": "Missing Respondent ID",
@@ -51,10 +51,10 @@ status_codes_name: Dict[str, str] = {
"5.10": "Daily Limit",
}
-status_map: Dict[str, Status] = defaultdict(
+status_map: dict[str, Status] = defaultdict(
lambda: Status.FAIL, **{"1.0": Status.COMPLETE, "1.1": Status.COMPLETE}
)
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["1.0", "1.1"],
StatusCode1.BUYER_FAIL: ["2.2", "3.2"],
StatusCode1.BUYER_QUALITY_FAIL: ["5.1", "5.2"],
@@ -88,10 +88,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
"5.10",
],
}
-ext_status_code_map: Dict[str, StatusCode1] = dict()
+ext_status_code_map: dict[str, StatusCode1] = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -100,9 +100,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: this is from the callback url params:
disposition and status, '.'-joined
@@ -117,16 +117,15 @@ def annotate_status_code(
return status, status_code, None
-def stop_marketplace_session(status_code_1: StatusCode1, ext_status_code_1) -> bool:
+def stop_marketplace_session(
+ status_code_1: StatusCode1, ext_status_code_1: str
+) -> bool:
if ext_status_code_1.startswith("5"):
# '5.10' is the user hit a Daily Limit, so they should not be sent in again today
return True
- if status_code_1 in {
+ return status_code_1 in {
StatusCode1.PS_QUALITY,
StatusCode1.BUYER_QUALITY_FAIL,
StatusCode1.PS_BLOCKED,
- }:
- return True
-
- return False
+ }
diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py
index eda4d9c..cd9fdff 100644
--- a/generalresearch/wall_status_codes/fullcircle.py
+++ b/generalresearch/wall_status_codes/fullcircle.py
@@ -7,11 +7,11 @@ we'll try to infer based on the time spent in survey.
from collections import defaultdict
from datetime import timedelta
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_map: Dict[str, str] = {
+status_codes_map: dict[str, str] = {
"1": "Complete",
"2": "Terminate",
"3": "Over-quota",
@@ -19,7 +19,7 @@ status_codes_map: Dict[str, str] = {
}
status_map = defaultdict(lambda: Status.FAIL, **{"1": Status.COMPLETE})
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["1"],
StatusCode1.BUYER_FAIL: ["2", "3"],
StatusCode1.BUYER_QUALITY_FAIL: ["4"],
@@ -29,10 +29,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_FAIL: [],
StatusCode1.PS_OVERQUOTA: [],
}
-ext_status_code_map: Dict[str, StatusCode1] = dict()
+ext_status_code_map: dict[str, StatusCode1] = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -41,9 +41,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: this is from the callback url param 's'
:params ext_status_code_2: not used
diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py
index b650ade..e3d2468 100644
--- a/generalresearch/wall_status_codes/innovate.py
+++ b/generalresearch/wall_status_codes/innovate.py
@@ -9,11 +9,11 @@ can map directly, and some we have to look at the category.
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_innovate: Dict[str, str] = {
+status_codes_innovate: dict[str, str] = {
"1": "Complete",
"2": "Buyer Fail",
"3": "Buyer Over Quota",
@@ -29,7 +29,7 @@ status_map = defaultdict(
lambda: Status.FAIL,
**{"1": Status.COMPLETE, "0": Status.ABANDON, "6": Status.ABANDON},
)
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.BUYER_FAIL: ["2", "3"],
StatusCode1.BUYER_QUALITY_FAIL: ["4"],
StatusCode1.PS_BLOCKED: [],
@@ -38,12 +38,12 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_FAIL: ["5"],
StatusCode1.PS_OVERQUOTA: ["7"],
}
-ext_status_code_map = dict()
+ext_status_code_map = {}
for k, v in status_codes_ext_map.items():
for vv in v:
ext_status_code_map[status_codes_ext_map.get(vv, vv)] = k
-category_innovate: Dict[str, StatusCode1] = {
+category_innovate: dict[str, StatusCode1] = {
"Selected threat potential score at joblevel not allow the survey": StatusCode1.PS_QUALITY,
"OE Validation": StatusCode1.PS_QUALITY,
"Unique IP": StatusCode1.PS_DUPLICATE,
@@ -78,9 +78,9 @@ category_innovate: Dict[str, StatusCode1] = {
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
Only quality terminate (4 and 8), and PS term (5) return a term_reason (af=).
diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py
index c0098e6..c4c5e90 100644
--- a/generalresearch/wall_status_codes/lucid.py
+++ b/generalresearch/wall_status_codes/lucid.py
@@ -5,11 +5,11 @@ https://support.lucidhq.com/s/article/Collecting-Data-From-Redirects
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-mp_codes: Dict[str, str] = {
+mp_codes: dict[str, str] = {
"-6": "Pre-Client Intermediary Page Drop Off",
"-5": "Failure in the Post Answer Behavior",
"-1": "Failure to Load the Lucid Marketplace",
@@ -54,15 +54,15 @@ mp_codes: Dict[str, str] = {
}
# todo: finish, there's a bunch more
-client_status_map: Dict[str, StatusCode1] = {
+client_status_map: dict[str, StatusCode1] = {
"30": StatusCode1.BUYER_QUALITY_FAIL,
"33": StatusCode1.BUYER_QUALITY_FAIL,
"34": StatusCode1.BUYER_QUALITY_FAIL,
"35": StatusCode1.BUYER_QUALITY_FAIL,
}
-status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE})
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE)
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: [],
StatusCode1.BUYER_FAIL: ["3"],
StatusCode1.BUYER_QUALITY_FAIL: [],
@@ -102,10 +102,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_OVERQUOTA: ["40", "41", "42"],
}
-ext_status_code_map: Dict[str, StatusCode1] = dict()
+ext_status_code_map: dict[str, StatusCode1] = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -115,9 +115,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: this indicates which callback url was hit. possible values {'s', *anything else*}
:params ext_status_code_2: this is from the callback url params: InitialStatus
diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py
index 318d1b2..6f63b82 100644
--- a/generalresearch/wall_status_codes/morning.py
+++ b/generalresearch/wall_status_codes/morning.py
@@ -1,5 +1,5 @@
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
@@ -17,7 +17,7 @@ timeout: The respondent completed the survey after the timeout period had expire
in_progress: The respondent interview session is still in progress, such as in the prescreener or survey.
"""
-short_code_to_status_codes_morning: Dict[str, str] = {
+short_code_to_status_codes_morning: dict[str, str] = {
"att_che": "attention_check",
"banned": "banned",
"bid_clo": "bid_closed",
@@ -52,9 +52,9 @@ short_code_to_status_codes_morning: Dict[str, str] = {
"sur_tim": "survey_timeout",
"tem_ban": "temporarily_banned",
}
-status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE})
+status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE)
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["complete"],
StatusCode1.BUYER_FAIL: [
"in_survey_failure",
@@ -97,10 +97,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
"quota_invalid_for_bid",
],
}
-ext_status_code_map: Dict[str, StatusCode1] = dict()
+ext_status_code_map: dict[str, StatusCode1] = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -109,9 +109,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: from callback url params: &sti={{status_id}}
:params ext_status_code_2: from callback url params: &sdi={{status_detail_id}}
diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py
index 361785d..e1ad12a 100644
--- a/generalresearch/wall_status_codes/pollfish.py
+++ b/generalresearch/wall_status_codes/pollfish.py
@@ -1,9 +1,9 @@
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_map: Dict[str, str] = {
+status_codes_map: dict[str, str] = {
"quo_ful": "quota_full",
"sur_clo": "survey_closed",
"profilin": "profiling",
@@ -28,8 +28,8 @@ status_codes_map: Dict[str, str] = {
"su_al_ta": "survey_already_taken",
"complete": "complete",
}
-status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE})
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE)
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["complete"],
StatusCode1.BUYER_FAIL: ["third_party_termination", "screenout"],
StatusCode1.BUYER_QUALITY_FAIL: [
@@ -58,10 +58,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
],
StatusCode1.PS_OVERQUOTA: ["quota_full", "survey_closed", "survey_expired"],
}
-ext_status_code_map = dict()
+ext_status_code_map = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -70,9 +70,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: from callback url params: &sti={{status_id}}
:params ext_status_code_2: from callback url params: &sdi={{status_detail_id}}
diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py
index ffeaeca..4466fa6 100644
--- a/generalresearch/wall_status_codes/precision.py
+++ b/generalresearch/wall_status_codes/precision.py
@@ -7,11 +7,11 @@ f - client approved the Preliminary complete as Final Complete
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_precision: Dict[str, str] = {
+status_codes_precision: dict[str, str] = {
"10": "Complete",
"20": "Client Terminate",
"21": "PS Terminate",
@@ -45,8 +45,8 @@ status_codes_precision: Dict[str, str] = {
"60": "Client Reject",
"80": "Final Complete",
}
-status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE})
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE)
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["10"],
StatusCode1.BUYER_FAIL: ["20", "30"],
StatusCode1.BUYER_QUALITY_FAIL: ["60"],
@@ -73,10 +73,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_FAIL: ["21", "22"],
StatusCode1.PS_OVERQUOTA: ["31", "32", "23"],
}
-ext_status_code_map = dict()
+ext_status_code_map = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -85,9 +85,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: from callback url params: status
:params ext_status_code_2: from callback url params: code
diff --git a/generalresearch/wall_status_codes/prodege.py b/generalresearch/wall_status_codes/prodege.py
index a1e25ea..058ce28 100644
--- a/generalresearch/wall_status_codes/prodege.py
+++ b/generalresearch/wall_status_codes/prodege.py
@@ -3,12 +3,12 @@ https://developer.prodege.com/surveys-feed/term-reasons
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
status_map = defaultdict(lambda: Status.FAIL, **{"1": Status.COMPLETE})
-status_code_map: Dict[StatusCode1, List[str]] = {
+status_code_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: [],
StatusCode1.BUYER_FAIL: ["1", "2"],
StatusCode1.BUYER_QUALITY_FAIL: ["10", "12"],
@@ -31,10 +31,10 @@ status_code_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_OVERQUOTA: ["13", "28", "29", "30", "31", "38"],
}
-status_class = dict()
+status_class = {}
for k, v in status_code_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -43,9 +43,9 @@ for k, v in status_code_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: status from redirect url
:params ext_status_code_2: termreason from redirect url
diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py
index c502532..0828b4c 100644
--- a/generalresearch/wall_status_codes/repdata.py
+++ b/generalresearch/wall_status_codes/repdata.py
@@ -3,11 +3,11 @@ Status codes are in a xlsx file. See thl-repdata readme
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_name: Dict[str, str] = {
+status_codes_name: dict[str, str] = {
"2": "Search Failed",
"3": "Activity Failed",
"4": "Review Failed",
@@ -26,7 +26,7 @@ status_codes_name: Dict[str, str] = {
"6003": "In-Survey maximum exceeded (Research Desk)",
}
# See: 02, and 13 are de-dupes
-rd_threat_name: Dict[str, str] = {
+rd_threat_name: dict[str, str] = {
"02": "Duplicate entrant into survey",
"03": "Emulator Usage",
"04": "VPN usage detected",
@@ -46,8 +46,8 @@ rd_threat_name: Dict[str, str] = {
"18": "MaxMind Failure",
}
-status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE})
-status_code_map: Dict[StatusCode1, List[str]] = {
+status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE)
+status_code_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["1000"],
StatusCode1.BUYER_FAIL: ["2000", "4000"],
StatusCode1.BUYER_QUALITY_FAIL: ["3000"],
@@ -58,10 +58,10 @@ status_code_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_OVERQUOTA: ["5001", "6001", "6002", "6003"],
}
-status_class = dict()
+status_class = {}
for k, v in status_code_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -70,9 +70,9 @@ for k, v in status_code_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: the redirect urls category (as defined in url param 549f3710b)
{'term', 'overquota', 'fraud', 'complete'}
diff --git a/generalresearch/wall_status_codes/sago.py b/generalresearch/wall_status_codes/sago.py
index 9c8710f..1a55a91 100644
--- a/generalresearch/wall_status_codes/sago.py
+++ b/generalresearch/wall_status_codes/sago.py
@@ -3,11 +3,11 @@ https://developer-beta.market-cube.com/api-details#api=definition-api&operation=
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_schlesinger: Dict[str, str] = {
+status_codes_schlesinger: dict[str, str] = {
"1": "Complete",
"2": "Buyer Fail",
"3": "Buyer Fail",
@@ -20,7 +20,7 @@ status_codes_schlesinger: Dict[str, str] = {
"11": "Abandon", # really it is "Buyer Abandon"
}
-status_reason_name: Dict[str, str] = {
+status_reason_name: dict[str, str] = {
"1": "Not a Unique Sample Cube User",
"4": "GeoIP - wrong country",
"7": "Duplicate - not a unique IP",
@@ -121,7 +121,7 @@ status_map = defaultdict(
lambda: Status.FAIL, **{"1": Status.COMPLETE, "0": Status.ABANDON}
)
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["48"],
StatusCode1.BUYER_FAIL: ["16", "29", "49", "50", "78", "114", "110", "114"],
StatusCode1.BUYER_QUALITY_FAIL: ["26", "52", "68", "81", "84"],
@@ -167,10 +167,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
StatusCode1.PS_FAIL: ["7", "29", "36", "47", "56", "58", "64"],
StatusCode1.PS_OVERQUOTA: ["29", "46", "33", "31"],
}
-ext_status_code_map: Dict[str, StatusCode1] = dict()
+ext_status_code_map: dict[str, StatusCode1] = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -179,9 +179,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: from callback url params: scstatus
:params ext_status_code_2: from callback url params: scsecuritystatus
diff --git a/generalresearch/wall_status_codes/spectrum.py b/generalresearch/wall_status_codes/spectrum.py
index 9e9e0a2..7dee1e8 100644
--- a/generalresearch/wall_status_codes/spectrum.py
+++ b/generalresearch/wall_status_codes/spectrum.py
@@ -3,11 +3,11 @@ https://purespectrum.atlassian.net/wiki/spaces/PA/pages/33613201/Minimizing+Clic
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
-status_codes_spectrum: Dict[str, str] = {
+status_codes_spectrum: dict[str, str] = {
"11": "PS Drop",
"12": "PS Quota Full Core",
"13": "PS Termination Core",
@@ -80,7 +80,7 @@ status_codes_spectrum: Dict[str, str] = {
"88": "PS_Supplier_Allocation_Throttle",
}
status_map = defaultdict(lambda: Status.FAIL, **{"21": Status.COMPLETE})
-status_codes_ext_map: Dict[StatusCode1, List[str]] = {
+status_codes_ext_map: dict[StatusCode1, list[str]] = {
StatusCode1.COMPLETE: ["21"],
StatusCode1.BUYER_FAIL: ["16", "17", "18", "19", "30", "59", "84"],
StatusCode1.BUYER_QUALITY_FAIL: ["20", "31"],
@@ -140,10 +140,10 @@ status_codes_ext_map: Dict[StatusCode1, List[str]] = {
"72",
],
}
-ext_status_code_map = dict()
+ext_status_code_map = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[str]
+ v: list[str]
for vv in v:
vv: str
@@ -152,9 +152,9 @@ for k, v in status_codes_ext_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[Any]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, Any | None]:
"""
:params ext_status_code_1: from url params: ps_rstatus
https://purespectrum.atlassian.net/wiki/spaces/PA/pages/33613201/Minimizing+Clickwaste+with+ps+rstatus
diff --git a/generalresearch/wall_status_codes/wxet.py b/generalresearch/wall_status_codes/wxet.py
index e7cf67d..ad25ff3 100644
--- a/generalresearch/wall_status_codes/wxet.py
+++ b/generalresearch/wall_status_codes/wxet.py
@@ -1,5 +1,4 @@
from collections import defaultdict
-from typing import Dict, List, Optional, Tuple
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.wxet.models.definitions import (
@@ -8,10 +7,10 @@ from generalresearch.wxet.models.definitions import (
WXETStatusCode2,
)
-status_map: Dict[WXETStatus, Status] = defaultdict(
+status_map: dict[WXETStatus, Status] = defaultdict(
lambda: Status.FAIL, **{WXETStatus.COMPLETE: Status.COMPLETE}
)
-status_codes_ext_map: Dict[StatusCode1, List[WXETStatusCode1]] = {
+status_codes_ext_map: dict[StatusCode1, list[WXETStatusCode1]] = {
StatusCode1.COMPLETE: [WXETStatusCode1.COMPLETE],
StatusCode1.BUYER_FAIL: [
WXETStatusCode1.BUYER_DUPLICATE,
@@ -30,16 +29,16 @@ status_codes_ext_map: Dict[StatusCode1, List[WXETStatusCode1]] = {
StatusCode1.UNKNOWN: [],
StatusCode1.MARKETPLACE_FAIL: [WXETStatusCode1.BUYER_POSTBACK_NOT_RECEIVED],
}
-ext_status_code_map = dict()
+ext_status_code_map = {}
for k, v in status_codes_ext_map.items():
k: StatusCode1
- v: List[WXETStatusCode1]
+ v: list[WXETStatusCode1]
for vv in v:
vv: WXETStatusCode1
ext_status_code_map[vv] = k
-status_code2_map: Dict[StatusCode1, List[WXETStatusCode2]] = {
+status_code2_map: dict[StatusCode1, list[WXETStatusCode2]] = {
StatusCode1.PS_QUALITY: [],
StatusCode1.PS_DUPLICATE: [
WXETStatusCode2.WORKER_INELIGIBLE,
@@ -59,7 +58,7 @@ status_code2_map: Dict[StatusCode1, List[WXETStatusCode2]] = {
WXETStatusCode2.TASK_VERSION_MISMATCH,
],
}
-ext_status_code2_map = dict()
+ext_status_code2_map = {}
for k, v in status_code2_map.items():
for vv in v:
ext_status_code2_map[vv] = k
@@ -67,9 +66,9 @@ for k, v in status_code2_map.items():
def annotate_status_code(
ext_status_code_1: str,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
-) -> Tuple[Status, StatusCode1, Optional[WXETStatusCode2]]:
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+) -> tuple[Status, StatusCode1, WXETStatusCode2 | None]:
"""
:params ext_status_code_1: WXETStatus
:params ext_status_code_2: WXETStatusCode1
diff --git a/generalresearch/wxet/models/definitions.py b/generalresearch/wxet/models/definitions.py
index 8d0d6b8..ea0b51f 100644
--- a/generalresearch/wxet/models/definitions.py
+++ b/generalresearch/wxet/models/definitions.py
@@ -1,11 +1,10 @@
-from enum import Enum
-from typing import Optional, Tuple
+from enum import IntEnum, StrEnum
from generalresearch.currency import USDMill
from generalresearch.utils.enum import ReprEnumMeta
-class IncExcFilterType(str, Enum, metaclass=ReprEnumMeta):
+class IncExcFilterType(StrEnum, metaclass=ReprEnumMeta):
INCLUDE = "include"
EXCLUDE = "exclude"
@@ -13,7 +12,7 @@ class IncExcFilterType(str, Enum, metaclass=ReprEnumMeta):
# Note: This is exactly the same as the generalresearch:models/thl/definitions.py:Status.
# Keeping this because the comments (and as a result, the documentation)
# is slightly different, and specific to wxet.
-class WXETStatus(str, Enum, metaclass=ReprEnumMeta):
+class WXETStatus(StrEnum, metaclass=ReprEnumMeta):
"""
The outcome of a task attempt. If the attempt is still in progress, the status will be NULL.
"""
@@ -36,7 +35,7 @@ class WXETStatus(str, Enum, metaclass=ReprEnumMeta):
# Basically same note as for WxetStatus for WallAdjustedStatus
-class WXETAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class WXETAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
# Task was reconciled to complete
ADJUSTED_TO_COMPLETE = "ac"
@@ -52,7 +51,7 @@ class WXETAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
POSTBACK_COMPLETE = "pc"
-class WXETStatusCode1(int, Enum, metaclass=ReprEnumMeta):
+class WXETStatusCode1(IntEnum, metaclass=ReprEnumMeta):
"""
__High level status code for outcome of the attempt.__
This should only be NULL if the WXETStatus is ABANDON or TIMEOUT
@@ -103,10 +102,10 @@ class WXETStatusCode1(int, Enum, metaclass=ReprEnumMeta):
"""This property helper indicates if the WXET Attempt made it into
the WXET Account's (eg: the "buyer"'s) Task.
"""
- return False if self.value > 10 else True
+ return not self.value > 10
-class WXETStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class WXETStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
__Status Detail__
These are generally only set if the StatusCode1 is WXET_FAIL,
@@ -166,8 +165,8 @@ class WXETStatusCode2(int, Enum, metaclass=ReprEnumMeta):
def check_wxet_status_consistent(
status: WXETStatus,
- status_code_1: Optional[WXETStatusCode1] = None,
- status_code_2: Optional[WXETStatusCode2] = None,
+ status_code_1: WXETStatusCode1 | None = None,
+ status_code_2: WXETStatusCode2 | None = None,
) -> bool:
"""
Raises an AssertionError if inconsistent
@@ -203,13 +202,13 @@ def check_wxet_status_consistent(
def check_wxet_adjusted_status_attempt_consistent(
status: WXETStatus,
- status_code_1: Optional[WXETStatusCode1] = None,
- cpi: Optional[USDMill] = None,
- adjusted_status: Optional[WXETAdjustedStatus] = None,
- adjusted_cpi: Optional[USDMill] = None,
- new_adjusted_status: Optional[WXETAdjustedStatus] = None,
- new_adjusted_cpi: Optional[USDMill] = None,
-) -> Tuple[bool, str]:
+ status_code_1: WXETStatusCode1 | None = None,
+ cpi: USDMill | None = None,
+ adjusted_status: WXETAdjustedStatus | None = None,
+ adjusted_cpi: USDMill | None = None,
+ new_adjusted_status: WXETAdjustedStatus | None = None,
+ new_adjusted_cpi: USDMill | None = None,
+) -> tuple[bool, str]:
"""
Raises an AssertionError if inconsistent.
- status, status_code_1, adjusted_status, adjusted_cpi, cpi are the attempt's CURRENT values
@@ -233,12 +232,12 @@ def check_wxet_adjusted_status_attempt_consistent(
def _check_wxet_adjusted_status_attempt_consistent(
status: WXETStatus,
- status_code_1: Optional[WXETStatusCode1] = None,
- cpi: Optional[USDMill] = None,
- adjusted_status: Optional[WXETAdjustedStatus] = None,
- adjusted_cpi: Optional[USDMill] = None,
- new_adjusted_status: Optional[WXETAdjustedStatus] = None,
- new_adjusted_cpi: Optional[USDMill] = None,
+ status_code_1: WXETStatusCode1 | None = None,
+ cpi: USDMill | None = None,
+ adjusted_status: WXETAdjustedStatus | None = None,
+ adjusted_cpi: USDMill | None = None,
+ new_adjusted_status: WXETAdjustedStatus | None = None,
+ new_adjusted_cpi: USDMill | None = None,
) -> None:
"""
Raises an AssertionError if inconsistent.
@@ -297,8 +296,8 @@ def _check_wxet_adjusted_status_attempt_consistent(
def _check_wxet_adjusted_status_consistent(
- adjusted_status: Optional[WXETAdjustedStatus] = None,
- adjusted_cpi: Optional[USDMill] = None,
+ adjusted_status: WXETAdjustedStatus | None = None,
+ adjusted_cpi: USDMill | None = None,
) -> None:
"""
Raises an AssertionError if inconsistent.
diff --git a/generalresearch/wxet/models/finish_type.py b/generalresearch/wxet/models/finish_type.py
index af60fe6..17923d6 100644
--- a/generalresearch/wxet/models/finish_type.py
+++ b/generalresearch/wxet/models/finish_type.py
@@ -1,11 +1,10 @@
-from enum import Enum
-from typing import Optional, Set
+from enum import StrEnum
from generalresearch.utils.enum import ReprEnumMeta
from generalresearch.wxet.models.definitions import WXETStatus, WXETStatusCode1
-class FinishType(str, Enum, metaclass=ReprEnumMeta):
+class FinishType(StrEnum, metaclass=ReprEnumMeta):
"""A Task can be classified as "finished" based on different outcomes.
<br/>
This controls how the `Task.required_finish_count` value
@@ -33,7 +32,7 @@ class FinishType(str, Enum, metaclass=ReprEnumMeta):
FAIL = "fail"
@property
- def finish_statuses(self) -> Set[Optional[WXETStatus]]:
+ def finish_statuses(self) -> set[WXETStatus | None]:
"""For this particular FinishType, what are the different WXETStatus
values that are consider
"""
@@ -64,9 +63,9 @@ class FinishType(str, Enum, metaclass=ReprEnumMeta):
def is_a_finish(
- status: Optional[WXETStatus],
- status_code_1: Optional[WXETStatusCode1],
- finish_type: Optional[FinishType],
+ status: WXETStatus | None,
+ status_code_1: WXETStatusCode1 | None,
+ finish_type: FinishType | None,
) -> bool:
"""Determines if a wall event should be considered a finish or not.
diff --git a/pyproject.toml b/pyproject.toml
index 0cc4227..5719a43 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -9,15 +9,16 @@ description = "Python Utilities for General Research"
readme = "README.md"
requires-python = ">=3.10"
dependencies = [
-# "fastapi",
"Faker",
"PyMySQL",
"psycopg",
"cachetools",
"decorator",
+ "influxdb",
"limits",
"more-itertools",
"numpy",
+ "pandas",
"pandera",
"protobuf",
"pyarrow",
@@ -34,6 +35,7 @@ dependencies = [
"scipy",
"sentry-sdk",
"slackclient",
+ "grip-client",
"tldextract",
"ua-parser",
"user-agents",
@@ -41,7 +43,7 @@ dependencies = [
]
[project.optional-dependencies]
django = ["Django>=5.2", "psycopg>=3.1"]
-
+dask = ["dask>=2026.7.1", "distributed>=2026.7.1"]
[tool.setuptools.packages.find]
where = ["."]
@@ -54,3 +56,13 @@ include = ["generalresearch", "generalresearch.*", "test_utils", "test_utils.*"]
[tool.pytest.ini_options]
testpaths = ["tests"]
addopts = "-v --tb=short"
+
+[tool.ruff]
+target-version = "py314"
+exclude = [
+ "generalresearch/thl_django",
+]
+
+[tool.pylint.messages_control]
+disable = ["all"]
+enable = ["cyclic-import"] \ No newline at end of file
diff --git a/requirements.txt b/requirements.txt
deleted file mode 100644
index 6c04995..0000000
--- a/requirements.txt
+++ /dev/null
@@ -1,112 +0,0 @@
-aiohappyeyeballs==2.6.1
-aiohttp==3.12.15
-aiosignal==1.4.0
-annotated-types==0.7.0
-anyio==4.10.0
-attrs==25.3.0
-boto3==1.40.19
-botocore==1.40.19
-CacheControl==0.14.3
-cachetools==6.1.0
-certifi==2025.8.3
-cffi==1.17.1
-charset-normalizer==3.4.3
-click==8.2.2
-cloudpickle==3.1.1
-coverage==7.10.5
-cryptography==45.0.6
-dask==2025.7.0
-decorator==5.2.1
-Deprecated==1.2.18
-distributed==2025.7.0
-dnspython==2.7.0
-Django>=5.2
-ecdsa==0.19.1
-email-validator==2.3.0
-Faker==37.6.0
-filelock==3.25.1
-frozenlist==1.7.0
-fsspec==2025.7.0
-geoip2==4.7.0
-idna==3.10
-importlib_metadata==8.7.0
-iniconfig==2.1.0
-Jinja2==3.1.6
-jmespath==1.0.1
-jsonpickle==5.0.0rc1
-limits==5.5.0
-locket==1.0.0
-MarkupSafe==3.0.2
-maxminddb==2.8.2
-more-itertools==10.7.0
-msgpack==1.1.1
-multidict==6.6.4
-mypy_extensions==1.1.0
-numpy==2.3.2
-opentelemetry-api==1.36.0
-opentelemetry-sdk==1.36.0
-opentelemetry-semantic-conventions==0.57b0
-outcome==1.3.0.post0
-packaging==25.0
-pandas==2.3.2
-pandera==0.26.1
-partd==1.4.2
-phonenumbers==9.0.12
-pluggy==1.6.0
-propcache==0.3.2
-protobuf==6.32.0
-psutil==7.0.0
-psycopg==3.2.9
-psycopg-binary==3.2.9
-pyarrow==21.0.0
-pyasn1==0.6.1
-pycountry==24.6.1
-pycparser==2.22
-pydantic==2.11.7
-pydantic-extra-types==2.10.5
-pydantic-settings==2.10.1
-pydantic_core==2.33.2
-Pygments==2.19.2
-pylibmc==1.6.3
-pymemcache==4.0.0
-PyMySQL==1.1.1
-pytest==8.4.1
-pytest-anyio==0.0.0
-pytest-cov==6.2.1
-python-dateutil==2.9.0.post0
-python-dotenv==1.1.1
-python-jose==3.5.0
-pytz==2025.2
-PyYAML==6.0.2
-redis==6.4.0
-requests==2.32.5
-requests-file==3.0.1
-rsa==4.9.1
-s3transfer==0.13.1
-scipy==1.16.1
-sentry-sdk==3.0.0a5
-setuptools==80.9.0
-six==1.17.0
-slackclient==2.9.4
-sniffio==1.3.1
-sortedcontainers==2.4.0
-tblib==3.1.0
-tldextract==5.3.1
-toolz==1.0.0
-tornado==6.5.2
-trio==0.30.0
-typeguard==4.4.4
-typing-inspect==0.9.0
-typing-inspection==0.4.1
-typing_extensions==4.15.0
-tzdata==2025.2
-ua-parse==1.0.1
-ua-parser==1.0.1
-ua-parser-builtins==0.19.0.dev79
-urllib3==2.5.0
-user-agents==2.2.0
-wheel==0.46.1
-wrapt==1.17.3
-yarl==1.20.1
-zict==3.0.0
-zipp==3.23.0
diff --git a/test_utils/conftest.py b/test_utils/conftest.py
index 378b9cc..33a7e77 100644
--- a/test_utils/conftest.py
+++ b/test_utils/conftest.py
@@ -5,24 +5,27 @@ import shutil
import stat
import subprocess
import sys
-import tempfile
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable, Generator
+from datetime import UTC, datetime, timedelta
from os.path import join as pjoin
from pathlib import Path
-from typing import Callable, Generator
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from _pytest.config import Config
from dotenv import load_dotenv
from pydantic import MariaDBDsn, PostgresDsn, TypeAdapter
+from pytest import FixtureRequest, TempPathFactory
-from generalresearch.config import GRLBaseSettings
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import InternalHostname, PostgresDict
-from generalresearch.pg_helper import PostgresConfig
from generalresearch.sql_helper import SqlHelper
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.pg_helper import PostgresConfig
+
@pytest.fixture(scope="session")
def env_file_path(pytestconfig: Config) -> Path:
@@ -93,7 +96,7 @@ def postgres_instance(settings: GRLBaseSettings) -> Generator[PostgresDsn]:
from psycopg import connect
from psycopg.sql import SQL, Identifier
- now = datetime.now(timezone.utc)
+ now = datetime.now(UTC)
ts: str = now.strftime("%Y-%m-%d")
db_name = f"unittest-{ts}-{uuid4().hex[:6]}"
@@ -152,40 +155,51 @@ def postgres_instance_host(
yield value
-# @pytest.fixture(scope="session")
-# def git_key_path(settings: GRLBaseSettings) -> Path:
-# return Path('/tmp/')
-
-
@pytest.fixture(scope="session")
def git_key_path(
+ tmp_path_factory: TempPathFactory,
settings: GRLBaseSettings,
) -> Generator[Path]:
+ # We are using the tmp_path_factory because unlike the tmp_path (which
+ # is function scoped), this is session scoped.
- assert settings.git_creds
- with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix="_id_rsa") as f:
- f.write(settings.git_creds)
- key_path = f.name
-
- os.chmod(key_path, stat.S_IRUSR | stat.S_IWUSR)
+ assert settings.git_creds, "Must define key to download alternative models"
+ fn = tmp_path_factory.mktemp("keys") / "git_creds"
+ key_content = settings.git_creds.replace("\\n", "\n")
+ fn.write_text(key_content, encoding="utf-8")
+ os.chmod(fn, stat.S_IRUSR | stat.S_IWUSR)
- yield Path(key_path)
+ yield Path(fn)
- os.unlink(key_path)
+ os.unlink(fn)
@pytest.fixture(scope="session")
-def gr_repo(git_key_path: Path) -> Callable[..., Path]:
- repo_url = "ssh://code.g-r-l.com/general-research/gr-carer.git"
- repo_path = Path("/tmp/gr-carer")
+def gr_repo(
+ git_key_path: Path,
+ tmp_path_factory: TempPathFactory,
+) -> Callable[..., Path | None]:
+ repo_url = "ssh://code.g-r-l.com:6611/general-research/gr-carer.git"
+
+ _ran = {}
+
+ fn = tmp_path_factory.mktemp("repos")
+ repo_path = fn / "gr-carer"
def _inner() -> Path:
+
+ if _ran.get(repo_url, False):
+ print(f"Already ran django_db_factory.{repo_url}")
+ return repo_path
+
+ _ran[repo_url] = True
+
ssh_cmd = (
- f"ssh -i {git_key_path} "
+ f'ssh -i "{git_key_path}" '
"-o IdentitiesOnly=yes "
- "-o StrictHostKeyChecking=no " # or accept-new, see note below
+ "-o StrictHostKeyChecking=no "
)
- env = {"GIT_SSH_COMMAND": ssh_cmd}
+ env = {**os.environ, "GIT_SSH_COMMAND": ssh_cmd}
if repo_path.exists():
subprocess.run(["git", "-C", str(repo_path), "pull"], check=True, env=env)
@@ -202,51 +216,140 @@ def gr_repo(git_key_path: Path) -> Callable[..., Path]:
@pytest.fixture(scope="session")
+def django_settings_file(
+ postgres_instance_dict: PostgresDict,
+) -> Callable[..., Path]:
+
+ def _inner(
+ settings_dir: Path, extra_installed_apps: list[str] | None = None
+ ) -> Path:
+ installed_apps = [
+ "django.contrib.postgres",
+ "django.contrib.contenttypes",
+ ] + (extra_installed_apps or [])
+ """
+ This returns the directory path of where the settings file is in,
+ not the path of the settings file itself
+ """
+
+ settings_content = f"""DATABASES = {{
+ "default": {{
+ "ENGINE": "django.db.backends.postgresql",
+ "NAME": {postgres_instance_dict["name"]!r},
+ "USER": {postgres_instance_dict["username"]!r},
+ "PASSWORD": {postgres_instance_dict["password"]!r},
+ "HOST": {postgres_instance_dict["host"]!r},
+ "PORT": {postgres_instance_dict["port"]!r},
+ }}
+}}
+INSTALLED_APPS = {installed_apps!r}
+DEFAULT_AUTO_FIELD = "django.db.models.BigAutoField"
+LANGUAGE_CODE = "en-us"
+TIME_ZONE = "UTC"
+USE_I18N = True
+USE_L10N = True
+USE_TZ = True
+"""
+ settings_file_path = settings_dir / "test_settings.py"
+ settings_file_path.unlink(missing_ok=True)
+ settings_file_path.write_text(settings_content)
+
+ return settings_dir
+
+ return _inner
+
+
+@pytest.fixture(scope="session")
def django_db_factory(
+ request: FixtureRequest,
postgres_instance: PostgresDsn,
- postgres_instance_dict: PostgresDict,
gr_repo: Callable[..., Path],
-) -> Callable[..., PostgresDsn]:
-
- import django
- from django.conf import settings as django_settings
- from django.core.management import call_command
-
- def _inner(django_project: str = "generalresearch.thl_django"):
-
- if "gr" in django_project:
- # We need model files that are NOT in this repo.
- gr_path = gr_repo()
- sys.path.insert(0, str(gr_path))
-
- print(sys.path)
-
- # 1. Bootstrapping Django settings
- if not django_settings.configured:
- django_settings.configure(
- DATABASES={
- "default": {
- "ENGINE": "django.db.backends.postgresql",
- "NAME": postgres_instance_dict["name"],
- "USER": postgres_instance_dict["username"],
- "PASSWORD": postgres_instance_dict["password"],
- "HOST": postgres_instance_dict["host"],
- "PORT": postgres_instance_dict["port"],
- }
- },
- INSTALLED_APPS=[
- "django.contrib.postgres",
- "django.contrib.contenttypes",
- django_project,
+ django_settings_file: Callable[..., Path],
+ postgres_instance_dict: PostgresDict,
+ tmp_path_factory: TempPathFactory,
+) -> Callable[..., PostgresDsn | None]:
+
+ _ran = {}
+
+ def _inner(
+ django_project: str = "generalresearch.thl_django",
+ ) -> PostgresDsn | None:
+
+ if _ran.get(django_project, False):
+ print(f"Already ran django_db_factory:{django_project}")
+ return postgres_instance
+ _ran[django_project] = True
+
+ # This is the generalresearch project root path, it's
+ # 1 directory up from test_utils/, or tests/
+ base_dir = Path(request.config.rootpath).parent
+
+ if django_project == "generalresearch.thl_django":
+ _cwd = base_dir
+ _manage_path = "generalresearch.thl_django.app.manage"
+ _settings_dir = base_dir / "generalresearch/thl_django/app"
+ _settings_module = "generalresearch.thl_django.app.test_settings"
+ django_settings_file(
+ settings_dir=_settings_dir,
+ extra_installed_apps=[
+ "generalresearch.thl_django",
],
)
- django.setup()
- # for model in apps.get_models():
- # print(f"Discovered model: {model._meta.label}")
+ elif django_project == "gr.common":
+ _cwd = gr_repo()
+ _manage_path = "gr.app.manage"
+ _settings_dir = gr_repo() / "gr/app"
+ _settings_module = "gr.app.test_settings"
+ django_settings_file(
+ settings_dir=_settings_dir, extra_installed_apps=["gr.common"]
+ )
+
+ else:
+ raise ValueError("Not implemented yet.")
+
+ assert _settings_dir
+
+ env = {"DJANGO_SETTINGS_MODULE": str(_settings_module)}
+ res1 = subprocess.run(
+ [
+ sys.executable,
+ "-m",
+ _manage_path,
+ "makemigrations",
+ f"--settings={_settings_module}",
+ ],
+ cwd=str(_cwd),
+ env=env,
+ capture_output=True,
+ text=True,
+ check=True,
+ )
+
+ if res1.returncode != 0:
+ print("STDOUT:", res1.stdout)
+ print("STDERR:", res1.stderr)
+ res1.check_returncode()
+
+ res2 = subprocess.run(
+ [
+ sys.executable,
+ "-m",
+ _manage_path,
+ "migrate",
+ f"--settings={_settings_module}",
+ ],
+ env=env,
+ cwd=str(_cwd),
+ capture_output=True,
+ text=True,
+ check=True,
+ )
- # 2. Run migrations directly during fixture activation
- call_command("migrate")
+ if res2.returncode != 0:
+ print("STDOUT:", res2.stdout)
+ print("STDERR:", res2.stderr)
+ res2.check_returncode()
# 3. Return the Dsn so the factory gives a way to connect
return postgres_instance
@@ -276,37 +379,37 @@ def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper:
@pytest.fixture
def start() -> datetime:
- return datetime(year=1900, month=1, day=1, tzinfo=timezone.utc)
+ return datetime(year=1900, month=1, day=1, tzinfo=UTC)
@pytest.fixture
def utc_now() -> datetime:
- return datetime.now(tz=timezone.utc)
+ return datetime.now(tz=UTC)
@pytest.fixture
def utc_hour_ago() -> datetime:
- return datetime.now(tz=timezone.utc) - timedelta(hours=1)
+ return datetime.now(tz=UTC) - timedelta(hours=1)
@pytest.fixture
def utc_day_ago() -> datetime:
- return datetime.now(tz=timezone.utc) - timedelta(hours=24)
+ return datetime.now(tz=UTC) - timedelta(hours=24)
@pytest.fixture
def utc_90days_ago() -> datetime:
- return datetime.now(tz=timezone.utc) - timedelta(days=90)
+ return datetime.now(tz=UTC) - timedelta(days=90)
@pytest.fixture
def utc_60days_ago() -> datetime:
- return datetime.now(tz=timezone.utc) - timedelta(days=60)
+ return datetime.now(tz=UTC) - timedelta(days=60)
@pytest.fixture
def utc_30days_ago() -> datetime:
- return datetime.now(tz=timezone.utc) - timedelta(days=30)
+ return datetime.now(tz=UTC) - timedelta(days=30)
# === Clean up ===
@@ -317,12 +420,12 @@ def delete_df_collection(
thl_web_rw: PostgresConfig, create_main_accounts: Callable[..., None]
) -> Callable[..., None]:
- from generalresearch.incite.collections import (
+ from generalresearch.incite.collections.base import (
DFCollection,
DFCollectionType,
)
- def _inner(coll: "DFCollection"):
+ def _inner(coll: DFCollection):
match coll.data_type:
case DFCollectionType.LEDGER:
for table in [
@@ -337,16 +440,15 @@ def delete_df_collection(
create_main_accounts()
case DFCollectionType.WALL | DFCollectionType.SESSION:
- with thl_web_rw.make_connection() as conn:
- with conn.cursor() as c:
- c.execute("SET CONSTRAINTS ALL DEFERRED")
- for table in [
- "thl_wall",
- "thl_session",
- ]:
- c.execute(
- query=f"DELETE FROM {table};",
- )
+ with thl_web_rw.make_connection() as conn, conn.cursor() as c:
+ c.execute("SET CONSTRAINTS ALL DEFERRED")
+ for table in [
+ "thl_wall",
+ "thl_session",
+ ]:
+ c.execute(
+ query=f"DELETE FROM {table};",
+ )
case DFCollectionType.USER:
for table in ["thl_usermetadata", "thl_user"]:
diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py
index e8175a5..249b068 100644
--- a/test_utils/grliq/conftest.py
+++ b/test_utils/grliq/conftest.py
@@ -1,27 +1,34 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
-from typing import Callable
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING, Any
from uuid import uuid4
import pytest
from pydantic import PostgresDsn
-from generalresearch.config import GRLBaseSettings
-from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA
from generalresearch.grliq.managers.forensic_data import (
GrlIqDataManager,
)
-from generalresearch.grliq.managers.forensic_events import (
- GrlIqEventManager,
-)
from generalresearch.grliq.managers.forensic_results import (
GrlIqCategoryResultsReader,
)
from generalresearch.grliq.models.forensic_data import GrlIqData
+from generalresearch.grliq.models.forensic_result import (
+ GrlIqCheckerResults,
+ GrlIqForensicCategoryResult,
+)
from generalresearch.pg_helper import PostgresConfig
-# === Miscellaneous ===
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.grliq.managers.forensic_events import (
+ GrlIqEventManager,
+ )
+
+
+# --- Assets ---
@pytest.fixture(scope="function")
@@ -42,56 +49,28 @@ def grliq_db(postgres_instance: PostgresDsn) -> PostgresConfig:
)
-# === Managers ===
+# --- GRLIQ Data ---
@pytest.fixture(scope="session")
-def grliq_dm(grliq_db: PostgresConfig) -> GrlIqDataManager:
+def grliq_data_manager(grliq_db: PostgresConfig) -> GrlIqDataManager:
assert grliq_db.dsn.path
assert "/unittest-" in grliq_db.dsn.path
return GrlIqDataManager(postgres_config=grliq_db)
@pytest.fixture(scope="session")
-def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager:
- assert grliq_db.dsn.path
- assert "/unittest-" in grliq_db.dsn.path
-
- from generalresearch.grliq.managers.forensic_events import (
- GrlIqEventManager,
- )
-
- return GrlIqEventManager(postgres_config=grliq_db)
-
-
-@pytest.fixture(scope="session")
-def grliq_crr(grliq_db: PostgresConfig) -> GrlIqCategoryResultsReader:
- assert grliq_db.dsn.path
- assert "/unittest-" in grliq_db.dsn.path
-
- return GrlIqCategoryResultsReader(postgres_config=grliq_db)
-
-
-# === Models ===
-
-
-@pytest.fixture(scope="function")
-def grliq_data() -> GrlIqData:
- from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA
-
- g: GrlIqData = DUMMY_GRLIQ_DATA[1]["data"]
-
- g.id = None
- g.uuid = uuid4().hex
- g.created_at = datetime.now(tz=timezone.utc)
- g.timestamp = g.created_at - timedelta(seconds=10)
- return g
+def grliq_dm(grliq_data_manager: GrlIqDataManager) -> GrlIqDataManager:
+ return grliq_data_manager
@pytest.fixture
-def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]:
+def grliq_data_factory(
+ grliq_data_manager: GrlIqDataManager, grliq_data_list: list[dict[str, Any]]
+) -> Callable[..., GrlIqData]:
def _inner(
+ save: bool = True,
is_attempt_allowed: bool = True,
product_id: str | None = None,
product_user_id: str | None = None,
@@ -109,30 +88,127 @@ def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]:
:param mid: the thl_session:uuid / mid for the attempt.
:return:
"""
- import copy
-
- res: GrlIqData = copy.deepcopy(DUMMY_GRLIQ_DATA[int(is_attempt_allowed)])
-
- product_id = product_id or uuid4().hex
- product_user_id = product_user_id or uuid4().hex
- uuid = uuid or uuid4().hex
- mid = mid or uuid4().hex
- created_at = created_at or datetime.now(tz=timezone.utc)
-
- res["data"].product_id = product_id
- res["data"].product_user_id = product_user_id
- res["data"].uuid = uuid
- res["data"].mid = mid
- res["data"].created_at = created_at
- res["result_data"].uuid = uuid
- res["category_result"].uuid = uuid
-
- return grliq_dm.create(
- iq_data=res["data"],
- result_data=res["result_data"],
- category_result=res["category_result"],
- fraud_score=res["category_result"].fraud_score,
- is_attempt_allowed=res["category_result"].is_attempt_allowed(),
- )
+
+ if save:
+ res: dict = grliq_data_list[int(is_attempt_allowed)]
+
+ product_id = product_id or uuid4().hex
+ product_user_id = product_user_id or uuid4().hex
+ uuid = uuid or uuid4().hex
+ mid = mid or uuid4().hex
+ created_at = created_at or datetime.now(tz=UTC)
+
+ res["data"].product_id = product_id
+ res["data"].product_user_id = product_user_id
+ res["data"].uuid = uuid
+ res["data"].mid = mid
+ res["data"].created_at = created_at
+ res["result_data"].uuid = uuid
+ res["category_result"].uuid = uuid
+
+ return grliq_data_manager.create(
+ iq_data=res["data"],
+ result_data=res["result_data"],
+ category_result=res["category_result"],
+ fraud_score=res["category_result"].fraud_score,
+ is_attempt_allowed=res["category_result"].is_attempt_allowed(),
+ )
+ else:
+ raise ValueError("Unsaved GRLIQ Data not supported yet")
return _inner
+
+
+@pytest.fixture(scope="function")
+def grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData:
+
+ g: GrlIqData = grliq_data_list[1]["data"]
+
+ g.id = None
+ g.uuid = uuid4().hex
+ g.created_at = datetime.now(tz=UTC)
+ g.timestamp = g.created_at - timedelta(seconds=10)
+ return g
+
+
+@pytest.fixture(scope="function")
+def unsaved_grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData:
+ raise ValueError("Not supported")
+
+
+# --- GRLIQ Event ---
+
+
+@pytest.fixture(scope="session")
+def grliq_event_manager(grliq_db: PostgresConfig) -> GrlIqEventManager:
+ assert grliq_db.dsn.path
+ assert "/unittest-" in grliq_db.dsn.path
+
+ from generalresearch.grliq.managers.forensic_events import (
+ GrlIqEventManager,
+ )
+
+ return GrlIqEventManager(postgres_config=grliq_db)
+
+
+@pytest.fixture(scope="session")
+def grliq_em(grliq_event_manager: GrlIqEventManager) -> GrlIqEventManager:
+ return grliq_event_manager
+
+
+# --- GRLIQ Category Results Reader ---
+
+
+@pytest.fixture(scope="session")
+def grliq_category_results_reader(
+ grliq_db: PostgresConfig,
+) -> GrlIqCategoryResultsReader:
+ assert grliq_db.dsn.path
+ assert "/unittest-" in grliq_db.dsn.path
+
+ return GrlIqCategoryResultsReader(postgres_config=grliq_db)
+
+
+@pytest.fixture(scope="session")
+def grliq_crr(
+ grliq_category_results_reader: GrlIqCategoryResultsReader,
+) -> GrlIqCategoryResultsReader:
+ return grliq_category_results_reader
+
+
+# === Models ===
+
+
+# === Miscellaneous ===
+
+
+@pytest.fixture(scope="session")
+def grliq_data_list() -> list[dict[str, Any]]:
+ return [
+ {
+ "data": GrlIqData.model_validate_json(
+ """{"mid": "3722ed29314940fabd37b42d808dcf5a", "uuid": "b11441da5a854dfbb8401d4c32e56db5", "phase": "offerwall-enter", "events": null, "vendor": "Google Inc.", "app_name": "Netscape", "calendar": "gregory", "language": "en-US", "platform": "Linux x86_64", "timezone": "America/Mexico_City", "client_ip": "131.196.250.250", "timestamp": "2025-02-27T16:05:34-06:00", "webrtc_ip": "131.196.250.250", "created_at": "2025-02-27T22:05:35.370589Z", "language_2": "en-US", "language_3": null, "platform_2": "Linux x86_64", "platform_3": null, "prefetched": true, "product_id": "d0606a0b5d034a8d81b1e3579d1f76fd", "webgl_flag": true, "webgl_hash": "da27e1b9b660057a3f5e185d3f5deabe", "canvas_hash": "14ed764326ec454d976c322261d99f16", "color_gamut": "3", "country_iso": "mx", "inner_width": 612, "outer_width": 1813, "product_sub": "20030107", "audio_codecs": "1,1,1,1,1,3,1,3,1,3,3,1,1,3,3,3,3,1,3,3,3,2,1,1", "cookie_check": "", "graphics_api": "WebKit WebGL", "inner_height": 1174, "mouse_events": null, "ontouchstart": false, "outer_height": 1261, "plugins_hash": "4c05fa2f766a444d4f253ead792c8b0e|2", "screen_width": 2560, "video_codecs": "1,3,3,3,3,3,3,3,3,3,1,1,1,1,1,1,3,1,1,1,3,3,1", "webgl_hash_2": "fc73fd5db75e2c36222fe34251be3971", "webrtc_error": false, "window_opera": false, "battery_level": 0.9, "canvas_hash_2": "bd11ebbf5c26fd20e0217820b4159752", "dynamic_range": false, "error_message": "Cannot read", "forced_colors": false, "math_result_1": "1.9275814160560204e-50", "math_result_2": "1.6182817135715877", "screen_height": 1440, "webgl_check_1": true, "webgl_context": "webgl2", "window_chrome": true, "connection_rtt": 150, "history_length": 16, "user_agent_str": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "web_sql_exists": false, "calender_locale": "en-US", "connection_type": "", "inverted_colors": true, "navigator_brave": false, "product_user_id": "d1d55df1-959e-4740-b77c-fa1f4fc457ae", "request_headers": {"host": "test", "accept": "*/*", "connection": "keep-alive", "user-agent": "python-httpx/0.27.0", "content-length": "3646", "accept-encoding": "gzip, deflate", "x-forwarded-for": "131.196.250.250"}, "timezone_offset": 360, "webrtc_local_ip": "50486637-6b64-4812-b10a-0a75337c31bd.local", "battery_charging": true, "client_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "131.196.250.250", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "mx", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "max_touch_points": 0, "numbering_system": "latn", "path_fingerprint": 3252, "prefers_contrast": "0", "rendering_engine": "WebKit", "timezone_success": "pass", "user_agent_hints": {"model": null, "brands": [{"brand": "Google Chrome", "version": "131"}, {"brand": "Chromium", "version": "131"}, {"brand": "Not_A Brand", "version": "24"}], "mobile": false, "bitness": "64", "platform": "Linux", "brands_full": [{"brand": "Google Chrome", "version": "131.0.6778.204"}, {"brand": "Chromium", "version": "131.0.6778.204"}, {"brand": "Not_A Brand", "version": "24.0.0.0"}], "architecture": "x86", "platform_version": "6.2.0"}, "user_agent_str_2": null, "webgl_extensions": "EXT_clip_control|EXT_color_buffer_float|EXT_color_buffer_half_float|EXT_conservative_depth|EXT_depth_clamp|EXT_disjoint_timer_query_webgl2|EXT_float_blend|EXT_polygon_offset_clamp|EXT_render_snorm|EXT_texture_compression_bptc|EXT_texture_compression_rgtc|EXT_texture_filter_anisotropic|EXT_texture_mirror_clamp_to_edge|EXT_texture_norm16|KHR_parallel_shader_compile|NV_shader_noperspective_interpolation|OES_draw_buffers_indexed|OES_sample_variables|OES_shader_multisample_interpolation|OES_texture_float_linear|OVR_multiview2|WEBGL_blend_func_extended|WEBGL_clip_cull_distance|WEBGL_compressed_texture_astc|WEBGL_compressed_texture_etc|WEBGL_compressed_texture_etc1|WEBGL_compressed_texture_s3tc|WEBGL_compressed_texture_s3tc_srgb|WEBGL_debug_renderer_info|WEBGL_debug_shaders|WEBGL_lose_context|WEBGL_multi_draw|WEBGL_polygon_mode|WEBGL_provoking_vertex|WEBGL_stencil_texturing", "webrtc_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "131.196.250.250", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "mx", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "chrome_extensions": "", "execution_time_ms": 371.0999999642372, "graphics_renderer": "WebGL 2.0 (OpenGL ES 3.0 Chromium)", "keyboard_detected": true, "mime_types_length": 2, "request_fs_exists": true, "audio_context_flag": "pass", "audio_context_hash": "9307303774dec3248c18a939392090da", "canvas_fingerprint": 258, "canvas_pixel_check": false, "device_pixel_ratio": 1.0, "indexedDbData_blob": true, "navigator_keys_len": 79, "no_edge_pdf_plugin": false, "screen_avail_width": 2560, "webdriver_detected": false, "window_orientation": 0, "connection_downlink": 10.0, "navigator_webdriver": false, "non_native_function": false, "screen_avail_height": 1400, "supported_fonts_str": "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0|2147483648|109543424|1880099872|268435471", "text_2d_fingerprint": "bfcce91c9e71d11af7b14dbee4c75f83", "webrtc_is_supported": "pass", "canvas_support_level": "full", "do_not_track_enabled": "1", "hardware_concurrency": 12, "keyboard_layout_size": 48, "prefers_color_scheme": false, "webgl_max_anisotropy": 16, "battery_charging_time": 0.0, "browser_by_properties": "c", "eval_to_string_length": 33, "performance_loop_time": 0.09999996423721313, "session_storage_check": "pass", "unmasked_vendor_webgl": "Google Inc. (Intel)", "hardware_concurrency_2": 12, "hardware_concurrency_3": null, "localStorage_available": true, "memory_jsHeapSizeLimit": 4294705152, "mozilla_web_app_exists": false, "navigator_deviceMemory": 8.0, "navigator_java_enabled": false, "prefers_reduced_motion": false, "storage_estimate_quota": 1178717110272, "webdriver_detected_msg": "", "window_active_x_object": false, "window_external_exists": true, "color_depth_pixel_depth": "24-24", "indexedDbData_available": true, "navigator_cookieEnabled": true, "unmasked_renderer_webgl": "ANGLE (Intel, Mesa Intel(R) Graphics (RPL-P), OpenGL 4.6)", "battery_discharging_time": 0.0, "connection_effectiveType": "4g", "non_native_function_flag": "", "speech_synthesis_voice_1": "Google Bahasa Indonesia", "window_client_information": true, "audio_compressor_reduction": 20.538288116455078, "navigator_mediaDevices_len": 3, "audio_intensity_fingerprint": 124.04347527516074, "speech_synthesis_voice_hash": "8010ee3313813de521e48e63bd5a6f13", "microsoft_credentials_exists": false, "window_installTrigger_exists": false, "speech_synthesis_voices_count": 19, "webgl_shading_language_version": "WebGL GLSL ES 3.00 (OpenGL ES GLSL ES 3.0 Chromium)", "error_message_stack_access_count": 0, "speech_synthesis_avail_voices_count": 19, "error_message_stack_access_count_worker": 0}"""
+ ),
+ "result_data": GrlIqCheckerResults.model_validate_json(
+ """{"uuid": "b11441da5a854dfbb8401d4c32e56db5", "check_codecs": {"score": 0}, "check_timezone": {"score": 0}, "check_timestamp": {"score": 0}, "check_user_type": {"score": 0}, "check_ip_changes": {"score": 0}, "check_ip_country": {"score": 0}, "check_environment": {"score": 0}, "check_ip_timezone": {"score": 0}, "check_isp_changes": {"score": 0}, "check_useragent_js": {"score": 0}, "check_required_fonts": {"score": 0}, "check_user_anonymous": {"score": 0}, "check_webrtc_success": {"score": 0}, "check_seen_timestamps": {"msg": "duplicate timestamp", "score": 100}, "check_country_timezone": {"score": 0}, "check_prohibited_fonts": {"score": 0}, "check_timezone_changes": {"score": 0}, "check_execution_time_ms": {"msg": "duplicate execution_time_ms", "score": 100}, "check_fingerprint_reuse": {"score": 0}, "check_fingerprint_cycling": {"score": 0}, "check_ip_webrtc_ip_detail": {"score": 0}, "check_environment_critical": {"score": 0}, "check_useragent_other_enums": {"score": 0}, "check_useragent_ip_properties": {"score": 0}, "check_useragent_data_properties": {"score": 0}, "check_useragent_device_family_brand": {"score": 0}}"""
+ ),
+ "category_result": GrlIqForensicCategoryResult.model_validate_json(
+ """{"uuid": "b11441da5a854dfbb8401d4c32e56db5", "is_bot": 0, "is_tampered": 100, "is_velocity": 0, "is_anonymous": 0, "suspicious_ip": 0, "is_oscillating": 0, "is_teleporting": 0, "is_inconsistent": 0, "platform_ip_inconsistent": 0}"""
+ ),
+ "fraud_score": 100,
+ "is_attempt_allowed": False,
+ },
+ {
+ "data": GrlIqData.model_validate_json(
+ """{"mid": "35f6f5c30bc74ea7ac4aca7b40a02352", "uuid": "d54509f2f310499f8ab74839b10b2a41", "phase": "offerwall-enter", "events": null, "vendor": "Google Inc.", "app_name": "Netscape", "calendar": "gregory", "language": "en-US", "platform": "Linux x86_64", "timezone": "America/Los_Angeles", "client_ip": "104.9.125.144", "timestamp": "2025-02-28T11:34:39-08:00", "webrtc_ip": "172.56.209.195", "created_at": "2025-02-28T19:34:39.681872Z", "language_2": "en-US", "language_3": null, "platform_2": "Linux x86_64", "platform_3": null, "prefetched": true, "product_id": "d0606a0b5d034a8d81b1e3579d1f76fd", "webgl_flag": true, "webgl_hash": "da27e1b9b660057a3f5e185d3f5deabe", "canvas_hash": "e6e4d17da26050ce85ad00d3c6ea999e", "color_gamut": "3", "country_iso": "us", "inner_width": 841, "outer_width": 1680, "product_sub": "20030107", "audio_codecs": "1,1,1,1,1,3,1,3,1,3,3,1,1,3,3,3,3,1,3,3,3,2,1,1", "cookie_check": "", "graphics_api": "WebKit WebGL", "inner_height": 891, "mouse_events": null, "ontouchstart": false, "outer_height": 978, "plugins_hash": "4c05fa2f766a444d4f253ead792c8b0e|2", "screen_width": 1680, "video_codecs": "1,3,3,3,3,3,3,3,3,3,1,1,1,1,1,1,3,1,1,1,3,3,1", "webgl_hash_2": "fc73fd5db75e2c36222fe34251be3971", "webrtc_error": false, "window_opera": false, "battery_level": 0.41, "canvas_hash_2": "e0559d49b1864985cafc0d1c3a6b053c", "dynamic_range": false, "error_message": "Cannot read", "forced_colors": false, "math_result_1": "1.9275814160560204e-50", "math_result_2": "1.6182817135715877", "screen_height": 1050, "webgl_check_1": true, "webgl_context": "webgl2", "window_chrome": true, "connection_rtt": 100, "history_length": 11, "user_agent_str": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "web_sql_exists": false, "calender_locale": "en-US", "connection_type": "", "inverted_colors": true, "navigator_brave": false, "product_user_id": "test-unit", "request_headers": {"dnt": "1", "host": "127.0.0.1:8081", "accept": "application/json, lk/null q=0.1", "origin": "http://127.0.0.1:8080", "referer": "http://127.0.0.1:8080/", "sec-ch-ua": "\\"Google Chrome\\";v=\\"131\\", \\"Chromium\\";v=\\"131\\", \\"Not_A Brand\\";v=\\"24\\"", "connection": "keep-alive", "user-agent": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36", "content-type": "application/json", "content-length": "3313", "sec-fetch-dest": "empty", "sec-fetch-mode": "cors", "sec-fetch-site": "same-site", "accept-encoding": "gzip, deflate, br, zstd", "accept-language": "en-US,en;q=0.9", "sec-ch-ua-mobile": "?0", "sec-ch-ua-platform": "\\"Linux\\""}, "timezone_offset": 480, "webrtc_local_ip": "10.253.217.45,[2607:fb91:20c5:c6af:cda0:10b4:830a:a85e]", "battery_charging": false, "client_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "104.9.125.144", "isp": "AT&T Internet", "latitude": 37.3897, "city_name": "Mountain View", "longitude": -122.083, "time_zone": "America/Los_Angeles", "user_type": "residential", "country_iso": "us", "postal_code": "94041", "is_anonymous": false, "accuracy_radius": 5, "static_ip_score": 40.3, "subdivision_1_iso": "CA", "subdivision_2_iso": null, "subdivision_1_name": "California", "subdivision_2_name": null, "registered_country_iso": "us"}, "max_touch_points": 0, "numbering_system": "latn", "path_fingerprint": 3252, "prefers_contrast": "0", "rendering_engine": "WebKit", "timezone_success": "pass", "user_agent_hints": {"model": null, "brands": [{"brand": "Google Chrome", "version": "131"}, {"brand": "Chromium", "version": "131"}, {"brand": "Not_A Brand", "version": "24"}], "mobile": false, "bitness": "64", "platform": "Linux", "brands_full": [{"brand": "Google Chrome", "version": "131.0.6778.204"}, {"brand": "Chromium", "version": "131.0.6778.204"}, {"brand": "Not_A Brand", "version": "24.0.0.0"}], "architecture": "x86", "platform_version": "6.2.0"}, "user_agent_str_2": null, "webgl_extensions": "EXT_clip_control|EXT_color_buffer_float|EXT_color_buffer_half_float|EXT_conservative_depth|EXT_depth_clamp|EXT_disjoint_timer_query_webgl2|EXT_float_blend|EXT_polygon_offset_clamp|EXT_render_snorm|EXT_texture_compression_bptc|EXT_texture_compression_rgtc|EXT_texture_filter_anisotropic|EXT_texture_mirror_clamp_to_edge|EXT_texture_norm16|KHR_parallel_shader_compile|NV_shader_noperspective_interpolation|OES_draw_buffers_indexed|OES_sample_variables|OES_shader_multisample_interpolation|OES_texture_float_linear|OVR_multiview2|WEBGL_blend_func_extended|WEBGL_clip_cull_distance|WEBGL_compressed_texture_astc|WEBGL_compressed_texture_etc|WEBGL_compressed_texture_etc1|WEBGL_compressed_texture_s3tc|WEBGL_compressed_texture_s3tc_srgb|WEBGL_debug_renderer_info|WEBGL_debug_shaders|WEBGL_lose_context|WEBGL_multi_draw|WEBGL_polygon_mode|WEBGL_provoking_vertex|WEBGL_stencil_texturing", "webrtc_ip_detail": {"continent_code": "EU", "continent_name": "Europe", "country_name": "France", "is_in_european_union": true, "ip": "172.56.209.195", "isp": null, "latitude": null, "city_name": null, "longitude": null, "time_zone": null, "user_type": null, "country_iso": "us", "postal_code": null, "is_anonymous": null, "accuracy_radius": null, "static_ip_score": null, "subdivision_1_iso": null, "subdivision_2_iso": null, "subdivision_1_name": null, "subdivision_2_name": null, "registered_country_iso": null}, "chrome_extensions": "", "execution_time_ms": 924.5, "graphics_renderer": "WebGL 2.0 (OpenGL ES 3.0 Chromium)", "keyboard_detected": true, "mime_types_length": 2, "request_fs_exists": true, "audio_context_flag": "pass", "audio_context_hash": "9307303774dec3248c18a939392090da", "canvas_fingerprint": 258, "canvas_pixel_check": false, "device_pixel_ratio": 1.0, "indexedDbData_blob": true, "navigator_keys_len": 79, "no_edge_pdf_plugin": false, "screen_avail_width": 1680, "webdriver_detected": false, "window_orientation": 0, "connection_downlink": 10.0, "navigator_webdriver": false, "non_native_function": false, "screen_avail_height": 1010, "supported_fonts_str": "72|17152|327680|1073741824|0|0|540736|73728|7340032|1342177280|117446657|256|16|0|262687|4290797636|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0|2147483648|109543680|1880099888|301989903", "text_2d_fingerprint": "bfcce91c9e71d11af7b14dbee4c75f83", "webrtc_is_supported": "pass", "canvas_support_level": "full", "do_not_track_enabled": "1", "hardware_concurrency": 12, "keyboard_layout_size": 48, "prefers_color_scheme": false, "webgl_max_anisotropy": 16, "battery_charging_time": 0.0, "browser_by_properties": "c", "eval_to_string_length": 33, "performance_loop_time": 0.09999999962747097, "session_storage_check": "pass", "unmasked_vendor_webgl": "Google Inc. (Intel)", "hardware_concurrency_2": 12, "hardware_concurrency_3": null, "localStorage_available": true, "memory_jsHeapSizeLimit": 4294705152, "mozilla_web_app_exists": false, "navigator_deviceMemory": 8.0, "navigator_java_enabled": false, "prefers_reduced_motion": false, "storage_estimate_quota": 1178717110272, "webdriver_detected_msg": "", "window_active_x_object": false, "window_external_exists": true, "color_depth_pixel_depth": "24-24", "indexedDbData_available": true, "navigator_cookieEnabled": true, "unmasked_renderer_webgl": "ANGLE (Intel, Mesa Intel(R) Graphics (RPL-P), OpenGL 4.6)", "battery_discharging_time": 4844.0, "connection_effectiveType": "4g", "non_native_function_flag": "", "speech_synthesis_voice_1": "Google Bahasa Indonesia", "window_client_information": true, "audio_compressor_reduction": 20.538288116455078, "navigator_mediaDevices_len": 8, "audio_intensity_fingerprint": 124.04347527516074, "speech_synthesis_voice_hash": "8010ee3313813de521e48e63bd5a6f13", "microsoft_credentials_exists": false, "window_installTrigger_exists": false, "speech_synthesis_voices_count": 19, "webgl_shading_language_version": "WebGL GLSL ES 3.00 (OpenGL ES GLSL ES 3.0 Chromium)", "error_message_stack_access_count": 2, "speech_synthesis_avail_voices_count": 19, "error_message_stack_access_count_worker": 2}"""
+ ),
+ "result_data": GrlIqCheckerResults.model_validate_json(
+ """{"uuid": "d54509f2f310499f8ab74839b10b2a41", "check_codecs": {"score": 0}, "check_timezone": {"score": 0}, "check_timestamp": {"score": 0}, "check_user_type": {"score": 0}, "check_ip_changes": {"score": 0}, "check_ip_country": {"score": 0}, "check_environment": {"msg": "error_message_stack_access_count: 2", "score": 100}, "check_ip_timezone": {"score": 0}, "check_isp_changes": {"score": 0}, "check_useragent_js": {"score": 0}, "check_required_fonts": {"score": 0}, "check_user_anonymous": {"score": 0}, "check_webrtc_success": {"score": 0}, "check_seen_timestamps": {"score": 0}, "check_country_timezone": {"score": 0}, "check_prohibited_fonts": {"score": 0}, "check_timezone_changes": {"score": 0}, "check_execution_time_ms": {"score": 0}, "check_fingerprint_reuse": {"score": 0}, "check_fingerprint_cycling": {"score": 0}, "check_ip_webrtc_ip_detail": {"score": 0}, "check_environment_critical": {"score": 0}, "check_useragent_other_enums": {"score": 0}, "check_useragent_ip_properties": {"score": 0}, "check_useragent_data_properties": {"score": 0}, "check_useragent_device_family_brand": {"score": 0}}"""
+ ),
+ "category_result": GrlIqForensicCategoryResult.model_validate_json(
+ """{"uuid": "d54509f2f310499f8ab74839b10b2a41", "is_bot": 0, "is_tampered": 0, "is_velocity": 0, "is_anonymous": 0, "suspicious_ip": 0, "is_oscillating": 0, "is_teleporting": 0, "is_inconsistent": 10, "platform_ip_inconsistent": 0}"""
+ ),
+ "fraud_score": 10,
+ "is_attempt_allowed": True,
+ },
+ ]
diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py
index 88eef72..499f90b 100644
--- a/test_utils/incite/collections/conftest.py
+++ b/test_utils/incite/collections/conftest.py
@@ -1,16 +1,16 @@
from __future__ import annotations
+from collections.abc import Callable
from datetime import datetime, timedelta
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.pg_helper import PostgresConfig
from test_utils.conftest import clear_directory
if TYPE_CHECKING:
from generalresearch.incite.base import DFCollectionType, GRLDatasets
- from generalresearch.incite.collections import DFCollection
+ from generalresearch.incite.collections.base import DFCollection
from generalresearch.incite.collections.thl_web import (
AuditLogDFCollection,
LedgerDFCollection,
@@ -19,6 +19,7 @@ if TYPE_CHECKING:
UserDFCollection,
WallDFCollection,
)
+ from generalresearch.pg_helper import PostgresConfig
@pytest.fixture
@@ -196,7 +197,7 @@ def df_collection(
utc_90days_ago: datetime,
thl_web_rr: PostgresConfig,
) -> DFCollection:
- from generalresearch.incite.collections import DFCollection
+ from generalresearch.incite.collections.base import DFCollection
start = utc_90days_ago.replace(microsecond=0)
diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py
index 12e57c5..bcf0511 100644
--- a/test_utils/incite/conftest.py
+++ b/test_utils/incite/conftest.py
@@ -1,11 +1,12 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from os.path import join as pjoin
from pathlib import Path
from random import choice as randchoice
from shutil import rmtree
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -15,11 +16,11 @@ from faker import Faker
if TYPE_CHECKING:
from generalresearch.config import GRLBaseSettings
from generalresearch.incite.base import GRLDatasets
- from generalresearch.incite.collections import (
+ from generalresearch.incite.collections.base import (
DFCollectionItem,
DFCollectionType,
)
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.models.admin.request import (
ReportRequest,
)
@@ -130,14 +131,14 @@ def duration() -> timedelta | None:
@pytest.fixture
def df_collection_data_type() -> DFCollectionType:
- from generalresearch.incite.collections import DFCollectionType
+ from generalresearch.incite.collections.base import DFCollectionType
return DFCollectionType.TEST
@pytest.fixture
def merge_type() -> MergeType:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
return MergeType.TEST
@@ -155,7 +156,7 @@ def incite_item_factory(
observations: int = 3,
user: User | None = None,
):
- from generalresearch.incite.collections import (
+ from generalresearch.incite.collections.base import (
DFCollection,
DFCollectionType,
)
@@ -166,7 +167,7 @@ def incite_item_factory(
for _ in range(5):
item_time = fake.date_time_between(
- start_date=item.start, end_date=item.finish, tzinfo=timezone.utc
+ start_date=item.start, end_date=item.finish, tzinfo=UTC
)
match data_type:
diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py
index e9970c2..fb95c81 100644
--- a/test_utils/incite/mergers/conftest.py
+++ b/test_utils/incite/mergers/conftest.py
@@ -1,38 +1,41 @@
from __future__ import annotations
+from collections.abc import Callable
from datetime import datetime, timedelta
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.incite.base import GRLDatasets
-from generalresearch.incite.mergers import MergeType
-from generalresearch.incite.mergers.foundations.enriched_session import (
- EnrichedSessionMerge,
-)
-from generalresearch.incite.mergers.foundations.enriched_task_adjust import (
- EnrichedTaskAdjustMerge,
-)
-from generalresearch.incite.mergers.foundations.enriched_wall import (
- EnrichedWallMerge,
-)
-from generalresearch.incite.mergers.foundations.user_id_product import (
- UserIdProductMerge,
-)
-from generalresearch.incite.mergers.pop_ledger import (
- PopLedgerMerge,
- PopLedgerMergeItem,
-)
-from generalresearch.incite.mergers.ym_survey_wall import (
- YMSurveyWallMerge,
- YMSurveyWallMergeCollectionItem,
-)
-from generalresearch.incite.mergers.ym_wall_summary import (
- YMWallSummaryMerge,
- YMWallSummaryMergeItem,
-)
from test_utils.conftest import clear_directory
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.mergers.base import MergeType
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_task_adjust import (
+ EnrichedTaskAdjustMerge,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+ )
+ from generalresearch.incite.mergers.foundations.user_id_product import (
+ UserIdProductMerge,
+ )
+ from generalresearch.incite.mergers.pop_ledger import (
+ PopLedgerMerge,
+ PopLedgerMergeItem,
+ )
+ from generalresearch.incite.mergers.ym_survey_wall import (
+ YMSurveyWallMerge,
+ YMSurveyWallMergeCollectionItem,
+ )
+ from generalresearch.incite.mergers.ym_wall_summary import (
+ YMWallSummaryMerge,
+ YMWallSummaryMergeItem,
+ )
+
# --------------------------
# Merges
# --------------------------
@@ -55,7 +58,7 @@ def pop_ledger_merge(
duration: timedelta,
) -> PopLedgerMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
return PopLedgerMerge(
@@ -85,7 +88,7 @@ def ym_survey_wall_merge(
mnt_filepath: GRLDatasets,
start: datetime,
) -> YMSurveyWallMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge
return YMSurveyWallMerge(
@@ -116,7 +119,7 @@ def ym_wall_summary_merge(
duration: timedelta,
start: datetime,
) -> YMWallSummaryMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.ym_wall_summary import YMWallSummaryMerge
return YMWallSummaryMerge(
@@ -152,7 +155,7 @@ def enriched_session_merge(
duration: timedelta,
start: datetime,
) -> EnrichedSessionMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.foundations.enriched_session import (
EnrichedSessionMerge,
)
@@ -172,7 +175,7 @@ def enriched_task_adjust_merge(
duration: timedelta,
start: datetime,
) -> EnrichedTaskAdjustMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.foundations.enriched_task_adjust import (
EnrichedTaskAdjustMerge,
)
@@ -194,7 +197,7 @@ def enriched_wall_merge(
duration: timedelta,
start: datetime,
) -> EnrichedWallMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.foundations.enriched_wall import (
EnrichedWallMerge,
)
@@ -214,7 +217,7 @@ def user_id_product_merge(
offset: str,
start: datetime,
) -> UserIdProductMerge:
- from generalresearch.incite.mergers import MergeType
+ from generalresearch.incite.mergers.base import MergeType
from generalresearch.incite.mergers.foundations.user_id_product import (
UserIdProductMerge,
)
@@ -240,7 +243,7 @@ def merge_collection(
duration: timedelta,
start: datetime,
):
- from generalresearch.incite.mergers import MergeCollection
+ from generalresearch.incite.mergers.base import MergeCollection
return MergeCollection(
merge_type=merge_type,
diff --git a/test_utils/managers/cashout_methods.py b/test_utils/managers/cashout_methods.py
index b201e8c..e69de29 100644
--- a/test_utils/managers/cashout_methods.py
+++ b/test_utils/managers/cashout_methods.py
@@ -1,75 +0,0 @@
-import random
-from uuid import uuid4
-
-from generalresearch.models.thl.wallet import Currency, PayoutType
-from generalresearch.models.thl.wallet.cashout_method import (
- CashoutMethod,
- TangoCashoutMethodData,
-)
-
-
-def random_ext_id(base: str = "U02"):
- suffix = random.randint(0, 99999)
- return f"{base}{suffix:05d}"
-
-
-EXAMPLE_TANGO_CASHOUT_METHODS = [
- CashoutMethod(
- id=uuid4().hex,
- last_updated="2021-06-23T20:45:38.239182Z",
- is_live=True,
- type=PayoutType.TANGO,
- ext_id=random_ext_id(),
- name="Safeway eGift Card $25",
- data=TangoCashoutMethodData(
- value_type="fixed", countries=["US"], utid=random_ext_id()
- ),
- user=None,
- image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b694446-1200w-326ppi.png",
- original_currency=Currency.USD,
- min_value=2500,
- max_value=2500,
- ),
- CashoutMethod(
- id=uuid4().hex,
- last_updated="2021-06-23T20:45:38.239182Z",
- is_live=True,
- type=PayoutType.TANGO,
- ext_id=random_ext_id(),
- name="Amazon.it Gift Certificate",
- data=TangoCashoutMethodData(
- value_type="variable", countries=["IT"], utid="U006961"
- ),
- user=None,
- image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b405753-1200w-326ppi.png",
- original_currency=Currency.EUR,
- min_value=1,
- max_value=10000,
- ),
-]
-
-# AMT_ASSIGNMENT_CASHOUT_METHOD = CashoutMethod(
-# id=uuid4().hex,
-# last_updated="2021-06-23T20:45:38.239182Z",
-# is_live=True,
-# type=PayoutType.AMT,
-# ext_id=None,
-# name="AMT Assignment",
-# data=AmtCashoutMethodData(),
-# user=None,
-# min_value=1,
-# max_value=5,
-# )
-
-# AMT_BONUS_CASHOUT_METHOD = CashoutMethod(
-# id=uuid4().hex,
-# last_updated="2021-06-23T20:45:38.239182Z",
-# is_live=True,
-# type=PayoutType.AMT,
-# ext_id=None,
-# name="AMT Bonus",
-# data=AmtCashoutMethodData(),
-# user=None,
-# min_value=7,
-# max_value=4000,
-# )
diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py
index d2e5d20..391e6bf 100644
--- a/test_utils/managers/conftest.py
+++ b/test_utils/managers/conftest.py
@@ -1,39 +1,42 @@
from __future__ import annotations
-from typing import Callable
+import random
+from collections.abc import Callable
+from datetime import datetime
+from typing import TYPE_CHECKING
+from uuid import uuid4
import pytest
-from generalresearch.managers.gr.business import (
- BusinessAddressManager,
- BusinessBankAccountManager,
- BusinessManager,
+from generalresearch.managers.thl.cashout_method import (
+ CashoutMethodManager,
)
-from generalresearch.managers.gr.team import (
- MembershipManager,
- TeamManager,
+from generalresearch.managers.thl.user_streak import (
+ UserStreakManager,
)
-from generalresearch.managers.spectrum.survey import SpectrumSurveyManager
-from generalresearch.managers.thl.buyer import BuyerManager
-from generalresearch.managers.thl.ipinfo import (
- GeoIpInfoManager,
- IPGeonameManager,
- IPInformationManager,
-)
-from generalresearch.managers.thl.profiling.uqa import UQAManager
-from generalresearch.managers.thl.userhealth import (
- AuditLogManager,
- IPRecordManager,
- UserIpHistoryManager,
-)
-from generalresearch.models import Source
-from generalresearch.models.thl.user import User
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
-from generalresearch.sql_helper import SqlHelper
-from test_utils.managers.cashout_methods import (
- EXAMPLE_TANGO_CASHOUT_METHODS,
+from generalresearch.models.definitions import Source
+from generalresearch.models.thl.wallet.cashout_method import (
+ CashoutMethod,
+ TangoCashoutMethodData,
)
+from generalresearch.models.thl.wallet.definitions import Currency, PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.managers.spectrum.survey import SpectrumSurveyManager
+ from generalresearch.managers.thl.buyer import BuyerManager
+ from generalresearch.managers.thl.ipinfo import (
+ GeoIpInfoManager,
+ IPGeonameManager,
+ )
+ from generalresearch.managers.thl.userhealth import (
+ AuditLogManager,
+ IPRecordManager,
+ UserIpHistoryManager,
+ )
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
+ from generalresearch.sql_helper import SqlHelper
# === THL ===
@@ -59,16 +62,6 @@ def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager:
@pytest.fixture(scope="session")
-def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager:
- assert thl_web_rw.dsn.path
- assert "/unittest-" in thl_web_rw.dsn.path
-
- from generalresearch.managers.thl.ipinfo import IPInformationManager
-
- return IPInformationManager(pg_config=thl_web_rw)
-
-
-@pytest.fixture(scope="session")
def ip_record_manager(
thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig
) -> IPRecordManager:
@@ -95,7 +88,7 @@ def user_iphistory_manager(
@pytest.fixture(scope="function")
-def user_iphistory_manager_clear_cache(user_iphistory_manager, user):
+def user_iphistory_manager_clear_cache(user_iphistory_manager, user: User):
# On successive py-test/jenkins runs, the cache may contain
# the previous run's info (keyed under the same user_id)
user_iphistory_manager.delete_user_ip_history_cache(user_id=user.user_id)
@@ -116,12 +109,9 @@ def geoipinfo_manager(
@pytest.fixture(scope="session")
-def cashout_method_manager(thl_web_rw: PostgresConfig):
+def cashout_method_manager(thl_web_rw: PostgresConfig) -> CashoutMethodManager:
assert thl_web_rw.dsn.path
assert "/unittest-" in thl_web_rw.dsn.path
- from generalresearch.managers.thl.cashout_method import (
- CashoutMethodManager,
- )
return CashoutMethodManager(pg_config=thl_web_rw)
@@ -134,12 +124,9 @@ def event_manager(thl_redis_config: RedisConfig):
@pytest.fixture(scope="session")
-def user_streak_manager(thl_web_rw: PostgresConfig):
+def user_streak_manager(thl_web_rw: PostgresConfig) -> UserStreakManager:
assert thl_web_rw.dsn.path
assert "/unittest-" in thl_web_rw.dsn.path
- from generalresearch.managers.thl.user_streak import (
- UserStreakManager,
- )
return UserStreakManager(pg_config=thl_web_rw)
@@ -171,123 +158,117 @@ def delete_cashoutmethod_db(thl_web_rw: PostgresConfig) -> Callable[..., None]:
@pytest.fixture(scope="session")
-def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db):
- delete_cashoutmethod_db()
- for x in EXAMPLE_TANGO_CASHOUT_METHODS:
- cashout_method_manager.create(x)
-
- # TODO: convert these ids into instances to use.
- # settings.amt_bonus_cashout_method_id
- # settings.amt_assignment_cashout_method_id
-
- # cashout_method_manager.create(AMT_ASSIGNMENT_CASHOUT_METHOD)
- # cashout_method_manager.create(AMT_BONUS_CASHOUT_METHOD)
- raise NotImplementedError("Need to implement setup_cashoutmethod_db")
-
-
-# === THL: Marketplaces ===
+def setup_cashoutmethod_db(
+ cashout_method_manager: CashoutMethodManager,
+ delete_cashoutmethod_db: Callable[..., None],
+ example_tango_cashout_methods: list[CashoutMethod],
+) -> Callable[..., None]:
+ def _inner():
+ delete_cashoutmethod_db()
-@pytest.fixture(scope="session")
-def spectrum_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager:
- from generalresearch.managers.spectrum.survey import (
- SpectrumSurveyManager,
- )
+ for x in example_tango_cashout_methods:
+ cashout_method_manager.create(x)
- return SpectrumSurveyManager(sql_helper=spectrum_rw)
+ return _inner
-# === GR ===
@pytest.fixture(scope="session")
-def business_manager(
- gr_db: PostgresConfig, gr_redis_config: RedisConfig
-) -> BusinessManager:
- from generalresearch.redis_helper import RedisConfig
+def random_ext_id_factory(base: str = "U02") -> Callable[..., str]:
- assert gr_db.dsn.path
- assert "/unittest-" in gr_db.dsn.path
- assert isinstance(gr_redis_config, RedisConfig)
+ def _inner() -> str:
+ suffix = random.randint(0, 99999)
+ return f"{base}{suffix:05d}"
- from generalresearch.managers.gr.business import BusinessManager
-
- return BusinessManager(
- pg_config=gr_db,
- redis_config=gr_redis_config,
- )
+ return _inner
@pytest.fixture(scope="session")
-def business_address_manager(gr_db: PostgresConfig) -> BusinessAddressManager:
- assert gr_db.dsn.path
- assert "/unittest-" in gr_db.dsn.path
+def example_tango_cashout_methods(
+ random_ext_id_factory: Callable[..., str],
+) -> list[CashoutMethod]:
+ return [
+ CashoutMethod(
+ id=uuid4().hex,
+ last_updated=datetime.fromisoformat("2021-06-23T20:45:38.239182Z"),
+ is_live=True,
+ type=PayoutType.TANGO,
+ ext_id='U025035',
+ name="Safeway eGift Card $25",
+ data=TangoCashoutMethodData(
+ value_type="fixed", countries=["US"], utid='U025035'
+ ),
+ user=None,
+ image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b694446-1200w-326ppi.png",
+ original_currency=Currency.USD,
+ min_value=2500,
+ max_value=2500,
+ ),
+ CashoutMethod(
+ id=uuid4().hex,
+ last_updated=datetime.fromisoformat("2021-06-23T20:45:38.239182Z"),
+ is_live=True,
+ type=PayoutType.TANGO,
+ ext_id='U006961',
+ name="Amazon.it Gift Certificate",
+ data=TangoCashoutMethodData(
+ value_type="variable", countries=["IT"], utid="U006961"
+ ),
+ user=None,
+ image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b405753-1200w-326ppi.png",
+ original_currency=Currency.EUR,
+ min_value=1,
+ max_value=10000,
+ ),
+ ]
- from generalresearch.managers.gr.business import BusinessAddressManager
- return BusinessAddressManager(pg_config=gr_db)
+# === THL: Marketplaces ===
@pytest.fixture(scope="session")
-def business_bank_account_manager(
- gr_db: PostgresConfig,
-) -> BusinessBankAccountManager:
- assert gr_db.dsn.path
- assert "/unittest-" in gr_db.dsn.path
-
- from generalresearch.managers.gr.business import (
- BusinessBankAccountManager,
+def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager:
+ from generalresearch.managers.spectrum.survey import (
+ SpectrumSurveyManager,
)
- return BusinessBankAccountManager(pg_config=gr_db)
-
-
-@pytest.fixture(scope="session")
-def team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> TeamManager:
- assert gr_db.dsn.path
- assert "/unittest-" in gr_db.dsn.path
-
- from generalresearch.managers.gr.team import TeamManager
-
- return TeamManager(pg_config=gr_db, redis_config=gr_redis_config)
+ return SpectrumSurveyManager(sql_helper=spectrum_rw)
@pytest.fixture(scope="session")
-def membership_manager(gr_db: PostgresConfig) -> MembershipManager:
- assert gr_db.dsn.path
- assert "/unittest-" in gr_db.dsn.path
-
- from generalresearch.managers.gr.team import MembershipManager
-
- return MembershipManager(pg_config=gr_db)
+def delete_buyers_surveys(
+ thl_web_rw: PostgresConfig, buyer_manager: BuyerManager
+) -> Callable[..., None]:
+ def _inner():
+ # assert "/unittest-" in thl_web_rw.dsn.path
+ thl_web_rw.execute_write(
+ """
+ DELETE FROM marketplace_surveystat
+ WHERE survey_id IN (
+ SELECT id
+ FROM marketplace_survey
+ WHERE source = %(source)s
+ );""",
+ params={"source": Source.TESTING.value},
+ )
+ thl_web_rw.execute_write(
+ """
+ DELETE FROM marketplace_survey
+ WHERE buyer_id IN (
+ SELECT id
+ FROM marketplace_buyer
+ WHERE source = %(source)s
+ );""",
+ params={"source": Source.TESTING.value},
+ )
+ thl_web_rw.execute_write(
+ """
+ DELETE from marketplace_buyer
+ WHERE source=%(source)s;
+ """,
+ params={"source": Source.TESTING.value},
+ )
+ buyer_manager.populate_caches()
-@pytest.fixture(scope="session")
-def delete_buyers_surveys(thl_web_rw: PostgresConfig, buyer_manager: BuyerManager):
- # assert "/unittest-" in thl_web_rw.dsn.path
- thl_web_rw.execute_write(
- """
- DELETE FROM marketplace_surveystat
- WHERE survey_id IN (
- SELECT id
- FROM marketplace_survey
- WHERE source = %(source)s
- );""",
- params={"source": Source.TESTING.value},
- )
- thl_web_rw.execute_write(
- """
- DELETE FROM marketplace_survey
- WHERE buyer_id IN (
- SELECT id
- FROM marketplace_buyer
- WHERE source = %(source)s
- );""",
- params={"source": Source.TESTING.value},
- )
- thl_web_rw.execute_write(
- """
- DELETE from marketplace_buyer
- WHERE source=%(source)s;
- """,
- params={"source": Source.TESTING.value},
- )
- buyer_manager.populate_caches()
+ return _inner
diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py
index 67935e7..b29cf18 100644
--- a/test_utils/managers/contest/conftest.py
+++ b/test_utils/managers/contest/conftest.py
@@ -1,8 +1,14 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.base import Permission
from generalresearch.managers.thl.contest_manager import ContestManager
-from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
@pytest.fixture(scope="session")
@@ -11,8 +17,6 @@ def contest_manager(thl_web_rw: PostgresConfig) -> ContestManager:
assert thl_web_rw.dsn.path
assert "/unittest-" in thl_web_rw.dsn.path
- from generalresearch.managers.thl.contest_manager import ContestManager
-
return ContestManager(
pg_config=thl_web_rw,
permissions=[
diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py
index 37da164..09e08f5 100644
--- a/test_utils/managers/gr/conftest.py
+++ b/test_utils/managers/gr/conftest.py
@@ -1,64 +1,76 @@
from __future__ import annotations
-from typing import Callable
+import subprocess
+from collections.abc import Callable, Generator
+from random import randint
+from typing import TYPE_CHECKING
import pytest
-import redis.asyncio as redis_async
+import redis
from pydantic import PostgresDsn
-from redis import Redis
-from generalresearch.config import GRLBaseSettings
-from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
from generalresearch.managers.gr.business import (
BusinessAddressManager,
BusinessBankAccountManager,
BusinessManager,
)
+from generalresearch.managers.gr.team import MembershipManager
from generalresearch.pg_helper import PostgresConfig
from generalresearch.redis_helper import RedisConfig
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
+ from generalresearch.managers.gr.team import TeamManager
+
# === Msc ===
-@pytest.fixture(scope="session")
-def gr_redis(settings: GRLBaseSettings) -> Redis:
- assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis)
- return Redis.from_url(
- url=str(settings.gr_redis),
- decode_responses=True,
- socket_timeout=settings.redis_timeout,
- socket_connect_timeout=settings.redis_timeout,
- )
-@pytest.fixture
-def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis:
- assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis)
+@pytest.fixture(scope="session")
+def gr_redis_config_db() -> str:
+ # need to update 'databases' in /etc/redis/redis.conf
+ # or this won't work and you'll have no indication why ...
+ return str(randint(99, 1_023))
- return redis_async.Redis.from_url(
- str(settings.gr_redis),
- decode_responses=True,
- socket_timeout=0.20,
- socket_connect_timeout=0.20,
+
+@pytest.fixture(scope="session")
+def gr_redis_config(
+ settings: GRLBaseSettings, gr_redis_config_db: str
+) -> Generator[RedisConfig]:
+ assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str(
+ settings.testing_redis
)
+ uri = f"redis://{settings.testing_redis}/{gr_redis_config_db}"
-@pytest.fixture(scope="session")
-def gr_redis_config(settings: GRLBaseSettings) -> RedisConfig:
- assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis)
+ res = subprocess.run(
+ ["redis-cli", "-u", uri, "SET", "jenkins_lock", "1", "NX", "EX", "3600"],
+ check=True,
+ text=True,
+ capture_output=True,
+ )
+
+ if res.stdout.strip() != "OK":
+ raise ValueError("Redis already locked... aborting.")
- return RedisConfig(
- dsn=settings.gr_redis,
+ yield RedisConfig(
+ dsn=uri,
decode_responses=True,
socket_timeout=settings.redis_timeout,
socket_connect_timeout=settings.redis_timeout,
)
+ r = redis.from_url(uri)
+ r.flushdb()
+
@pytest.fixture(scope="session")
def gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig:
+ _dsn = django_db_factory("gr.common")
return PostgresConfig(
- dsn=django_db_factory("gr_carer"),
+ dsn=_dsn,
connect_timeout=1,
statement_timeout=5,
)
@@ -80,7 +92,17 @@ def gr_user_manager(
@pytest.fixture(scope="session")
-def gr_team_manager(gr_db: PostgresConfig) -> GRTokenManager:
+def gr_team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> TeamManager:
+ assert gr_db.dsn.path
+ assert "/unittest-" in gr_db.dsn.path
+
+ from generalresearch.managers.gr.team import TeamManager
+
+ return TeamManager(pg_config=gr_db, redis_config=gr_redis_config)
+
+
+@pytest.fixture(scope="session")
+def gr_token_manager(gr_db: PostgresConfig) -> GRTokenManager:
assert gr_db.dsn.path
assert "/unittest-" in gr_db.dsn.path
@@ -108,3 +130,10 @@ def gr_business_address_manager(
gr_db: PostgresConfig,
) -> BusinessAddressManager:
return BusinessAddressManager(pg_config=gr_db)
+
+
+@pytest.fixture(scope="session")
+def gr_membership_manager(
+ gr_db: PostgresConfig,
+) -> MembershipManager:
+ return MembershipManager(pg_config=gr_db)
diff --git a/test_utils/managers/ledger/conftest.py b/test_utils/managers/ledger/conftest.py
index ce8348e..c60ee1b 100644
--- a/test_utils/managers/ledger/conftest.py
+++ b/test_utils/managers/ledger/conftest.py
@@ -1,18 +1,24 @@
from __future__ import annotations
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.base import Permission
from generalresearch.managers.thl.ledger_manager.ledger import (
- LedgerAccountManager,
LedgerManager,
- LedgerTransactionManager,
)
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import (
+ LedgerAccountManager,
+ LedgerTransactionManager,
+ )
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
# --- Ledger ---
diff --git a/test_utils/managers/network/__init__.py b/test_utils/managers/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/test_utils/managers/network/__init__.py
+++ /dev/null
diff --git a/test_utils/managers/network/conftest.py b/test_utils/managers/network/conftest.py
deleted file mode 100644
index e69de29..0000000
--- a/test_utils/managers/network/conftest.py
+++ /dev/null
diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py
index 5b70961..355a39d 100644
--- a/test_utils/managers/thl/conftest.py
+++ b/test_utils/managers/thl/conftest.py
@@ -1,44 +1,108 @@
from __future__ import annotations
-from typing import Callable
+import subprocess
+from collections.abc import Callable, Generator
+from random import randint
+from typing import TYPE_CHECKING
import pytest
+import redis
from pydantic import PostgresDsn
-from generalresearch.config import GRLBaseSettings
from generalresearch.managers.base import Permission
-from generalresearch.managers.thl.buyer import BuyerManager
-from generalresearch.managers.thl.category import CategoryManager
-from generalresearch.managers.thl.payout import (
- BrokerageProductPayoutEventManager,
- BusinessPayoutEventManager,
- PayoutEventManager,
- UserPayoutEventManager,
+from generalresearch.managers.thl.user_manager.mysql_user_manager import (
+ MysqlUserManager,
)
-from generalresearch.managers.thl.product import ProductManager
-from generalresearch.managers.thl.session import SessionManager
-from generalresearch.managers.thl.task_adjustment import (
- TaskAdjustmentManager,
-)
-from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
-)
-from generalresearch.managers.thl.user_manager.user_metadata_manager import (
- UserMetadataManager,
-)
-from generalresearch.managers.thl.wall import (
- WallCacheManager,
- WallManager,
+from generalresearch.managers.thl.user_manager.redis_user_manager import (
+ RedisUserManager,
)
from generalresearch.pg_helper import PostgresConfig
from generalresearch.redis_helper import RedisConfig
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.managers.thl.buyer import BuyerManager
+ from generalresearch.managers.thl.category import CategoryManager
+ from generalresearch.managers.thl.ipinfo import (
+ IPGeonameManager,
+ IPInformationManager,
+ )
+ from generalresearch.managers.thl.payout import (
+ BrokerageProductPayoutEventManager,
+ BusinessPayoutEventManager,
+ PayoutEventManager,
+ UserPayoutEventManager,
+ )
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.task_adjustment import (
+ TaskAdjustmentManager,
+ )
+ from generalresearch.managers.thl.user_manager.user_manager import (
+ UserManager,
+ )
+ from generalresearch.managers.thl.user_manager.user_metadata_manager import (
+ UserMetadataManager,
+ )
+ from generalresearch.managers.thl.userhealth import (
+ AuditLogManager,
+ IPRecordManager,
+ )
+ from generalresearch.managers.thl.wall import (
+ WallCacheManager,
+ WallManager,
+ )
+
+# === Msc ===
+
+
+@pytest.fixture(scope="session")
+def thl_redis_config_db() -> str:
+ return str(randint(99, 1_023))
+
+
+@pytest.fixture(scope="session")
+def thl_redis_config(
+ settings: GRLBaseSettings, thl_redis_config_db: str
+) -> Generator[RedisConfig]:
+ assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str(
+ settings.testing_redis
+ )
+
+ uri = f"redis://{settings.testing_redis}/{thl_redis_config_db}"
+
+ res = subprocess.run(
+ ["redis-cli", "-u", uri, "SET", "jenkins_lock", "1", "NX", "EX", "3600"],
+ check=True,
+ text=True,
+ capture_output=True,
+ )
+
+ if res.stdout.strip() != "OK":
+ raise ValueError("Redis already locked... aborting.")
+
+ yield RedisConfig(
+ dsn=uri,
+ decode_responses=True,
+ socket_timeout=settings.redis_timeout,
+ socket_connect_timeout=settings.redis_timeout,
+ )
+
+ r = redis.from_url(uri)
+ r.flushdb()
+
+
+@pytest.fixture(scope="session")
+def thl_redis_client(thl_redis_config):
+ return thl_redis_config.create_redis_client()
+
@pytest.fixture(scope="session")
def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig:
+ _dsn = django_db_factory("generalresearch.thl_django")
return PostgresConfig(
- dsn=django_db_factory("generalresearch.thl_django"),
+ dsn=_dsn,
connect_timeout=1,
statement_timeout=5,
)
@@ -49,14 +113,7 @@ def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig:
return thl_web_rr
-@pytest.fixture(scope="session")
-def thl_redis_config(settings: GRLBaseSettings) -> RedisConfig:
- return RedisConfig(
- dsn=settings.thl_redis,
- decode_responses=True,
- socket_timeout=settings.redis_timeout,
- socket_connect_timeout=settings.redis_timeout,
- )
+# === Managers ===
@pytest.fixture(scope="session")
@@ -109,6 +166,13 @@ def brokerage_product_payout_event_manager(
)
+@pytest.fixture()
+def audit_log_manager(thl_web_rw: PostgresConfig) -> AuditLogManager:
+ from generalresearch.managers.thl.userhealth import AuditLogManager
+
+ return AuditLogManager(pg_config=thl_web_rw)
+
+
@pytest.fixture(scope="session")
def business_payout_event_manager(
thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig
@@ -140,7 +204,10 @@ def product_manager(thl_web_rw: PostgresConfig) -> ProductManager:
@pytest.fixture(scope="session")
def user_manager(
- settings: GRLBaseSettings, thl_web_rw: PostgresConfig, thl_web_rr: PostgresConfig
+ settings: GRLBaseSettings,
+ thl_web_rw: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
) -> UserManager:
assert thl_web_rw.dsn
assert thl_web_rw.dsn.path
@@ -149,16 +216,32 @@ def user_manager(
assert "/unittest-" in thl_web_rw.dsn.path
assert "/unittest-" in thl_web_rr.dsn.path
+ from generalresearch.managers.thl.user_manager.rate_limit import UserManagerLimiter
from generalresearch.managers.thl.user_manager.user_manager import (
UserManager,
)
- return UserManager(
+ um = UserManager(
pg_config=thl_web_rw,
pg_config_rr=thl_web_rr,
redis=settings.redis,
)
+ # rc = thl_redis_config.create_redis_client()
+ um.user_manager_limiter = UserManagerLimiter(redis=thl_redis_config.dsn)
+
+ return um
+
+
+@pytest.fixture(scope="session")
+def mysql_user_manager(thl_web_rw: PostgresConfig) -> MysqlUserManager:
+ return MysqlUserManager(pg_config=thl_web_rw, is_read_replica=False)
+
+
+@pytest.fixture(scope="session")
+def redis_user_manager(thl_redis_config: RedisConfig) -> RedisUserManager:
+ return RedisUserManager(redis_dsn=thl_redis_config.dsn)
+
@pytest.fixture(scope="session")
def user_metadata_manager(thl_web_rw: PostgresConfig) -> UserMetadataManager:
@@ -256,3 +339,41 @@ def surveypenalty_manager(thl_redis_config: RedisConfig):
from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager
return SurveyPenaltyManager(redis_config=thl_redis_config)
+
+
+# --- IP Geolocation ---
+
+
+@pytest.fixture
+def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager:
+ from generalresearch.managers.thl.ipinfo import IPGeonameManager
+
+ return IPGeonameManager(pg_config=thl_web_rw)
+
+
+# --- IP Information ---
+
+
+@pytest.fixture(scope="session")
+def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager:
+ assert thl_web_rw.dsn.path
+ assert "/unittest-" in thl_web_rw.dsn.path
+
+ from generalresearch.managers.thl.ipinfo import IPInformationManager
+
+ return IPInformationManager(pg_config=thl_web_rw)
+
+
+# --- IP Record ---
+
+
+@pytest.fixture(scope="session")
+def ip_record_manager(
+ thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig
+) -> IPRecordManager:
+ assert thl_web_rw.dsn.path
+ assert "/unittest-" in thl_web_rw.dsn.path
+
+ from generalresearch.managers.thl.userhealth import IPRecordManager
+
+ return IPRecordManager(pg_config=thl_web_rw, redis_config=thl_redis_config)
diff --git a/test_utils/managers/upk/conftest.py b/test_utils/managers/upk/conftest.py
index d8f956c..23af1b3 100644
--- a/test_utils/managers/upk/conftest.py
+++ b/test_utils/managers/upk/conftest.py
@@ -1,4 +1,5 @@
-from typing import Callable, Generator
+from collections.abc import Callable, Generator
+from typing import TYPE_CHECKING
import pytest
diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py
index 9925a9e..ffce272 100644
--- a/test_utils/models/conftest.py
+++ b/test_utils/models/conftest.py
@@ -1,331 +1,27 @@
from __future__ import annotations
-from datetime import datetime, timedelta, timezone
-from decimal import Decimal
-from random import choice as randchoice
-from random import randint
-from typing import TYPE_CHECKING, Callable
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from fastapi import Request
-from pydantic import AwareDatetime, PositiveInt
+from pytest import FixtureRequest as Request
-from generalresearch.models import Source
-from generalresearch.models.thl.definitions import (
- WALL_ALLOWED_STATUS_STATUS_CODE,
- Status,
-)
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.survey.model import Buyer, Survey
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
if TYPE_CHECKING:
- from generalresearch.currency import USDCent
- from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
- from generalresearch.managers.gr.business import (
- BusinessAddressManager,
- BusinessBankAccountManager,
- BusinessManager,
- )
- from generalresearch.managers.gr.team import MembershipManager, TeamManager
from generalresearch.managers.thl.buyer import BuyerManager
- from generalresearch.managers.thl.ipinfo import (
- IPGeonameManager,
- IPInformationManager,
- )
- from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
- from generalresearch.managers.thl.payout import (
- BusinessPayoutEventManager,
- )
- from generalresearch.managers.thl.product import ProductManager
- from generalresearch.managers.thl.session import SessionManager
from generalresearch.managers.thl.survey import SurveyManager
- from generalresearch.managers.thl.user_manager.user_manager import UserManager
- from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager
- from generalresearch.managers.thl.wall import WallManager
- from generalresearch.models.gr.authentication import GRToken, GRUser
- from generalresearch.models.gr.business import (
- Business,
- BusinessAddress,
- BusinessBankAccount,
- )
- from generalresearch.models.gr.team import Membership, Team
- from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation
- from generalresearch.models.thl.payout import (
- BrokerageProductPayoutEvent,
- )
from generalresearch.models.thl.product import (
PayoutConfig,
Product,
)
- from generalresearch.models.thl.session import Session, Wall
- from generalresearch.models.thl.user import User
- from generalresearch.models.thl.user_iphistory import IPRecord
- from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
# === THL ===
-@pytest.fixture
-def user(
- request,
- product_manager: ProductManager,
- user_manager: UserManager,
- thl_web_rr: PostgresConfig,
-) -> User:
- product = getattr(request, "product", None)
-
- if product is None:
- product = product_manager.create_dummy()
-
- u = user_manager.create_dummy(product_id=product.id)
- u.prefetch_product(pg_config=thl_web_rr)
-
- return u
-
-
-@pytest.fixture
-def user_with_wallet(
- user_factory: Callable[..., User],
- product_user_wallet_yes: Product,
-) -> User:
- # A user on a product with user wallet enabled, but they have no money
- return user_factory(product=product_user_wallet_yes)
-
-
-@pytest.fixture
-def user_with_wallet_amt(
- user_factory: Callable[..., User], product_amt_true: Product
-) -> User:
- # A user on a product with user wallet enabled, on AMT, but they have no money
- return user_factory(product=product_amt_true)
-
-
-@pytest.fixture(scope="function")
-def user_factory(
- user_manager: UserManager, thl_web_rr: PostgresConfig
-) -> Callable[..., User]:
-
- def _inner(product: Product, created: datetime | None = None) -> User:
- u = user_manager.create_dummy(product=product, created=created)
- u.prefetch_product(pg_config=thl_web_rr)
-
- return u
-
- return _inner
-
-
-@pytest.fixture
-def wall_factory(wall_manager: WallManager) -> Callable[..., Wall]:
-
- def _inner(
- session: Session, wall_status: Status, req_cpi: Decimal | None = None
- ) -> Wall:
-
- assert session.started <= datetime.now(
- tz=timezone.utc
- ), "Session can't start in the future"
-
- if session.wall_events:
- # Subsequent Wall events
- wall = session.wall_events[-1]
- assert not wall.finished, "Can't add new Walls until prior finishes"
- # wall_started = last_wall.started + timedelta(milliseconds=1)
- else:
- # First Wall Event in a session
- wall_started = session.started + timedelta(milliseconds=1)
-
- wall = wall_manager.create_dummy(
- session_id=session.id,
- user_id=session.user_id,
- started=wall_started,
- req_cpi=req_cpi,
- )
- session.append_wall_event(w=wall)
-
- options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(wall_status, {}))
- wall.finish(
- finished=wall.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
- status=wall_status,
- status_code_1=randchoice(options),
- )
-
- return wall
-
- return _inner
-
-
-@pytest.fixture
-def wall(session: Session, user: User, wall_manager: WallManager) -> Wall | None:
- from generalresearch.models.thl.task_status import StatusCode1
-
- wall = wall_manager.create_dummy(session_id=session.id, user_id=user.user_id)
- # thl_session.append_wall_event(wall)
- wall.finish(
- finished=wall.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
- status=Status.COMPLETE,
- status_code_1=StatusCode1.COMPLETE,
- )
- return wall
-
-
-@pytest.fixture
-def session_factory(
- session_manager: SessionManager,
- wall_manager: WallManager,
- utc_hour_ago: datetime,
-) -> Callable[..., Session]:
- from generalresearch.models.thl.session import Source
-
- def _inner(
- user: User,
- # Wall details
- wall_count: int = 5,
- wall_req_cpi: Decimal = Decimal(".50"),
- wall_req_cpis: list[Decimal] | None = None,
- wall_statuses: list[Status] | None = None,
- wall_source: Source = Source.TESTING,
- # Session details
- final_status: Status = Status.COMPLETE,
- started: datetime = utc_hour_ago,
- ) -> Session:
- if wall_req_cpis:
- assert len(wall_req_cpis) == wall_count
- if wall_statuses:
- assert len(wall_statuses) == wall_count
-
- s = session_manager.create_dummy(started=started, user=user, country_iso="us")
- for idx in range(wall_count):
- if idx == 0:
- # First Wall Event in a session
- wall_started = s.started + timedelta(milliseconds=1)
- else:
- # Subsequent Wall events
- last_wall = s.wall_events[-1]
- assert last_wall.finished, "Can't add new Walls until prior finishes"
- wall_started = last_wall.started + timedelta(milliseconds=1)
-
- w = wall_manager.create_dummy(
- session_id=s.id,
- source=wall_source,
- user_id=s.user_id,
- started=wall_started,
- req_cpi=wall_req_cpis[idx] if wall_req_cpis else wall_req_cpi,
- )
- s.append_wall_event(w=w)
-
- # If it's the last wall in the session, respect the final_status
- # value for the Session
- if wall_statuses:
- _final_status = wall_statuses[idx]
- else:
- _final_status = final_status if idx == wall_count - 1 else Status.FAIL
-
- options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(_final_status, {}))
- wall_manager.finish(
- wall=w,
- status=_final_status,
- status_code_1=randchoice(options),
- finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
- )
-
- return s
-
- return _inner
-
-
-@pytest.fixture(scope="function")
-def finished_session_factory(
- session_factory: Callable[..., Session],
- session_manager: SessionManager,
- utc_hour_ago: datetime,
-) -> Callable[..., Session]:
- from generalresearch.models.thl.session import Source
-
- def _inner(
- user: User,
- # Wall details
- wall_count: int = 5,
- wall_req_cpi: Decimal = Decimal(".50"),
- wall_req_cpis: list[Decimal] | None = None,
- wall_statuses: list[Status] | None = None,
- wall_source: Source = Source.TESTING,
- # Session details
- final_status: Status = Status.COMPLETE,
- started: datetime = utc_hour_ago,
- ) -> Session:
- s: Session = session_factory(
- user=user,
- wall_count=wall_count,
- wall_req_cpi=wall_req_cpi,
- wall_req_cpis=wall_req_cpis,
- wall_statuses=wall_statuses,
- wall_source=wall_source,
- final_status=final_status,
- started=started,
- )
- status, status_code_1 = s.determine_session_status()
- _, _, bp_pay, user_pay = s.determine_payments()
- session_manager.finish_with_status(
- s,
- finished=s.wall_events[-1].finished,
- payout=bp_pay,
- user_payout=user_pay,
- status=status,
- status_code_1=status_code_1,
- )
- return s
-
- return _inner
-
-
-@pytest.fixture
-def session(
- user: User, session_manager: SessionManager, wall_manager: WallManager
-) -> Session:
-
- session: Session = session_manager.create_dummy(user=user, country_iso="us")
- wall: Wall = wall_manager.create_dummy(
- session_id=session.id,
- user_id=session.user_id,
- started=session.started,
- )
- session.append_wall_event(w=wall)
-
- return session
-
-
-@pytest.fixture
-def product(request: Request, product_manager: ProductManager) -> Product:
-
- team = getattr(request, "team", None)
- business = getattr(request, "business", None)
-
- return product_manager.create_dummy(
- team_id=team.uuid if team else None,
- business_id=business.uuid if business else None,
- )
-
-
-@pytest.fixture
-def product_factory(product_manager: ProductManager) -> Callable[..., Product]:
-
- def _inner(
- team: Team | None = None,
- business: Business | None = None,
- commission_pct: Decimal = Decimal("0.05"),
- ) -> Product:
- return product_manager.create_dummy(
- team_id=team.uuid if team else None,
- business_id=business.uuid if business else None,
- commission_pct=commission_pct,
- )
-
- return _inner
-
-
-@pytest.fixture
+@pytest.fixture()
def payout_config(request: Request) -> PayoutConfig:
from generalresearch.models.thl.product import (
PayoutConfig,
@@ -348,176 +44,38 @@ def payout_config(request: Request) -> PayoutConfig:
@pytest.fixture
def product_user_wallet_yes(
- payout_config: PayoutConfig, product_manager: ProductManager
+ product_factory: Callable[..., Product],
+ payout_config: PayoutConfig,
) -> Product:
from generalresearch.models.thl.product import UserWalletConfig
- return product_manager.create_dummy(
+ return product_factory(
payout_config=payout_config, user_wallet_config=UserWalletConfig(enabled=True)
)
@pytest.fixture
-def product_user_wallet_no(product_manager: ProductManager) -> Product:
+def product_user_wallet_no(
+ product_factory: Callable[..., Product],
+) -> Product:
from generalresearch.models.thl.product import UserWalletConfig
- return product_manager.create_dummy(
- user_wallet_config=UserWalletConfig(enabled=False)
- )
+ return product_factory(user_wallet_config=UserWalletConfig(enabled=False))
@pytest.fixture
def product_amt_true(
- product_manager: ProductManager, payout_config: PayoutConfig
+ product_factory: Callable[..., Product],
+ payout_config: PayoutConfig,
) -> Product:
from generalresearch.models.thl.product import UserWalletConfig
- return product_manager.create_dummy(
+ return product_factory(
user_wallet_config=UserWalletConfig(amt=True, enabled=True),
payout_config=payout_config,
)
-@pytest.fixture
-def bp_payout_factory(
- thl_lm: ThlLedgerManager,
- product_manager: ProductManager,
- business_payout_event_manager: BusinessPayoutEventManager,
-) -> Callable[..., BrokerageProductPayoutEvent]:
-
- def _inner(
- product: Product | None = None,
- amount: USDCent | None = None,
- ext_ref_id: str | None = None,
- created: AwareDatetime | None = None,
- skip_wallet_balance_check: bool = False,
- skip_one_per_day_check: bool = False,
- ) -> BrokerageProductPayoutEvent:
- from generalresearch.currency import USDCent
-
- product = product or product_manager.create_dummy()
- amount = amount or USDCent(randint(1, 99_99))
-
- return business_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=amount,
- ext_ref_id=ext_ref_id,
- created=created,
- skip_wallet_balance_check=skip_wallet_balance_check,
- skip_one_per_day_check=skip_one_per_day_check,
- )
-
- return _inner
-
-
-# === GR ===
-
-
-@pytest.fixture
-def business(request, business_manager: BusinessManager) -> Business:
- return business_manager.create_dummy()
-
-
-@pytest.fixture
-def business_address(
- request, business: Business, business_address_manager: BusinessAddressManager
-) -> BusinessAddress:
- return business_address_manager.create_dummy(business_id=business.id)
-
-
-@pytest.fixture
-def business_bank_account(
- request,
- business: Business,
- business_bank_account_manager: BusinessBankAccountManager,
-) -> BusinessBankAccount:
- return business_bank_account_manager.create_dummy(business_id=business.id)
-
-
-@pytest.fixture
-def team(request, team_manager: TeamManager) -> Team:
- return team_manager.create_dummy()
-
-
-@pytest.fixture
-def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog:
-
- return audit_log_manager.create_dummy(user_id=user.user_id)
-
-
-@pytest.fixture
-def audit_log_factory(
- audit_log_manager: AuditLogManager,
-) -> Callable[..., AuditLog]:
-
- def _inner(
- user_id: PositiveInt,
- level: AuditLogLevel | None = None,
- event_type: str | None = None,
- event_msg: str | None = None,
- event_value: float | None = None,
- ) -> AuditLog:
- return audit_log_manager.create_dummy(
- user_id=user_id,
- level=level,
- event_type=event_type,
- event_msg=event_msg,
- event_value=event_value,
- )
-
- return _inner
-
-
-@pytest.fixture
-def ip_geoname(ip_geoname_manager: IPGeonameManager) -> IPGeoname:
- return ip_geoname_manager.create_dummy()
-
-
-@pytest.fixture
-def ip_information(
- ip_information_manager: IPInformationManager, ip_geoname: IPGeoname
-) -> IPInformation:
- return ip_information_manager.create_dummy(
- geoname_id=ip_geoname.geoname_id, country_iso=ip_geoname.country_iso
- )
-
-
-@pytest.fixture
-def ip_information_factory(
- ip_information_manager: IPInformationManager,
-) -> Callable[..., IPInformation]:
-
- def _inner(ip: str, geoname: IPGeoname, **kwargs) -> IPInformation:
- return ip_information_manager.create_dummy(
- ip=ip,
- geoname_id=geoname.geoname_id,
- country_iso=geoname.country_iso,
- **kwargs,
- )
-
- return _inner
-
-
-@pytest.fixture
-def ip_record(
- ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User
-) -> IPRecord:
-
- return ip_record_manager.create_dummy(user_id=user.user_id)
-
-
-@pytest.fixture
-def ip_record_factory(
- ip_record_manager: IPRecordManager, user: User
-) -> Callable[..., IPRecord]:
-
- def _inner(user_id: PositiveInt, ip: str | None = None) -> IPRecord:
- return ip_record_manager.create_dummy(user_id=user_id, ip=ip)
-
- return _inner
-
-
@pytest.fixture(scope="session")
def buyer(buyer_manager: BuyerManager) -> Buyer:
buyer_code = uuid4().hex
diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py
index e750076..18a8e5f 100644
--- a/test_utils/models/contest/conftest.py
+++ b/test_utils/models/contest/conftest.py
@@ -1,50 +1,73 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from fastapi import Request
+from pytest import FixtureRequest as Request
from generalresearch.currency import USDCent
-from generalresearch.managers.thl.contest_manager import ContestManager
-from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
-from generalresearch.models.thl.contest.contest import Contest
-from generalresearch.models.thl.contest.leaderboard import (
- LeaderboardContestCreate,
+from generalresearch.models.thl.contest import (
+ ContestEndCondition,
+ ContestPrize,
)
-from generalresearch.models.thl.contest.milestone import (
- MilestoneContestCreate,
+from generalresearch.models.thl.contest.definitions import (
+ ContestPrizeKind,
+ ContestType,
)
from generalresearch.models.thl.contest.raffle import (
+ ContestEntryType,
RaffleContestCreate,
)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.contest_manager import ContestManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.contest.contest import Contest
+ from generalresearch.models.thl.contest.leaderboard import (
+ LeaderboardContestCreate,
+ )
+ from generalresearch.models.thl.contest.milestone import (
+ MilestoneContestCreate,
+ )
+ from generalresearch.models.thl.contest.raffle import (
+ RaffleContest,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
# === Miscellaneous ===
# === Managers ===
+# --- Factories ---
+
+
+@pytest.fixture(scope="function")
+def raffle_contest_factory(
+ product_user_wallet_yes: Product,
+ raffle_contest_create: RaffleContestCreate,
+ contest_manager: ContestManager,
+) -> Callable[..., RaffleContest]:
+
+ def _inner(**kwargs):
+ raffle_contest_create.update(**kwargs)
+ return contest_manager.create(
+ product_id=product_user_wallet_yes.uuid,
+ contest_create=raffle_contest_create,
+ )
+
+ return _inner
+
+
# === Models ===
@pytest.fixture
def raffle_contest_create() -> RaffleContestCreate:
- from generalresearch.models.thl.contest import (
- ContestEndCondition,
- ContestPrize,
- )
- from generalresearch.models.thl.contest.definitions import (
- ContestPrizeKind,
- ContestType,
- )
- from generalresearch.models.thl.contest.raffle import (
- ContestEntryType,
- RaffleContestCreate,
- )
# This is what we'll get from the fastapi endpoint
return RaffleContestCreate(
@@ -84,23 +107,6 @@ def raffle_contest(
)
-@pytest.fixture(scope="function")
-def raffle_contest_factory(
- product_user_wallet_yes: Product,
- raffle_contest_create: RaffleContestCreate,
- contest_manager: ContestManager,
-) -> Callable[..., Contest]:
-
- def _inner(**kwargs):
- raffle_contest_create.update(**kwargs)
- return contest_manager.create(
- product_id=product_user_wallet_yes.uuid,
- contest_create=raffle_contest_create,
- )
-
- return _inner
-
-
@pytest.fixture
def milestone_contest_create() -> MilestoneContestCreate:
from generalresearch.models.thl.contest import (
@@ -135,7 +141,7 @@ def milestone_contest_create() -> MilestoneContestCreate:
),
],
end_condition=MilestoneContestEndCondition(
- ends_at=datetime(year=2030, month=1, day=1, tzinfo=timezone.utc),
+ ends_at=datetime(year=2030, month=1, day=1, tzinfo=UTC),
max_winners=5,
),
entry_trigger=ContestEntryTrigger.TASK_COMPLETE,
@@ -269,24 +275,26 @@ def user_with_money(
request: Request,
user_factory: Callable[..., User],
product_user_wallet_yes: Product,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
) -> User:
params = getattr(request, "param", {}) or {}
min_balance = int(params.get("min_balance", USDCent(1_00)))
user: User = user_factory(product=product_user_wallet_yes)
- wallet = thl_lm.get_account_or_create_user_wallet(user)
- balance = thl_lm.get_account_balance(wallet)
+ wallet = thl_ledger_manager.get_account_or_create_user_wallet(user)
+ balance = thl_ledger_manager.get_account_balance(wallet)
todo = min_balance - balance
if todo > 0:
# # Put money in user's wallet
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
ref_uuid=uuid4().hex,
description="bonus",
amount=Decimal(todo) / 100,
)
- print(f"wallet balance: {thl_lm.get_user_wallet_balance(user=user)}")
+ print(
+ f"wallet balance: {thl_ledger_manager.get_user_wallet_balance(user=user)}"
+ )
return user
diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py
index df97306..f5dcaa1 100644
--- a/test_utils/models/gr/conftest.py
+++ b/test_utils/models/gr/conftest.py
@@ -1,31 +1,34 @@
from __future__ import annotations
-from typing import Callable
+from collections.abc import Callable
+from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from pydantic import PositiveInt
from pydantic_extra_types.phone_numbers import PhoneNumber
-from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
-from generalresearch.managers.gr.business import (
- BusinessAddressManager,
- BusinessBankAccountManager,
- BusinessManager,
-)
-from generalresearch.managers.gr.team import MembershipManager, TeamManager
from generalresearch.models.custom_types import UUIDStr
-from generalresearch.models.gr.authentication import GRToken, GRUser
-from generalresearch.models.gr.business import (
- Business,
- BusinessAddress,
- BusinessBankAccount,
- BusinessType,
- TransferMethod,
-)
-from generalresearch.models.gr.team import Membership, Team
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
+from generalresearch.models.gr.definitions import TransferMethod
+
+if TYPE_CHECKING:
+ from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
+ from generalresearch.managers.gr.business import (
+ BusinessAddressManager,
+ BusinessBankAccountManager,
+ BusinessManager,
+ )
+ from generalresearch.managers.gr.team import MembershipManager, TeamManager
+ from generalresearch.models.gr.authentication import GRToken, GRUser
+ from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+ BusinessBankAccount,
+ )
+ from generalresearch.models.gr.team import Membership, Team
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
# --- Static ---
@@ -33,68 +36,57 @@ from generalresearch.redis_helper import RedisConfig
# --- Factory / Database ---
-@pytest.fixture
-def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]:
-
- def _inner(
- sub: str | None = None,
- is_superuser: bool = False,
- ) -> GRUser:
- sub = sub or f"{uuid4().hex}-{uuid4().hex}"
-
- return gr_user_manager.create(
- sub=sub,
- is_superuser=is_superuser,
- )
-
- return _inner
-
-
-@pytest.fixture
-def gr_user_cache(
- gr_user: GRUser,
- gr_db: PostgresConfig,
- thl_web_rr: PostgresConfig,
- gr_redis_config: RedisConfig,
-) -> GRUser:
- gr_user.set_cache(
- pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
- )
- return gr_user
+# --- Business Bank Account ---
@pytest.fixture
def gr_business_bank_account_factory(
- gr_bbam: BusinessBankAccountManager,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
) -> Callable[..., BusinessBankAccount]:
def _inner(
business_id: PositiveInt,
+ save: bool = True,
uuid: UUIDStr | None = None,
transfer_method: TransferMethod | None = None,
account_number: str | None = None,
routing_number: str | None = None,
iban: str | None = None,
swift: str | None = None,
- ):
- from generalresearch.models.gr.business import TransferMethod
-
- return gr_bbam.create(
- business_id=business_id,
- uuid=uuid or uuid4().hex,
- transfer_method=transfer_method or TransferMethod.ACH,
- account_number=account_number or uuid4().hex[:6],
- routing_number=routing_number or uuid4().hex[:6],
- iban=iban or uuid4().hex[:6],
- swift=swift or uuid4().hex[:6],
- )
+ **kwargs,
+ ) -> BusinessBankAccount:
+
+ if save:
+ return gr_business_bank_account_manager.create(
+ business_id=business_id,
+ uuid=uuid or uuid4().hex,
+ transfer_method=transfer_method or TransferMethod.ACH,
+ account_number=account_number or uuid4().hex[:6],
+ routing_number=routing_number or uuid4().hex[:6],
+ iban=iban or uuid4().hex[:6],
+ swift=swift or uuid4().hex[:6],
+ **kwargs,
+ )
+ else:
+ raise ValueError("Unsaved BusinessBankAccount not supported yet")
return _inner
@pytest.fixture
+def gr_business_bank_account(
+ gr_business_bank_account_factory: Callable[..., BusinessBankAccount],
+ gr_business: Business,
+) -> BusinessBankAccount:
+ return gr_business_bank_account_factory(save=True, business_id=gr_business.id)
+
+
+# --- Business Address ---
+
+
+@pytest.fixture
def gr_business_address_factory(
- gr_bam: BusinessAddressManager,
+ gr_business_address_manager: BusinessAddressManager,
) -> Callable[..., BusinessAddress]:
def _inner(
@@ -107,7 +99,7 @@ def gr_business_address_factory(
postal_code: str | None = None,
phone_number: PhoneNumber | None = None,
country: str | None = None,
- ):
+ ) -> BusinessAddress:
uuid = uuid or uuid4().hex
line_1 = line_1 or "abc"
line_2 = line_2 or "bczx"
@@ -117,7 +109,7 @@ def gr_business_address_factory(
phone_number = None
country = country or "US"
- return gr_bam.create(
+ return gr_business_address_manager.create(
business_id=business_id,
uuid=uuid,
line_1=line_1,
@@ -133,54 +125,175 @@ def gr_business_address_factory(
@pytest.fixture
+def gr_business_address(
+ gr_business_address_factory: Callable[..., BusinessAddress], gr_business: Business
+) -> BusinessAddress:
+ return gr_business_address_factory(business_id=gr_business.id)
+
+
+# --- Business ---
+
+
+@pytest.fixture
def gr_business_factory(
- gr_bm: BusinessManager,
+ gr_business_manager: BusinessManager,
) -> Callable[..., Business]:
def _inner(
+ save: bool = True, name: str | None = None, team: Team | None = None, **kwargs
+ ) -> Business:
+ name = name or f"<Unknown {uuid4().hex[:12]}>"
+ tax_number = str(randint(1, 999_999_999))
+
+ if save:
+ return gr_business_manager.create(
+ name=name,
+ kind="c",
+ uuid=uuid4().hex,
+ team=team,
+ tax_number=tax_number,
+ **kwargs,
+ )
+ else:
+ raise ValueError("Unsaved Business not supported yet")
+
+ return _inner
+
+
+@pytest.fixture
+def gr_business(gr_business_factory: Callable[..., Business]) -> Business:
+ return gr_business_factory(save=True)
+
+
+@pytest.fixture
+def unsaved_gr_business(gr_business_factory: Callable[..., Business]) -> Business:
+ return gr_business_factory(save=False)
+
+
+# --- GR Team ---
+
+
+@pytest.fixture
+def gr_team_factory(
+ gr_team_manager: TeamManager,
+) -> Callable[..., Team]:
+
+ def _inner(
+ save: bool = True,
uuid: UUIDStr | None = None,
name: str | None = None,
- team: Team | None = None,
- kind: BusinessType | None = None,
- tax_number: str | None = None,
- ) -> Business:
- from random import randint
+ **kwargs,
+ ) -> Team:
- uuid = uuid or uuid4().hex
- name = name or "< Unknown >"
- tax_number = tax_number or str(randint(1, 999_999_999))
+ name = name or f"<Team ({uuid4().hex[:6]})>"
- return gr_bm.create(
- uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number
- )
+ if save:
+ return gr_team_manager.create(name=name, uuid=uuid, **kwargs)
+
+ else:
+ raise ValueError("BusinessBankAccount Business not supported yet")
return _inner
@pytest.fixture
-def gr_team(
- gr_tm: TeamManager,
-) -> Callable[..., Team]:
+def gr_team(gr_team_factory: Callable[..., Team]) -> Team:
+ return gr_team_factory(save=True)
- def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team:
- uuid = uuid or uuid4().hex
- name = name or f"name-{uuid4().hex[:12]}"
- return gr_tm.create(uuid=uuid, name=name)
+@pytest.fixture
+def unsaved_gr_team(
+ gr_team_factory: Callable[..., Team],
+) -> Team:
+ return gr_team_factory(save=False)
+
+
+# --- GR User ---
+
+
+@pytest.fixture
+def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]:
+
+ def _inner(
+ save: bool = True,
+ sub: str | None = None,
+ is_superuser: bool = False,
+ ) -> GRUser:
+ sub = sub or f"{uuid4().hex}-{uuid4().hex}"
+
+ if save:
+ return gr_user_manager.create(
+ sub=sub,
+ is_superuser=is_superuser,
+ )
+ else:
+ raise ValueError("Unsaved GR User not supported yet")
return _inner
-@pytest.fixture()
-def gr_user_token(
- gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig
-) -> GRToken:
- gr_tm.create(user_id=gr_user.id)
- gr_user.prefetch_token(pg_config=gr_db)
+@pytest.fixture
+def gr_user_cache(
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
+) -> GRUser:
+ gr_user.set_cache(
+ pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
+ )
+ return gr_user
- res = gr_user.token
- assert res is not None, "GRToken should exist after creation and prefetching"
- return res
+
+@pytest.fixture
+def gr_user(gr_user_factory: Callable[..., GRUser]) -> GRUser:
+ return gr_user_factory(save=True)
+
+
+@pytest.fixture
+def unsaved_gr_user(
+ gr_user_factory: Callable[..., GRUser],
+) -> GRUser:
+ return gr_user_factory(save=False)
+
+
+# --- GR User Token ---
+
+
+@pytest.fixture
+def gr_user_token_factory(
+ gr_user: GRUser, gr_token_manager: GRTokenManager, gr_db: PostgresConfig
+) -> Callable[..., GRToken]:
+
+ def _inner(
+ save: bool = True,
+ ) -> GRToken:
+
+ if save:
+ assert gr_user.id
+ gr_token_manager.create(user_id=gr_user.id)
+ gr_user.prefetch_token(pg_config=gr_db)
+
+ res = gr_user.token
+ assert res is not None, (
+ "GRToken should exist after creation and prefetching"
+ )
+ return res
+
+ else:
+ raise ValueError("Unsaved GR User not supported yet")
+
+ return _inner
+
+
+@pytest.fixture
+def gr_user_token(gr_user_token_factory: Callable[..., GRToken]) -> GRToken:
+ return gr_user_token_factory(save=True)
+
+
+@pytest.fixture
+def unsaved_gr_user_token(gr_user_token_factory: Callable[..., GRToken]) -> GRToken:
+ return gr_user_token_factory(save=False)
@pytest.fixture()
@@ -188,26 +301,34 @@ def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]:
return gr_user_token.auth_header
-@pytest.fixture(scope="function")
-def membership(team: Team, gr_user: GRUser, team_manager: TeamManager) -> Membership:
- assert team.id, "Team must be saved"
- assert gr_user.id, "GRUser must be saved"
- return team_manager.add_user(team=team, gr_user=gr_user)
+# --- GR Membership ---
-@pytest.fixture(scope="function")
-def membership_factory(
- team: Team,
- gr_user: GRUser,
- membership_manager: MembershipManager,
- team_manager: TeamManager,
- gr_um: GRUserManager,
+@pytest.fixture()
+def gr_membership_factory(
+ gr_membership_manager: MembershipManager,
) -> Callable[..., Membership]:
- def _inner(**kwargs) -> Membership:
- _team = kwargs.get("team", team_manager.create_dummy())
- _gr_user = kwargs.get("gr_user", gr_um.create_dummy())
-
- return membership_manager.create(team=_team, gr_user=_gr_user)
+ def _inner(
+ gr_team: Team, gr_user: GRUser, save: bool = True, **kwargs
+ ) -> Membership:
+ if save:
+ return gr_membership_manager.create(team=gr_team, gr_user=gr_user, **kwargs)
+ else:
+ raise ValueError("Unsaved GR Membership not supported yet")
return _inner
+
+
+@pytest.fixture()
+def gr_membership(
+ gr_membership_factory: Callable[..., Membership], gr_team: Team, gr_user: GRUser
+) -> Membership:
+ return gr_membership_factory(gr_team=gr_team, gr_user=gr_user, save=True)
+
+
+@pytest.fixture()
+def unsaved_gr_membership(
+ gr_membership_factory: Callable[..., Membership],
+) -> Membership:
+ return gr_membership_factory(save=False)
diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py
index 5bef113..9ee0df2 100644
--- a/test_utils/models/ledger/conftest.py
+++ b/test_utils/models/ledger/conftest.py
@@ -1,43 +1,39 @@
from __future__ import annotations
+from collections.abc import Callable
from datetime import datetime
from decimal import Decimal
from random import randint
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from fastapi import Request
+from pytest import FixtureRequest as Request
from generalresearch.currency import USDCent
-from generalresearch.managers.base import PostgresManager
-from test_utils.models.conftest import (
- payout_config,
- product_amt_true,
- product_user_wallet_no,
- product_user_wallet_yes,
- session,
- session_factory,
- user_factory,
- wall,
- wall_factory,
-)
-
-_ = (
- user_factory,
- product_user_wallet_no,
- wall,
- product_amt_true,
- product_user_wallet_yes,
- session_factory,
- session,
- wall_factory,
- payout_config,
-)
-if TYPE_CHECKING:
+# from test_utils.models.conftest import (
+# payout_config,
+# product_amt_true,
+# product_user_wallet_no,
+# product_user_wallet_yes,
+# )
+
+# _ = (
+# user_factory,
+# product_user_wallet_no,
+# wall,
+# product_amt_true,
+# product_user_wallet_yes,
+# session_factory,
+# session,
+# wall_factory,
+# payout_config,
+# )
+if TYPE_CHECKING:
from generalresearch.currency import LedgerCurrency
+ from generalresearch.managers.base import PostgresManager
from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
@@ -62,7 +58,7 @@ if TYPE_CHECKING:
@pytest.fixture
def ledger_account(
- request: Request, lm: LedgerManager, currency: LedgerCurrency
+ request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
@@ -84,14 +80,14 @@ def ledger_account(
account_type=account_type,
normal_balance=direction,
)
- return lm.create_account(account=acct_model)
+ return ledger_manager.create_account(account=acct_model)
@pytest.fixture
def ledger_account_factory(
request: Request,
- thl_lm: ThlLedgerManager,
- lm: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
currency: LedgerCurrency,
) -> Callable[..., LedgerAccount]:
@@ -106,7 +102,7 @@ def ledger_account_factory(
account_type: AccountType = AccountType.CASH,
direction: Direction = Direction.CREDIT,
) -> LedgerAccount:
- thl_lm.get_account_or_create_bp_wallet(product=product)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
acct_uuid = uuid4().hex
qn = f"{currency}:{account_type}:{acct_uuid}"
@@ -118,14 +114,14 @@ def ledger_account_factory(
account_type=account_type,
normal_balance=direction,
)
- return lm.create_account(account=acct_model)
+ return ledger_manager.create_account(account=acct_model)
return _inner
@pytest.fixture
def ledger_account_credit(
- request: Request, lm: LedgerManager, currency: LedgerCurrency
+ request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import AccountType, Direction
@@ -143,12 +139,12 @@ def ledger_account_credit(
account_type=account_type,
normal_balance=Direction.CREDIT,
)
- return lm.create_account(account=acct_model)
+ return ledger_manager.create_account(account=acct_model)
@pytest.fixture
def ledger_account_debit(
- request: Request, lm: LedgerManager, currency: LedgerCurrency
+ request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import AccountType, Direction
@@ -166,11 +162,11 @@ def ledger_account_debit(
account_type=account_type,
normal_balance=Direction.DEBIT,
)
- return lm.create_account(account=acct_model)
+ return ledger_manager.create_account(account=acct_model)
@pytest.fixture
-def tag(request: Request, lm: LedgerManager) -> str:
+def tag(request: Request) -> str:
from generalresearch.currency import LedgerCurrency
return (
@@ -190,23 +186,24 @@ def usd_cent(request: Request) -> USDCent:
def bp_payout_event(
product: Product,
usd_cent: USDCent,
- business_payout_event_manager: BusinessPayoutEventManager,
- thl_lm: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEvent,
+ thl_ledger_manager: ThlLedgerManager,
) -> BrokerageProductPayoutEvent:
- return business_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ _ext_ref_id = f"tx-{uuid4().hex[:7]}"
+
+ return brokerage_product_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ ext_ref_id=_ext_ref_id,
product=product,
amount=usd_cent,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
@pytest.fixture
def bp_payout_event_factory(
brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
) -> Callable[..., BrokerageProductPayoutEvent]:
def _inner(
@@ -214,7 +211,7 @@ def bp_payout_event_factory(
) -> BrokerageProductPayoutEvent:
return brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ thl_ledger_manager=thl_ledger_manager,
product=product,
amount=usd_cent,
ext_ref_id=ext_ref_id,
@@ -226,10 +223,12 @@ def bp_payout_event_factory(
@pytest.fixture
-def currency(lm: LedgerManager) -> LedgerCurrency:
+def currency(ledger_manager: LedgerManager) -> LedgerCurrency:
# return request.param if hasattr(request, "currency") else LedgerCurrency.TEST
- assert lm.currency, "LedgerManager must have a currency specified for these tests"
- return lm.currency
+ assert (
+ ledger_manager.currency
+ ), "LedgerManager must have a currency specified for these tests"
+ return ledger_manager.currency
@pytest.fixture
@@ -249,7 +248,7 @@ def ledger_tx(
tag: str,
currency: LedgerCurrency,
tx_metadata: dict[str, str] | None,
- lm: LedgerManager,
+ ledger_manager: LedgerManager,
) -> LedgerTransaction:
from generalresearch.models.thl.ledger import Direction, LedgerEntry
@@ -268,12 +267,12 @@ def ledger_tx(
),
]
- return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata)
+ return ledger_manager.create_tx(entries=entries, tag=tag, metadata=tx_metadata)
@pytest.fixture
def create_main_accounts(
- lm: LedgerManager, currency: LedgerCurrency
+ ledger_manager: LedgerManager, currency: LedgerCurrency
) -> Callable[..., None]:
def _inner() -> None:
@@ -288,9 +287,9 @@ def create_main_accounts(
qualified_name=f"{currency.value}:revenue:task_complete",
normal_balance=Direction.CREDIT,
account_type=AccountType.REVENUE,
- currency=lm.currency,
+ currency=ledger_manager.currency,
)
- lm.get_account_or_create(account=account)
+ ledger_manager.get_account_or_create(account=account)
account = LedgerAccount(
display_name="Operating Cash Account",
@@ -300,7 +299,7 @@ def create_main_accounts(
currency=currency,
)
- lm.get_account_or_create(account=account)
+ ledger_manager.get_account_or_create(account=account)
return _inner
@@ -324,7 +323,7 @@ def delete_ledger_db(thl_web_rw: PostgresManager) -> Callable[..., None]:
@pytest.fixture
def wipe_main_accounts(
- thl_web_rw: PostgresManager, lm: LedgerManager, currency: LedgerCurrency
+ thl_web_rw: PostgresManager, ledger_manager: LedgerManager, currency: LedgerCurrency
) -> Callable[..., None]:
def _inner() -> None:
@@ -394,7 +393,9 @@ def wipe_main_accounts(
@pytest.fixture
-def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount:
+def account_cash(
+ ledger_manager: LedgerManager, currency: LedgerCurrency
+) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
Direction,
@@ -408,12 +409,12 @@ def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount:
account_type=AccountType.CASH,
currency=currency,
)
- return lm.get_account_or_create(account=account)
+ return ledger_manager.get_account_or_create(account=account)
@pytest.fixture
def account_revenue_task_complete(
- lm: LedgerManager, currency: LedgerCurrency
+ ledger_manager: LedgerManager, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
@@ -428,11 +429,13 @@ def account_revenue_task_complete(
account_type=AccountType.REVENUE,
currency=currency,
)
- return lm.get_account_or_create(account=account)
+ return ledger_manager.get_account_or_create(account=account)
@pytest.fixture
-def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount:
+def account_expense_tango(
+ ledger_manager: LedgerManager, currency: LedgerCurrency
+) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
Direction,
@@ -446,12 +449,12 @@ def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> Ledger
account_type=AccountType.EXPENSE,
currency=currency,
)
- return lm.get_account_or_create(account=account)
+ return ledger_manager.get_account_or_create(account=account)
@pytest.fixture
def user_account_user_wallet(
- lm: LedgerManager, user: User, currency: LedgerCurrency
+ ledger_manager: LedgerManager, user: User, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
@@ -468,12 +471,12 @@ def user_account_user_wallet(
reference_uuid=user.uuid,
currency=currency,
)
- return lm.get_account_or_create(account=account)
+ return ledger_manager.get_account_or_create(account=account)
@pytest.fixture
def product_account_bp_wallet(
- lm: LedgerManager, product: Product, currency: LedgerCurrency
+ ledger_manager: LedgerManager, product: Product, currency: LedgerCurrency
) -> LedgerAccount:
from generalresearch.models.thl.ledger import (
AccountType,
@@ -492,83 +495,86 @@ def product_account_bp_wallet(
"currency": currency,
}
)
- return lm.get_account_or_create(account=account)
+ return ledger_manager.get_account_or_create(account=account)
@pytest.fixture
def setup_accounts(
product_factory: Callable[..., Product],
- lm: LedgerManager,
+ ledger_manager: LedgerManager,
user: User,
currency: LedgerCurrency,
-) -> None:
+) -> Callable[..., None]:
from generalresearch.models.thl.ledger import (
AccountType,
Direction,
LedgerAccount,
)
- # BP's wallet and a revenue from their commissions account.
- p1 = product_factory()
+ def _inner():
+ # BP's wallet and a revenue from their commissions account.
+ p1 = product_factory()
- account = LedgerAccount(
- display_name=f"Revenue from {p1.name} commission",
- qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}",
- normal_balance=Direction.CREDIT,
- account_type=AccountType.REVENUE,
- reference_type="bp",
- reference_uuid=p1.uuid,
- currency=currency,
- )
- lm.get_account_or_create(account=account)
+ account = LedgerAccount(
+ display_name=f"Revenue from {p1.name} commission",
+ qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}",
+ normal_balance=Direction.CREDIT,
+ account_type=AccountType.REVENUE,
+ reference_type="bp",
+ reference_uuid=p1.uuid,
+ currency=currency,
+ )
+ ledger_manager.get_account_or_create(account=account)
+
+ account = LedgerAccount.model_validate(
+ {
+ "display_name": f"{p1.name} Wallet",
+ "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}",
+ "normal_balance": Direction.CREDIT,
+ "account_type": AccountType.BP_WALLET,
+ "reference_type": "bp",
+ "reference_uuid": p1.uuid,
+ "currency": currency,
+ }
+ )
+ ledger_manager.get_account_or_create(account=account)
- account = LedgerAccount.model_validate(
- {
- "display_name": f"{p1.name} Wallet",
- "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}",
- "normal_balance": Direction.CREDIT,
- "account_type": AccountType.BP_WALLET,
- "reference_type": "bp",
- "reference_uuid": p1.uuid,
- "currency": currency,
- }
- )
- lm.get_account_or_create(account=account)
+ # BP's wallet, user's wallet, and a revenue from their commissions account.
+ p2 = product_factory()
+ account = LedgerAccount(
+ display_name=f"Revenue from {p2.name} commission",
+ qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}",
+ normal_balance=Direction.CREDIT,
+ account_type=AccountType.REVENUE,
+ reference_type="bp",
+ reference_uuid=p2.uuid,
+ currency=currency,
+ )
+ ledger_manager.get_account_or_create(account)
- # BP's wallet, user's wallet, and a revenue from their commissions account.
- p2 = product_factory()
- account = LedgerAccount(
- display_name=f"Revenue from {p2.name} commission",
- qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}",
- normal_balance=Direction.CREDIT,
- account_type=AccountType.REVENUE,
- reference_type="bp",
- reference_uuid=p2.uuid,
- currency=currency,
- )
- lm.get_account_or_create(account)
+ account = LedgerAccount(
+ display_name=f"{p2.name} Wallet",
+ qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}",
+ normal_balance=Direction.CREDIT,
+ account_type=AccountType.BP_WALLET,
+ reference_type="bp",
+ reference_uuid=p2.uuid,
+ currency=currency,
+ )
+ ledger_manager.get_account_or_create(account)
- account = LedgerAccount(
- display_name=f"{p2.name} Wallet",
- qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}",
- normal_balance=Direction.CREDIT,
- account_type=AccountType.BP_WALLET,
- reference_type="bp",
- reference_uuid=p2.uuid,
- currency=currency,
- )
- lm.get_account_or_create(account)
+ account = LedgerAccount(
+ display_name=f"{user.uuid} Wallet",
+ qualified_name=f"{currency.value}:user_wallet:{user.uuid}",
+ normal_balance=Direction.CREDIT,
+ account_type=AccountType.USER_WALLET,
+ reference_type="user",
+ reference_uuid=user.uuid,
+ currency="test",
+ )
+ ledger_manager.get_account_or_create(account=account)
- account = LedgerAccount(
- display_name=f"{user.uuid} Wallet",
- qualified_name=f"{currency.value}:user_wallet:{user.uuid}",
- normal_balance=Direction.CREDIT,
- account_type=AccountType.USER_WALLET,
- reference_type="user",
- reference_uuid=user.uuid,
- currency="test",
- )
- lm.get_account_or_create(account=account)
+ return _inner
@pytest.fixture
@@ -577,7 +583,7 @@ def session_with_tx_factory(
session_manager: SessionManager,
wall_manager: WallManager,
utc_hour_ago: datetime,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
) -> Callable[..., Session]:
from generalresearch.models.thl.session import (
@@ -618,14 +624,16 @@ def session_with_tx_factory(
status_code_1=status_code_1,
)
- thl_lm.create_tx_task_complete(
+ thl_ledger_manager.create_tx_task_complete(
wall=last_wall,
user=user,
created=last_wall.finished,
force=True,
)
- thl_lm.create_tx_bp_payment(session=s, created=last_wall.finished, force=True)
+ thl_ledger_manager.create_tx_bp_payment(
+ session=s, created=last_wall.finished, force=True
+ )
return s
@@ -636,7 +644,7 @@ def session_with_tx_factory(
def adj_to_fail_with_tx_factory(
session_manager: SessionManager,
wall_manager: WallManager,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
) -> Callable[..., None]:
from datetime import timedelta
@@ -669,7 +677,7 @@ def adj_to_fail_with_tx_factory(
adjusted_timestamp=created,
)
- thl_lm.create_tx_task_adjustment(
+ thl_ledger_manager.create_tx_task_adjustment(
wall=w1,
user=session.user,
created=created + timedelta(milliseconds=1),
@@ -678,7 +686,7 @@ def adj_to_fail_with_tx_factory(
session.wall_events = wall_manager.get_wall_events(session_id=session.id)
session_manager.adjust_status(session=session)
- thl_lm.create_tx_bp_adjustment(
+ thl_ledger_manager.create_tx_bp_adjustment(
session=session, created=created + timedelta(milliseconds=2)
)
@@ -689,7 +697,7 @@ def adj_to_fail_with_tx_factory(
def adj_to_complete_with_tx_factory(
session_manager: SessionManager,
wall_manager: WallManager,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
) -> Callable[..., None]:
from datetime import timedelta
@@ -708,7 +716,7 @@ def adj_to_complete_with_tx_factory(
adjusted_timestamp=created,
)
- thl_lm.create_tx_task_adjustment(
+ thl_ledger_manager.create_tx_task_adjustment(
wall=w1,
user=session.user,
created=created + timedelta(milliseconds=1),
@@ -717,7 +725,7 @@ def adj_to_complete_with_tx_factory(
session.wall_events = wall_manager.get_wall_events(session_id=session.id)
session_manager.adjust_status(session=session)
- thl_lm.create_tx_bp_adjustment(
+ thl_ledger_manager.create_tx_bp_adjustment(
session=session, created=created + timedelta(milliseconds=2)
)
diff --git a/test_utils/models/network/__init__.py b/test_utils/models/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/test_utils/models/network/__init__.py
+++ /dev/null
diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py
deleted file mode 100644
index abfbc18..0000000
--- a/test_utils/models/network/conftest.py
+++ /dev/null
@@ -1,144 +0,0 @@
-import os
-from datetime import datetime, timedelta, timezone
-from uuid import uuid4
-
-import pytest
-from fastapi import Request
-
-from generalresearch.managers.network.label import IPLabelManager
-from generalresearch.managers.network.tool_run import ToolRunManager
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.mtr.parser import parse_mtr_output
-from generalresearch.models.network.mtr.result import MTRResult
-from generalresearch.models.network.nmap.parser import parse_nmap_xml
-from generalresearch.models.network.nmap.result import NmapResult
-from generalresearch.models.network.rdns.parser import parse_rdns_output
-from generalresearch.models.network.rdns.result import RDNSResult
-from generalresearch.models.network.tool_run import MTRRun, NmapRun, RDNSRun, Status
-from generalresearch.models.network.tool_run_command import (
- MTRRunCommand,
- MTRRunCommandOptions,
- NmapRunCommand,
- NmapRunCommandOptions,
- RDNSRunCommand,
- RDNSRunCommandOptions,
-)
-from generalresearch.pg_helper import PostgresConfig
-
-
-@pytest.fixture(scope="session")
-def scan_group_id() -> str:
- return uuid4().hex
-
-
-@pytest.fixture(scope="session")
-def iplabel_manager(thl_web_rw: PostgresConfig) -> IPLabelManager:
- return IPLabelManager(pg_config=thl_web_rw)
-
-
-@pytest.fixture(scope="session")
-def toolrun_manager(thl_web_rw: PostgresConfig) -> ToolRunManager:
- return ToolRunManager(pg_config=thl_web_rw)
-
-
-@pytest.fixture(scope="session")
-def nmap_raw_output(request: Request) -> str:
- fp = os.path.join(request.config.rootpath, "data/nmaprun1.xml")
- with open(fp) as f:
- data = f.read()
- return data
-
-
-@pytest.fixture(scope="session")
-def nmap_result(nmap_raw_output: str) -> NmapResult:
- return parse_nmap_xml(nmap_raw_output)
-
-
-@pytest.fixture(scope="session")
-def nmap_run(nmap_result: NmapResult, scan_group_id: str):
- r = nmap_result
- config = NmapRunCommand(
- command="nmap",
- options=NmapRunCommandOptions(
- ip=r.target_ip, ports="22-1000,11000,1100,3389,61232", top_ports=None
- ),
- )
- return NmapRun(
- tool_version=r.version,
- status=Status.SUCCESS,
- ip=r.target_ip,
- started_at=r.started_at,
- finished_at=r.finished_at,
- raw_command=config.to_command_str(),
- scan_group_id=scan_group_id,
- config=config,
- parsed=r,
- )
-
-
-@pytest.fixture(scope="session")
-def dig_raw_output() -> str:
- return "156.32.33.45.in-addr.arpa. 300 IN PTR scanme.nmap.org."
-
-
-@pytest.fixture(scope="session")
-def rdns_result(dig_raw_output: str) -> RDNSResult:
- return parse_rdns_output(ip="45.33.32.156", raw=dig_raw_output)
-
-
-@pytest.fixture(scope="session")
-def rdns_run(rdns_result: RDNSResult, scan_group_id: str):
- r = rdns_result
- ip = "45.33.32.156"
- utc_now = datetime.now(tz=timezone.utc)
- config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip))
- return RDNSRun(
- tool_version="1.2.3",
- status=Status.SUCCESS,
- ip=ip,
- started_at=utc_now,
- finished_at=utc_now + timedelta(seconds=1),
- raw_command=config.to_command_str(),
- scan_group_id=scan_group_id,
- config=config,
- parsed=r,
- )
-
-
-@pytest.fixture(scope="session")
-def mtr_raw_output(request: Request) -> str:
- fp = os.path.join(request.config.rootpath, "data/mtr_fatbeam.json")
- with open(fp) as f:
- data = f.read()
- return data
-
-
-@pytest.fixture(scope="session")
-def mtr_result(mtr_raw_output: str) -> MTRResult:
- return parse_mtr_output(mtr_raw_output, port=443, protocol=IPProtocol.TCP)
-
-
-@pytest.fixture(scope="session")
-def mtr_run(mtr_result: MTRResult, scan_group_id: str):
- r = mtr_result
- utc_now = datetime.now(tz=timezone.utc)
- config = MTRRunCommand(
- command="mtr",
- options=MTRRunCommandOptions(
- ip=r.destination, protocol=IPProtocol.TCP, port=443
- ),
- )
-
- return MTRRun(
- tool_version="1.2.3",
- status=Status.SUCCESS,
- ip=r.destination,
- started_at=utc_now,
- finished_at=utc_now + timedelta(seconds=1),
- raw_command=config.to_command_str(),
- scan_group_id=scan_group_id,
- config=config,
- parsed=r,
- facility_id=1,
- source_ip="1.2.3.4",
- )
diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py
index cf8d2fa..e09eadd 100644
--- a/test_utils/models/thl/conftest.py
+++ b/test_utils/models/thl/conftest.py
@@ -1,149 +1,386 @@
from __future__ import annotations
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import ROUND_DOWN, Decimal
from random import choice as rand_choice
-from random import choice as rchoice
from random import randint, random
-from typing import Any, Callable
+from typing import TYPE_CHECKING, Any
from uuid import uuid4
import faker
import pytest
+from grip_client.enums import AccessType
from pydantic import PositiveInt
-from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager
-from generalresearch.managers.thl.payout import UserPayoutEventManager
-from generalresearch.managers.thl.product import ProductManager
-from generalresearch.managers.thl.session import SessionManager
-from generalresearch.managers.thl.user_manager.user_manager import UserManager
-from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager
-from generalresearch.managers.thl.wall import WallManager
-from generalresearch.models import DeviceType
+from generalresearch.currency import USDCent
+from generalresearch.managers.thl.payout import (
+ BusinessPayoutEventManager,
+ UserPayoutEventManager,
+)
from generalresearch.models.custom_types import (
AwareDatetimeISO,
IPvAnyAddressStr,
UUIDStr,
)
-from generalresearch.models.legacy.bucket import Bucket
+from generalresearch.models.definitions import DeviceType, Source
from generalresearch.models.thl.definitions import (
+ WALL_ALLOWED_STATUS_STATUS_CODE,
PayoutStatus,
-)
-from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.product import (
- PayoutConfig,
- Product,
- ProfilingConfig,
- SessionConfig,
- SourcesConfig,
- SupplyConfig,
- UserCreateConfig,
- UserHealthConfig,
- UserWalletConfig,
-)
-from generalresearch.models.thl.session import (
- Session,
- Source,
Status,
- Wall,
)
+from generalresearch.models.thl.payout import UserPayoutEvent
from generalresearch.models.thl.user import User
-from generalresearch.models.thl.user_iphistory import IPRecord
-from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData
+from generalresearch.models.thl.userhealth import AuditLogLevel
+from generalresearch.models.thl.wallet.definitions import PayoutType
+from generalresearch.pg_helper import PostgresConfig
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ipinfo import (
+ IPGeonameManager,
+ IPInformationManager,
+ )
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
+ from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.custom_types import AwareDatetime
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.gr.team import Team
+ from generalresearch.models.legacy.bucket import Bucket
+ from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation
+ from generalresearch.models.thl.payout import BrokerageProductPayoutEvent
+ from generalresearch.models.thl.product import (
+ PayoutConfig,
+ Product,
+ ProfilingConfig,
+ SessionConfig,
+ SourcesConfig,
+ SupplyConfig,
+ UserCreateConfig,
+ UserHealthConfig,
+ UserWalletConfig,
+ )
+ from generalresearch.models.thl.session import (
+ Session,
+ Wall,
+ )
+ from generalresearch.models.thl.user_iphistory import IPRecord
+ from generalresearch.models.thl.userhealth import AuditLog
+ from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData
fake = faker.Faker()
+# --- Wall ---
+
+
+@pytest.fixture
+def wall_factory(
+ wall_manager: WallManager,
+ bare_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
+ user_factory: Callable[..., User],
+) -> Callable[..., Wall]:
+
+ def _inner(
+ wall_status: Status = Status.FAIL,
+ save: bool = True,
+ session: Session | None = None,
+ session_id: PositiveInt | None = None,
+ user: User | None = None,
+ started: datetime | None = None,
+ source: Source | None = None,
+ req_survey_id: str | None = None,
+ req_cpi: Decimal | None = None,
+ buyer_id: str | None = None,
+ uuid_id: str | None = None,
+ ) -> Wall:
+ """To be used in tests, where we don't care about certain fields"""
+ user = user or user_factory()
+ if save:
+ _wall_started = started or fake.date_time_between(
+ start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC),
+ end_date=datetime.now(tz=UTC),
+ tzinfo=UTC,
+ )
+
+ if session:
+ # If an existing Session was provided, we want to do some
+ # additional validation.
+
+ if session.wall_events:
+ # Subsequent Wall events
+ _last_wall = session.wall_events[-1]
+ assert not _last_wall.finished, (
+ "Can't add new Walls until prior finishes"
+ )
+ _wall_started = _last_wall.started + timedelta(milliseconds=1)
+ else:
+ # First Wall Event in a session
+ _wall_started = session.started + timedelta(milliseconds=1)
+ else:
+ # If a Session was NOT provided, either (1) try to retrieve it
+ # from an optionally provided session_id int, or (2) proceed
+ # forward and make one
+ session = (
+ session_manager.get_from_id(session_id=session_id)
+ if session_id
+ else None
+ ) or bare_session_factory(save=True, user=user)
+
+ assert session, "Wall factory requires Session"
+
+ source = source or rand_choice(list(Source))
+ req_survey_id = req_survey_id or uuid4().hex
+ req_cpi = req_cpi or Decimal(
+ fake.random_int(min=1, max=150) / 100
+ ).quantize(Decimal(".01"), rounding=ROUND_DOWN)
+
+ w = wall_manager.create(
+ session_id=session.id,
+ user_id=session.user_id,
+ started=_wall_started,
+ source=source,
+ req_survey_id=req_survey_id,
+ req_cpi=req_cpi,
+ buyer_id=buyer_id,
+ uuid_id=uuid_id,
+ )
+
+ _status_code_options = list(
+ WALL_ALLOWED_STATUS_STATUS_CODE.get(wall_status, {})
+ )
+ w.finish(
+ finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
+ status=wall_status,
+ status_code_1=rand_choice(_status_code_options),
+ )
+
+ session.append_wall_event(w=w)
+
+ return w
+
+ else:
+ raise ValueError("Unsaved Wall not yet supported")
+
+ return _inner
+
+
+@pytest.fixture
+def wall(wall_factory: Callable[..., Wall]) -> Wall:
+ return wall_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_wall(wall_factory: Callable[..., Wall]) -> Wall:
+ return wall_factory(save=False)
+
+
+# --- Wall: Enum(s) ---
+
@pytest.fixture
def wall_status() -> Status:
return Status.COMPLETE
+# --- Session ---
+
+
@pytest.fixture
-def user_factory(user_manager: UserManager) -> Callable[..., User]:
+def bare_session_factory(
+ session_manager: SessionManager, user_factory: Callable[..., User]
+):
+ # Create a session with no wall events
def _inner(
- # --- Create dummy "optional" --- #
- product_user_id: str | None = None,
- # --- Optional --- #
- product_id: UUIDStr | None = None,
- product: Product | None = None,
- created: datetime | None = None,
- ) -> User:
-
- product_user_id = product_user_id or uuid4().hex
+ save: bool = True,
+ # -- Create Dummy "optional" -- #
+ started: datetime | None = None,
+ user: User | None = None,
+ # -- Optional -- #
+ country_iso: str | None = None,
+ device_type: DeviceType | None = None,
+ ip: str | None = None,
+ bucket: Bucket | None = None,
+ url_metadata: dict[str, str] | None = None,
+ uuid_id: str | None = None,
+ ) -> Session:
- return user_manager.create_user(
- product_user_id=product_user_id,
- product_id=product_id,
- product=product,
- created=created,
- )
+ if save:
+ """To be used in tests, where we don't care about certain fields"""
+ started = started or fake.date_time_between(
+ start_date=datetime(year=1900, month=1, day=1, tzinfo=UTC),
+ end_date=datetime(year=2000, month=1, day=1, tzinfo=UTC),
+ tzinfo=UTC,
+ )
+ user = user or user_factory(save=True)
+ assert user.user_id, "Provided User must be saved to the database"
+
+ return session_manager.create(
+ started=started,
+ user=user,
+ country_iso=country_iso,
+ device_type=device_type,
+ ip=ip,
+ bucket=bucket,
+ url_metadata=url_metadata,
+ uuid_id=uuid_id,
+ )
+ else:
+ # user = User(
+ # user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex
+ # )
+ raise ValueError("Unsaved Session not yet supported")
return _inner
+@pytest.fixture()
+def bare_session(bare_session_factory: Callable[..., Session], user) -> Session:
+ # A session with no wall events
+ return bare_session_factory(user=user)
+
+
@pytest.fixture
-def wall_factory(
- wall_manager: WallManager, session_factory: Session
-) -> Callable[..., Wall]:
+def session(
+ bare_session: Session,
+ wall_factory: Callable[..., Wall],
+) -> Session:
+ s = bare_session.model_copy()
+ wall: Wall = wall_factory(
+ session_id=s.id,
+ user=s.user,
+ started=s.started,
+ )
+ s.append_wall_event(w=wall)
+ return s
+
+
+@pytest.fixture
+def session_factory(
+ wall_manager: WallManager,
+ utc_hour_ago: datetime,
+ bare_session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
+) -> Callable[..., Session]:
def _inner(
- session_id: int | None = None,
- user_id: int | None = None,
- started: datetime | None = None,
- source: Source | None = None,
- req_survey_id: str | None = None,
- req_cpi: Decimal | None = None,
- buyer_id: str | None = None,
- uuid_id: str | None = None,
- ):
- """To be used in tests, where we don't care about certain fields"""
+ user: User,
+ # Wall details
+ wall_count: int = 5,
+ wall_req_cpi: Decimal = Decimal(".50"),
+ wall_req_cpis: list[Decimal] | None = None,
+ wall_statuses: list[Status] | None = None,
+ wall_source: Source = Source.TESTING,
+ # Session details
+ final_status: Status = Status.COMPLETE,
+ started: datetime = utc_hour_ago,
+ ) -> Session:
+ if wall_req_cpis:
+ assert len(wall_req_cpis) == wall_count
+ if wall_statuses:
+ assert len(wall_statuses) == wall_count
+
+ s = bare_session_factory(started=started, user=user, country_iso="us")
+ for idx in range(wall_count):
+ if idx == 0:
+ # First Wall Event in a session
+ wall_started = s.started + timedelta(milliseconds=1)
+ else:
+ # Subsequent Wall events
+ last_wall = s.wall_events[-1]
+ assert last_wall.finished, "Can't add new Walls until prior finishes"
+ wall_started = last_wall.started + timedelta(milliseconds=1)
+
+ w = wall_factory(
+ session_id=s.id,
+ source=wall_source,
+ user=s.user,
+ started=wall_started,
+ req_cpi=wall_req_cpis[idx] if wall_req_cpis else wall_req_cpi,
+ )
+ s.append_wall_event(w=w)
+
+ # If it's the last wall in the session, respect the final_status
+ # value for the Session
+ if wall_statuses:
+ _final_status = wall_statuses[idx]
+ else:
+ _final_status = final_status if idx == wall_count - 1 else Status.FAIL
+
+ options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(_final_status, {}))
+ wall_manager.finish(
+ wall=w,
+ status=_final_status,
+ status_code_1=rand_choice(options),
+ finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
+ )
+
+ return s
- user_id = user_id or fake.random_int(min=1, max=2_147_483_648)
- started = started or fake.date_time_between(
- start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- end_date=datetime.now(tz=timezone.utc),
- tzinfo=timezone.utc,
- )
+ return _inner
- if session_id is None:
- # session = SessionManager(pg_config=self.pg_config).create_dummy(
- # started=started
- # )
- session = session_factory()
- session_id = session.id
- source = source or rchoice(list(Source))
- req_survey_id = req_survey_id or uuid4().hex
- req_cpi = req_cpi or Decimal(fake.random_int(min=1, max=150) / 100).quantize(
- Decimal(".01"), rounding=ROUND_DOWN
- )
+@pytest.fixture(scope="function")
+def finished_session_factory(
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
+) -> Callable[..., Session]:
- return wall_manager.create(
- session_id=session_id,
- user_id=user_id,
+ def _inner(
+ user: User,
+ # Wall details
+ wall_count: int = 5,
+ wall_req_cpi: Decimal = Decimal(".50"),
+ wall_req_cpis: list[Decimal] | None = None,
+ wall_statuses: list[Status] | None = None,
+ wall_source: Source = Source.TESTING,
+ # Session details
+ final_status: Status = Status.COMPLETE,
+ started: datetime = utc_hour_ago,
+ ) -> Session:
+ s: Session = session_factory(
+ user=user,
+ wall_count=wall_count,
+ wall_req_cpi=wall_req_cpi,
+ wall_req_cpis=wall_req_cpis,
+ wall_statuses=wall_statuses,
+ wall_source=wall_source,
+ final_status=final_status,
started=started,
- source=source,
- req_survey_id=req_survey_id,
- req_cpi=req_cpi,
- buyer_id=buyer_id,
- uuid_id=uuid_id,
)
+ status, status_code_1 = s.determine_session_status()
+ _, _, bp_pay, user_pay = s.determine_payments()
+ session_manager.finish_with_status(
+ s,
+ finished=s.wall_events[-1].finished,
+ payout=bp_pay,
+ user_payout=user_pay,
+ status=status,
+ status_code_1=status_code_1,
+ )
+ return s
return _inner
-@pytest.fixture
+# --- Product ---
+
+
+@pytest.fixture()
def product_factory(product_manager: ProductManager) -> Callable[..., Product]:
def _inner(
- product_id: UUIDStr | None = None,
+ save: bool = True,
+ team: Team | None = None,
team_id: UUIDStr | None = None,
+ business: Business | None = None,
business_id: UUIDStr | None = None,
+ product_id: UUIDStr | None = None,
name: str | None = None,
redirect_url: str | None = None,
harmonizer_domain: str | None = None,
@@ -157,74 +394,60 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]:
user_health_config: UserHealthConfig | None = None,
) -> Product:
"""To be used in tests, where we don't care about certain fields"""
+
product_id = product_id if product_id else uuid4().hex
- team_id = team_id if team_id else uuid4().hex
+
+ team_id = (team.uuid if team else None) or team_id or uuid4().hex
+ business_id = (
+ (business.uuid if business else None) or business_id or uuid4().hex
+ )
+
name = name if name else f"name-{product_id[:12]}"
redirect_url = redirect_url if redirect_url else "https://www.example.com/"
- return product_manager.create(
- product_id=product_id,
- team_id=team_id,
- business_id=business_id,
- name=name,
- redirect_url=redirect_url,
- harmonizer_domain=harmonizer_domain,
- commission_pct=commission_pct,
- sources_config=sources_config,
- payout_config=payout_config,
- session_config=session_config,
- profiling_config=profiling_config,
- user_wallet_config=user_wallet_config,
- user_create_config=user_create_config,
- user_health_config=user_health_config,
- )
+ if save:
+ return product_manager.create(
+ product_id=product_id,
+ team_id=team_id,
+ business_id=business_id,
+ name=name,
+ redirect_url=redirect_url,
+ harmonizer_domain=harmonizer_domain,
+ commission_pct=commission_pct,
+ sources_config=sources_config,
+ payout_config=payout_config,
+ session_config=session_config,
+ profiling_config=profiling_config,
+ user_wallet_config=user_wallet_config,
+ user_create_config=user_create_config,
+ user_health_config=user_health_config,
+ )
+ else:
+ raise ValueError("Unsaved Product not yet supported")
return _inner
-@pytest.fixture
-def session_factory(session_manager: SessionManager):
+@pytest.fixture()
+def product(product_factory: Callable[..., Product]) -> Product:
+ return product_factory(save=True)
- def _inner(
- # -- Create Dummy "optional" -- #
- started: datetime | None = None,
- user: User | None = None,
- # -- Optional -- #
- country_iso: str | None = None,
- device_type: DeviceType | None = None,
- ip: str | None = None,
- bucket: Bucket | None = None,
- url_metadata: dict[str, str] | None = None,
- uuid_id: str | None = None,
- ) -> Session:
- """To be used in tests, where we don't care about certain fields"""
- started = started or fake.date_time_between(
- start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- end_date=datetime(year=2000, month=1, day=1, tzinfo=timezone.utc),
- tzinfo=timezone.utc,
- )
- user = user or User(
- user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex
- )
- return session_manager.create(
- started=started,
- user=user,
- country_iso=country_iso,
- device_type=device_type,
- ip=ip,
- bucket=bucket,
- url_metadata=url_metadata,
- uuid_id=uuid_id,
- )
+@pytest.fixture()
+def unsaved_product(product_factory: Callable[..., Product]) -> Product:
+ return product_factory(save=False)
- return _inner
+
+# --- IP Geoname ---
@pytest.fixture
-def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]:
+def ip_geoname_factory(
+ ip_geoname_manager: IPGeonameManager,
+) -> Callable[..., IPGeoname]:
def _inner(
+ save: bool = True,
geoname_id: PositiveInt | None = None,
continent_code: str | None = None,
continent_name: str | None = None,
@@ -239,31 +462,48 @@ def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGe
time_zone: str | None = None,
is_in_european_union: bool | None = None,
) -> IPGeoname:
-
- return ipgeoname_manager.create(
- geoname_id=geoname_id or randint(1, 999_999_999),
- continent_code=continent_code or "na",
- continent_name=continent_name or "North America",
- country_iso=country_iso or "us",
- country_name=country_name or "United States",
- subdivision_1_iso=subdivision_1_iso or "fl",
- subdivision_1_name=subdivision_1_name or "Florida",
- subdivision_2_iso=subdivision_2_iso,
- subdivision_2_name=subdivision_2_name,
- city_name=city_name,
- metro_code=metro_code,
- time_zone=time_zone,
- is_in_european_union=is_in_european_union,
- )
+ if save:
+ return ip_geoname_manager.create(
+ geoname_id=geoname_id or randint(1, 999_999_999),
+ continent_code=continent_code or "na",
+ continent_name=continent_name or "North America",
+ country_iso=country_iso or "us",
+ country_name=country_name or "United States",
+ subdivision_1_iso=subdivision_1_iso or "fl",
+ subdivision_1_name=subdivision_1_name or "Florida",
+ subdivision_2_iso=subdivision_2_iso,
+ subdivision_2_name=subdivision_2_name,
+ city_name=city_name,
+ metro_code=metro_code,
+ time_zone=time_zone,
+ is_in_european_union=is_in_european_union,
+ )
+ else:
+ raise ValueError("Unsaved IPGeoname not yet supported")
return _inner
-def ipinformation_factory(
- ipinformation_manager: IPInformationManager,
+@pytest.fixture()
+def ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname:
+ return ip_geoname_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname:
+ return ip_geoname_factory(save=True)
+
+
+# --- IP Information ---
+
+
+@pytest.fixture
+def ip_information_factory(
+ ip_information_manager: IPInformationManager,
) -> Callable[..., IPInformation]:
def _inner(
+ save: bool = True,
ip: IPvAnyAddressStr | None = None,
geoname_id: PositiveInt | None = None,
country_iso: str | None = None,
@@ -283,43 +523,186 @@ def ipinformation_factory(
network: str | None = None,
organization: str | None = None,
static_ip_score: float | None = None,
- user_type: UserType | None = None,
+ user_type: AccessType | None = None,
postal_code: str | None = None,
latitude: Decimal | None = None,
longitude: Decimal | None = None,
accuracy_radius: int | None = None,
) -> IPInformation:
- return ipinformation_manager.create(
- ip=ip or fake.ipv4_public(),
- geoname_id=geoname_id,
- country_iso=country_iso or fake.country_code(),
- registered_country_iso=registered_country_iso,
- is_anonymous=is_anonymous,
- is_anonymous_vpn=is_anonymous_vpn,
- is_hosting_provider=is_hosting_provider,
- is_public_proxy=is_public_proxy,
- is_tor_exit_node=is_tor_exit_node,
- is_residential_proxy=is_residential_proxy,
- autonomous_system_number=autonomous_system_number,
- autonomous_system_organization=autonomous_system_organization,
- domain=domain,
- isp=isp,
- mobile_country_code=mobile_country_code,
- mobile_network_code=mobile_network_code,
- network=network,
- organization=organization,
- static_ip_score=static_ip_score,
- user_type=user_type,
- postal_code=postal_code,
- latitude=latitude,
- longitude=longitude,
- accuracy_radius=accuracy_radius,
- )
+ if save:
+ return ip_information_manager.create(
+ ip=ip or fake.ipv4_public(),
+ geoname_id=geoname_id,
+ country_iso=country_iso or fake.country_code(),
+ registered_country_iso=registered_country_iso,
+ is_anonymous=is_anonymous,
+ is_anonymous_vpn=is_anonymous_vpn,
+ is_hosting_provider=is_hosting_provider,
+ is_public_proxy=is_public_proxy,
+ is_tor_exit_node=is_tor_exit_node,
+ is_residential_proxy=is_residential_proxy,
+ autonomous_system_number=autonomous_system_number,
+ autonomous_system_organization=autonomous_system_organization,
+ domain=domain,
+ isp=isp,
+ mobile_country_code=mobile_country_code,
+ mobile_network_code=mobile_network_code,
+ network=network,
+ organization=organization,
+ static_ip_score=static_ip_score,
+ user_type=user_type,
+ postal_code=postal_code,
+ latitude=latitude,
+ longitude=longitude,
+ accuracy_radius=accuracy_radius,
+ )
+ else:
+ raise ValueError("Unsaved IP Information not supported yet")
+
+ return _inner
+
+
+@pytest.fixture
+def ip_information(
+ ip_information_factory: Callable[..., IPInformation],
+) -> IPInformation:
+ return ip_information_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_ip_information(
+ ip_information_factory: Callable[..., IPInformation],
+) -> IPInformation:
+ return ip_information_factory(save=False)
+
+
+# --- IP Record ---
+
+
+@pytest.fixture()
+def ip_record_factory(ip_record_manager: IPRecordManager) -> Callable[..., IPRecord]:
+
+ def _inner(
+ user_id: PositiveInt,
+ save: bool = True,
+ ip: IPvAnyAddressStr | None = None,
+ forwarded_ip1: IPvAnyAddressStr | None = None,
+ forwarded_ip2: IPvAnyAddressStr | None = None,
+ forwarded_ip3: IPvAnyAddressStr | None = None,
+ forwarded_ip4: IPvAnyAddressStr | None = None,
+ forwarded_ip5: IPvAnyAddressStr | None = None,
+ forwarded_ip6: IPvAnyAddressStr | None = None,
+ ) -> IPRecord:
+
+ if save:
+ return ip_record_manager.create(
+ user_id=user_id,
+ ip=ip or fake.ipv4_public(),
+ forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()),
+ forwarded_ip2=(
+ forwarded_ip2 or fake.ipv6() if random() < 0.5 else None
+ ),
+ forwarded_ip3=(
+ forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None
+ ),
+ forwarded_ip4=forwarded_ip4,
+ forwarded_ip5=forwarded_ip5,
+ forwarded_ip6=forwarded_ip6,
+ )
+ else:
+ raise ValueError("Unsaved IP Record not supported")
+
+ return _inner
+
+
+@pytest.fixture()
+def ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord:
+ return ip_record_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord:
+ return ip_record_factory(save=False)
+
+
+# --- User ---
+
+
+@pytest.fixture()
+def user_factory(
+ user_manager: UserManager,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+) -> Callable[..., User]:
+
+ def _inner(
+ save: bool = True,
+ # --- Create dummy "optional" --- #
+ product_user_id: str | None = None,
+ # --- Optional --- #
+ product_id: UUIDStr | None = None,
+ product: Product | None = None,
+ created: datetime | None = None,
+ ) -> User:
+ if save:
+ if product is None:
+ if product_id:
+ raise ValueError("this is broken")
+ product = product_factory()
+
+ product_user_id = product_user_id or uuid4().hex
+
+ u = user_manager.create_user(
+ product_user_id=product_user_id,
+ product_id=product_id,
+ product=product,
+ created=created,
+ )
+
+ u.prefetch_product(pg_config=thl_web_rr)
+ return u
+
+ else:
+ raise ValueError("Unsaved User not supported")
return _inner
+@pytest.fixture()
+def user(
+ user_factory: Callable[..., User],
+) -> User:
+ return user_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_user(
+ user_factory: Callable[..., User],
+) -> User:
+ return user_factory(save=False)
+
+
+@pytest.fixture
+def user_with_wallet(
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+) -> User:
+ # A user on a product with user wallet enabled, but they have no money
+ return user_factory(save=True, product=product_user_wallet_yes)
+
+
+@pytest.fixture
+def user_with_wallet_amt(
+ user_factory: Callable[..., User], product_amt_true: Product
+) -> User:
+ # A user on a product with user wallet enabled, on AMT, but they have no money
+ return user_factory(save=True, product=product_amt_true)
+
+
+# --- User Payout Event ---
+
+
@pytest.fixture
def user_payout_event_factory(
user_payout_event_manager: UserPayoutEventManager,
@@ -374,40 +757,56 @@ def user_payout_event_factory(
return _inner
+@pytest.fixture()
+def user_payout_event(
+ user_payout_event_factory: Callable[..., UserPayoutEvent],
+) -> UserPayoutEvent:
+ return user_payout_event_factory(save=True)
+
+
+@pytest.fixture()
+def unsaved_user_payout_event(
+ user_payout_event_factory: Callable[..., UserPayoutEvent],
+) -> UserPayoutEvent:
+ return user_payout_event_factory(save=True)
+
+
+# -- Brokerage Product Payout Event
+
+
@pytest.fixture
-def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecord]:
+def brokerage_product_payout_event_factory(
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ product_factory: Callable[..., Product],
+) -> Callable[..., BrokerageProductPayoutEvent]:
def _inner(
- user_id: PositiveInt,
- ip: IPvAnyAddressStr | None = None,
- forwarded_ip1: IPvAnyAddressStr | None = None,
- forwarded_ip2: IPvAnyAddressStr | None = None,
- forwarded_ip3: IPvAnyAddressStr | None = None,
- forwarded_ip4: IPvAnyAddressStr | None = None,
- forwarded_ip5: IPvAnyAddressStr | None = None,
- forwarded_ip6: IPvAnyAddressStr | None = None,
- ) -> IPRecord:
- return iprecord_manager.create(
- user_id=user_id,
- ip=ip or fake.ipv4_public(),
- forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()),
- forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None),
- forwarded_ip3=(
- forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None
- ),
- forwarded_ip4=forwarded_ip4,
- forwarded_ip5=forwarded_ip5,
- forwarded_ip6=forwarded_ip6,
+ product: Product | None = None,
+ amount: USDCent | None = None,
+ ext_ref_id: str | None = None,
+ created: AwareDatetime | None = None,
+ ) -> BrokerageProductPayoutEvent:
+
+ product = product or product_factory()
+ amount = amount or USDCent(randint(1, 99_99))
+
+ return business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ product=product,
+ amount=amount,
+ ext_ref_id=ext_ref_id or uuid4().hex,
+ created=created,
)
return _inner
-# class AuditLogManager(PostgresManager):
+# --- Audit Log Manager ---
-@pytest.fixture
-def auditlog_factory(audit_log_manager: AuditLogManager):
+@pytest.fixture()
+def audit_log_factory(audit_log_manager: AuditLogManager) -> Callable[..., AuditLog]:
def _inner(
user_id: PositiveInt,
@@ -425,10 +824,180 @@ def auditlog_factory(audit_log_manager: AuditLogManager):
return audit_log_manager.create(
user_id=user_id,
- level=level or rchoice(list(AuditLogLevel)),
- event_type=event_type or rchoice(list(event_types)),
+ level=level or rand_choice(list(AuditLogLevel)),
+ event_type=event_type or rand_choice(list(event_types)),
event_msg=event_msg,
event_value=event_value,
)
return _inner
+
+
+@pytest.fixture()
+def audit_log(audit_log_factory: Callable[..., AuditLog], user: User) -> AuditLog:
+ return audit_log_factory(user_id=user.user_id)
+
+
+# --- ---
+
+
+@pytest.fixture(scope="session")
+def profiling_info_json() -> str:
+ return (
+ '[{"property_label": "hispanic", "cardinality": "*", "prop_type": "i", "country_iso": "us", '
+ '"property_id": "05170ae296ab49178a075cab2a2073a6", "item_id": "7911ec1468b146ee870951f8ae9cbac1", '
+ '"item_label": "panamanian", "gold_standard": 1, "options": [{"id": "c358c11e72c74fa2880358f1d4be85ab", '
+ '"label": "not_hispanic"}, {"id": "b1d6c475770849bc8e0200054975dc9c", "label": "yes_hispanic"}, '
+ '{"id": "bd1eb44495d84b029e107c188003c2bd", "label": "other_hispanic"}, '
+ '{"id": "f290ad5e75bf4f4ea94dc847f57c1bd3", "label": "mexican"}, '
+ '{"id": "49f50f2801bd415ea353063bfc02d252", "label": "puerto_rican"}, '
+ '{"id": "dcbe005e522f4b10928773926601f8bf", "label": "cuban"}, '
+ '{"id": "467ef8ddb7ac4edb88ba9ef817cbb7e9", "label": "salvadoran"}, '
+ '{"id": "3c98e7250707403cba2f4dc7b877c963", "label": "dominican"}, '
+ '{"id": "981ee77f6d6742609825ef54fea824a8", "label": "guatemalan"}, '
+ '{"id": "81c8057b809245a7ae1b8a867ea6c91e", "label": "colombian"}, '
+ '{"id": "513656d5f9e249fa955c3b527d483b93", "label": "honduran"}, '
+ '{"id": "afc8cddd0c7b4581bea24ccd64db3446", "label": "ecuadorian"}, '
+ '{"id": "61f34b36e80747a89d85e1eb17536f84", "label": "argentinian"}, '
+ '{"id": "5330cfa681d44aa8ade3a6d0ea198e44", "label": "peruvian"}, '
+ '{"id": "e7bceaffd76e486596205d8545019448", "label": "nicaraguan"}, '
+ '{"id": "b7bbb2ebf8424714962e6c4f43275985", "label": "spanish"}, '
+ '{"id": "8bf539785e7a487892a2f97e52b1932d", "label": "venezuelan"}, '
+ '{"id": "7911ec1468b146ee870951f8ae9cbac1", "label": "panamanian"}], "category": [{"id": '
+ '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
+ '"adwords_vertical_id": null}]}, {"property_label": "ethnic_group", "cardinality": "*", "prop_type": '
+ '"i", "country_iso": "us", "property_id": "15070958225d4132b7f6674fcfc979f6", "item_id": '
+ '"64b7114cf08143949e3bcc3d00a5d8a0", "item_label": "other_ethnicity", "gold_standard": 1, "options": [{'
+ '"id": "a72e97f4055e4014a22bee4632cbf573", "label": "caucasians"}, '
+ '{"id": "4760353bc0654e46a928ba697b102735", "label": "black_or_african_american"}, '
+ '{"id": "20ff0a2969fa4656bbda5c3e0874e63b", "label": "asian"}, '
+ '{"id": "107e0a79e6b94b74926c44e70faf3793", "label": "native_hawaiian_or_other_pacific_islander"}, '
+ '{"id": "900fa12691d5458c8665bf468f1c98c1", "label": "native_americans"}, '
+ '{"id": "64b7114cf08143949e3bcc3d00a5d8a0", "label": "other_ethnicity"}], "category": [{"id": '
+ '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
+ '"adwords_vertical_id": null}]}, {"property_label": "educational_attainment", "cardinality": "?", '
+ '"prop_type": "i", "country_iso": "us", "property_id": "2637783d4b2b4075b93e2a156e16e1d8", "item_id": '
+ '"934e7b81d6744a1baa31bbc51f0965d5", "item_label": "other_education", "gold_standard": 1, "options": [{'
+ '"id": "df35ef9e474b4bf9af520aa86630202d", "label": "3rd_grade_completion"}, '
+ '{"id": "83763370a1064bd5ba76d1b68c4b8a23", "label": "8th_grade_completion"}, '
+ '{"id": "f0c25a0670c340bc9250099dcce50957", "label": "not_high_school_graduate"}, '
+ '{"id": "02ff74c872bd458983a83847e1a9f8fd", "label": "high_school_completion"}, '
+ '{"id": "ba8beb807d56441f8fea9b490ed7561c", "label": "vocational_program_completion"}, '
+ '{"id": "65373a5f348a410c923e079ddbb58e9b", "label": "some_college_completion"}, '
+ '{"id": "2d15d96df85d4cc7b6f58911fdc8d5e2", "label": "associate_academic_degree_completion"}, '
+ '{"id": "497b1fedec464151b063cd5367643ffa", "label": "bachelors_degree_completion"}, '
+ '{"id": "295133068ac84424ae75e973dc9f2a78", "label": "some_graduate_completion"}, '
+ '{"id": "e64f874faeff4062a5aa72ac483b4b9f", "label": "masters_degree_completion"}, '
+ '{"id": "cbaec19a636d476385fb8e7842b044f5", "label": "doctorate_degree_completion"}, '
+ '{"id": "934e7b81d6744a1baa31bbc51f0965d5", "label": "other_education"}], "category": [{"id": '
+ '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
+ '"adwords_vertical_id": null}]}, {"property_label": "household_spoken_language", "cardinality": "*", '
+ '"prop_type": "i", "country_iso": "us", "property_id": "5a844571073d482a96853a0594859a51", "item_id": '
+ '"62b39c1de141422896ad4ab3c4318209", "item_label": "dut", "gold_standard": 1, "options": [{"id": '
+ '"f65cd57b79d14f0f8460761ce41ec173", "label": "ara"}, {"id": "6d49de1f8f394216821310abd29392d9", '
+ '"label": "zho"}, {"id": "be6dc23c2bf34c3f81e96ddace22800d", "label": "eng"}, '
+ '{"id": "ddc81f28752d47a3b1c1f3b8b01a9b07", "label": "fre"}, {"id": "2dbb67b29bd34e0eb630b1b8385542ca", '
+ '"label": "ger"}, {"id": "a747f96952fc4b9d97edeeee5120091b", "label": "hat"}, '
+ '{"id": "7144b04a3219433baac86273677551fa", "label": "hin"}, {"id": "e07ff3e82c7149eaab7ea2b39ee6a6dc", '
+ '"label": "ita"}, {"id": "b681eff81975432ebfb9f5cc22dedaa3", "label": "jpn"}, '
+ '{"id": "5cb20440a8f64c9ca62fb49c1e80cdef", "label": "kor"}, {"id": "171c4b77d4204bc6ac0c2b81e38a10ff", '
+ '"label": "pan"}, {"id": "8c3ec18e6b6c4a55a00dd6052e8e84fb", "label": "pol"}, '
+ '{"id": "3ce074d81d384dd5b96f1fb48f87bf01", "label": "por"}, {"id": "6138dc951990458fa88a666f6ddd907b", '
+ '"label": "rus"}, {"id": "e66e5ecc07df4ebaa546e0b436f034bd", "label": "spa"}, '
+ '{"id": "5a981b3d2f0d402a96dd2d0392ec2fcb", "label": "tgl"}, {"id": "b446251bd211403487806c4d0a904981", '
+ '"label": "vie"}, {"id": "92fb3ee337374e2db875fb23f52eed46", "label": "xxx"}, '
+ '{"id": "8b1f590f12f24cc1924d7bdcbe82081e", "label": "ind"}, {"id": "bf3f4be556a34ff4b836420149fd2037", '
+ '"label": "tur"}, {"id": "87ca815c43ba4e7f98cbca98821aa508", "label": "zul"}, '
+ '{"id": "0adbf915a7a64d67a87bb3ce5d39ca54", "label": "may"}, {"id": "62b39c1de141422896ad4ab3c4318209", '
+ '"label": "dut"}], "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", '
+ '"path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": "gender", "cardinality": '
+ '"?", "prop_type": "i", "country_iso": "us", "property_id": "73175402104741549f21de2071556cd7", '
+ '"item_id": "093593e316344cd3a0ac73669fca8048", "item_label": "other_gender", "gold_standard": 1, '
+ '"options": [{"id": "b9fc5ea07f3a4252a792fd4a49e7b52b", "label": "male"}, '
+ '{"id": "9fdb8e5e18474a0b84a0262c21e17b56", "label": "female"}, '
+ '{"id": "093593e316344cd3a0ac73669fca8048", "label": "other_gender"}], "category": [{"id": '
+ '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
+ '"adwords_vertical_id": null}]}, {"property_label": "age_in_years", "cardinality": "?", "prop_type": '
+ '"n", "country_iso": "us", "property_id": "94f7379437874076b345d76642d4ce6d", "item_id": null, '
+ '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", '
+ '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": '
+ '"children_age_gender", "cardinality": "*", "prop_type": "i", "country_iso": "us", "property_id": '
+ '"e926142fcea94b9cbbe13dc7891e1e7f", "item_id": "b7b8074e95334b008e8958ccb0a204f1", "item_label": '
+ '"female_18", "gold_standard": 1, "options": [{"id": "16a6448ec24c48d4993d78ebee33f9b4", '
+ '"label": "male_under_1"}, {"id": "809c04cb2e3b4a3bbd8077ab62cdc220", "label": "female_under_1"}, '
+ '{"id": "295e05bb6a0843bc998890b24c99841e", "label": "no_children"}, '
+ '{"id": "142cb948d98c4ae8b0ef2ef10978e023", "label": "male_0"}, '
+ '{"id": "5a5c1b0e9abc48a98b3bc5f817d6e9d0", "label": "male_1"}, '
+ '{"id": "286b1a9afb884bdfb676dbb855479d1e", "label": "male_2"}, '
+ '{"id": "942ca3cda699453093df8cbabb890607", "label": "male_3"}, '
+ '{"id": "995818d432f643ec8dd17e0809b24b56", "label": "male_4"}, '
+ '{"id": "f38f8b57f25f4cdea0f270297a1e7a5c", "label": "male_5"}, '
+ '{"id": "975df709e6d140d1a470db35023c432d", "label": "male_6"}, '
+ '{"id": "f60bd89bbe0f4e92b90bccbc500467c2", "label": "male_7"}, '
+ '{"id": "6714ceb3ed5042c0b605f00b06814207", "label": "male_8"}, '
+ '{"id": "c03c2f8271d443cf9df380e84b4dea4c", "label": "male_9"}, '
+ '{"id": "11690ee0f5a54cb794f7ddd010d74fa2", "label": "male_10"}, '
+ '{"id": "17bef9a9d14b4197b2c5609fa94b0642", "label": "male_11"}, '
+ '{"id": "e79c8338fe28454f89ccc78daf6f409a", "label": "male_12"}, '
+ '{"id": "3a4f87acb3fa41f4ae08dfe2858238c1", "label": "male_13"}, '
+ '{"id": "36ffb79d8b7840a7a8cb8d63bbc8df59", "label": "male_14"}, '
+ '{"id": "1401a508f9664347aee927f6ec5b0a40", "label": "male_15"}, '
+ '{"id": "6e0943c5ec4a4f75869eb195e3eafa50", "label": "male_16"}, '
+ '{"id": "47d4b27b7b5242758a9fff13d3d324cf", "label": "male_17"}, '
+ '{"id": "9ce886459dd44c9395eb77e1386ab181", "label": "female_0"}, '
+ '{"id": "6499ccbf990d4be5b686aec1c7353fd8", "label": "female_1"}, '
+ '{"id": "d85ceaa39f6d492abfc8da49acfd14f2", "label": "female_2"}, '
+ '{"id": "18edb45c138e451d8cb428aefbb80f9c", "label": "female_3"}, '
+ '{"id": "bac6f006ed9f4ccf85f48e91e99fdfd1", "label": "female_4"}, '
+ '{"id": "5a6a1a8ad00c4ce8be52dcb267b034ff", "label": "female_5"}, '
+ '{"id": "6bff0acbf6364c94ad89507bcd5f4f45", "label": "female_6"}, '
+ '{"id": "d0d56a0a6b6f4516a366a2ce139b4411", "label": "female_7"}, '
+ '{"id": "bda6028468044b659843e2bef4db2175", "label": "female_8"}, '
+ '{"id": "dbb6d50325464032b456357b1a6e5e9c", "label": "female_9"}, '
+ '{"id": "b87a93d7dc1348edac5e771684d63fb8", "label": "female_10"}, '
+ '{"id": "11449d0d98f14e27ba47de40b18921d7", "label": "female_11"}, '
+ '{"id": "16156501e97b4263962cbbb743840292", "label": "female_12"}, '
+ '{"id": "04ee971c89a345cc8141a45bce96050c", "label": "female_13"}, '
+ '{"id": "e818d310bfbc4faba4355e5d2ed49d4f", "label": "female_14"}, '
+ '{"id": "440d25e078924ba0973163153c417ed6", "label": "female_15"}, '
+ '{"id": "78ff804cc9b441c5a524bd91e3d1f8bf", "label": "female_16"}, '
+ '{"id": "4b04d804d7d84786b2b1c22e4ed440f5", "label": "female_17"}, '
+ '{"id": "28bc848cd3ff44c3893c76bfc9bc0c4e", "label": "male_18"}, '
+ '{"id": "b7b8074e95334b008e8958ccb0a204f1", "label": "female_18"}], "category": [{"id": '
+ '"e18ba6e9d51e482cbb19acf2e6f505ce", "label": "Parenting", "path": "/People & Society/Family & '
+ 'Relationships/Family/Parenting", "adwords_vertical_id": "58"}]}, {"property_label": "home_postal_code", '
+ '"cardinality": "?", "prop_type": "x", "country_iso": "us", "property_id": '
+ '"f3b32ebe78014fbeb1ed6ff77d6338bf", "item_id": null, "item_label": null, "gold_standard": 1, '
+ '"category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
+ '"adwords_vertical_id": null}]}, {"property_label": "household_income", "cardinality": "?", "prop_type": '
+ '"n", "country_iso": "us", "property_id": "ff5b1d4501d5478f98de8c90ef996ac1", "item_id": null, '
+ '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", '
+ '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}]'
+ )
+
+
+@pytest.fixture(scope="session")
+def profiling_user_info_json() -> str:
+ return (
+ '{"user_profile_knowledge": [], "marketplace_profile_knowledge": [{"source": "d", "question_id": '
+ '"1", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "pr", '
+ '"question_id": "3", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": '
+ '"h", "question_id": "60", "answer": ["58"], "created": "2023-11-07T16:41:05.234096Z"}, '
+ '{"source": "c", "question_id": "43", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
+ '{"source": "s", "question_id": "211", "answer": ["111"], "created": '
+ '"2023-11-07T16:41:05.234096Z"}, {"source": "s", "question_id": "1843", "answer": ["111"], '
+ '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "h", "question_id": "13959", "answer": ['
+ '"244155"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "33092", '
+ '"answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender", '
+ '"answer": ["10682"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "e", "question_id": '
+ '"gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "f", '
+ '"question_id": "gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": '
+ '"i", "question_id": "gender", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
+ '{"source": "c", "question_id": "137510", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
+ '{"source": "m", "question_id": "gender", "answer": ["1"], "created": '
+ '"2023-11-07T16:41:05.234096Z"}, {"source": "o", "question_id": "gender", "answer": ["male"], '
+ '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender_plus", "answer": ['
+ '"7657644"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "i", "question_id": '
+ '"gender_plus", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", '
+ '"question_id": "income_level", "answer": ["9071"], "created": "2023-11-07T16:41:05.234096Z"}]}'
+ )
diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py
index c8855da..ad96bbb 100644
--- a/test_utils/models/upk/conftest.py
+++ b/test_utils/models/upk/conftest.py
@@ -2,17 +2,15 @@ from __future__ import annotations
import os
import time
-from typing import TYPE_CHECKING
+from collections.abc import Callable
from uuid import UUID
import pandas as pd
import pytest
+from generalresearch.managers.thl.category import CategoryManager
from generalresearch.pg_helper import PostgresConfig
-if TYPE_CHECKING:
- from generalresearch.managers.thl.category import CategoryManager
-
def insert_data_from_csv(
thl_web_rw: PostgresConfig,
@@ -169,9 +167,13 @@ def upk_data(
propertymarketplaceassociation_data,
propertyitemrange_data,
question_data,
-) -> None:
- # Wait a second to make sure the HarmonizerCache refresh loop pulls these in
- time.sleep(2)
+) -> Callable[..., None]:
+
+ def _inner():
+ # Wait a second to make sure the HarmonizerCache refresh loop pulls these in
+ time.sleep(2)
+
+ return _inner
def test_fixtures(upk_data):
diff --git a/generalresearch/managers/network/__init__.py b/test_utils/precision/__init__.py
index e69de29..e69de29 100644
--- a/generalresearch/managers/network/__init__.py
+++ b/test_utils/precision/__init__.py
diff --git a/test_utils/precision/conftest.py b/test_utils/precision/conftest.py
new file mode 100644
index 0000000..7acfe6a
--- /dev/null
+++ b/test_utils/precision/conftest.py
@@ -0,0 +1,129 @@
+from typing import Any
+
+import pytest
+
+
+@pytest.fixture(scope="session")
+def precision_survey_json() -> dict[str, Any]:
+ return {
+ "cpi": "1.44",
+ "country_isos": "ca",
+ "language_isos": "eng",
+ "country_iso": "ca",
+ "language_iso": "eng",
+ "buyer_id": "7047",
+ "bid_loi": 1200,
+ "bid_ir": 0.45,
+ "source": "e",
+ "used_question_ids": ["age", "country_iso", "gender", "gender_1"],
+ "survey_id": "0000",
+ "group_id": "633473",
+ "status": "open",
+ "name": "beauty survey",
+ "survey_guid": "c7f375c5077d4c6c8209ff0b539d7183",
+ "category_id": "-1",
+ "global_conversion": None,
+ "desired_count": 96,
+ "achieved_count": 0,
+ "allowed_devices": "1,2,3",
+ "entry_link": "https://www.opinionetwork.com/survey/entry.aspx?mid=[%MID%]&project=633473&key=%%key%%",
+ "excluded_surveys": "470358,633286",
+ "quotas": [
+ {
+ "name": "25-34,Male,Quebec",
+ "id": "2324110",
+ "guid": "23b5760d24994bc08de451b3e62e77c7",
+ "status": "open",
+ "desired_count": 48,
+ "achieved_count": 0,
+ "termination_count": 0,
+ "overquota_count": 0,
+ "condition_hashes": ["b41e1a3", "bc89ee8", "4124366", "9f32c61"],
+ },
+ {
+ "name": "25-34,Female,Quebec",
+ "id": "2324111",
+ "guid": "0706f1a88d7e4f11ad847c03012e68d2",
+ "status": "open",
+ "desired_count": 48,
+ "achieved_count": 0,
+ "termination_count": 4,
+ "overquota_count": 0,
+ "condition_hashes": ["b41e1a3", "0cdc304", "500af2c", "9f32c61"],
+ },
+ ],
+ "conditions": {
+ "b41e1a3": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "country_iso",
+ "values": ["ca"],
+ "criterion_hash": "b41e1a3",
+ "value_len": 1,
+ "sizeof": 2,
+ },
+ "bc89ee8": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "gender",
+ "values": ["male"],
+ "criterion_hash": "bc89ee8",
+ "value_len": 1,
+ "sizeof": 4,
+ },
+ "4124366": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "gender_1",
+ "values": ["male"],
+ "criterion_hash": "4124366",
+ "value_len": 1,
+ "sizeof": 4,
+ },
+ "9f32c61": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "age",
+ "values": ["25", "26", "27", "28", "29", "30", "31", "32", "33", "34"],
+ "criterion_hash": "9f32c61",
+ "value_len": 10,
+ "sizeof": 20,
+ },
+ "0cdc304": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "gender",
+ "values": ["female"],
+ "criterion_hash": "0cdc304",
+ "value_len": 1,
+ "sizeof": 6,
+ },
+ "500af2c": {
+ "logical_operator": "OR",
+ "value_type": 1,
+ "negate": False,
+ "question_id": "gender_1",
+ "values": ["female"],
+ "criterion_hash": "500af2c",
+ "value_len": 1,
+ "sizeof": 6,
+ },
+ },
+ "expected_end_date": "2024-06-28T10:40:33.000000Z",
+ "created": None,
+ "updated": None,
+ "is_live": True,
+ "all_hashes": [
+ "0cdc304",
+ "b41e1a3",
+ "9f32c61",
+ "bc89ee8",
+ "4124366",
+ "500af2c",
+ ],
+ }
diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py
index 0afc3f5..cc91cff 100644
--- a/test_utils/spectrum/conftest.py
+++ b/test_utils/spectrum/conftest.py
@@ -1,7 +1,9 @@
-import logging
+from __future__ import annotations
+
import time
-from datetime import datetime, timezone
-from typing import TYPE_CHECKING
+from datetime import UTC, datetime
+from decimal import Decimal
+from typing import TYPE_CHECKING, Any
import pytest
@@ -9,21 +11,24 @@ from generalresearch.managers.spectrum.survey import (
SpectrumCriteriaManager,
SpectrumSurveyManager,
)
-from generalresearch.models.spectrum.survey import SpectrumSurvey
+from generalresearch.models.definitions import (
+ LogicalOperator,
+)
+from generalresearch.models.spectrum.survey import (
+ SpectrumCondition,
+ SpectrumSurvey,
+)
+from generalresearch.models.thl.survey.condition import ConditionValueType
from generalresearch.sql_helper import SqlHelper
-from .surveys_json import CONDITIONS, SURVEYS_JSON
-
if TYPE_CHECKING:
from generalresearch.config import GRLBaseSettings
@pytest.fixture(scope="session")
-def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper:
- logging.info(f"{settings.spectrum_rw_db=}")
-
+def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper:
assert settings.spectrum_rw_db is not None
- assert "/unittest-" in settings.spectrum_rw_db.path
+ assert "/unittest-" in str(settings.spectrum_rw_db.path)
return SqlHelper(
dsn=settings.spectrum_rw_db,
@@ -35,27 +40,36 @@ def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper:
@pytest.fixture(scope="session")
def spectrum_criteria_manager(spectrum_rw: SqlHelper) -> SpectrumCriteriaManager:
+ assert spectrum_rw.dsn
+ assert spectrum_rw.dsn.path
assert "/unittest-" in spectrum_rw.dsn.path
return SpectrumCriteriaManager(spectrum_rw)
@pytest.fixture(scope="session")
def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager:
+ assert spectrum_rw.dsn
+ assert spectrum_rw.dsn.path
assert "/unittest-" in spectrum_rw.dsn.path
return SpectrumSurveyManager(spectrum_rw)
@pytest.fixture(scope="session")
def setup_spectrum_surveys(
- spectrum_rw: SqlHelper, spectrum_survey_manager, spectrum_criteria_manager
+ spectrum_rw: SqlHelper,
+ spectrum_survey_manager: SpectrumSurveyManager,
+ spectrum_criteria_manager: SpectrumCriteriaManager,
+ spectrum_conditions: list[SpectrumCondition],
+ spectrum_api_surveys_json: list[str],
) -> None:
- now = datetime.now(timezone.utc)
+ now = datetime.now(UTC)
# make sure these example surveys exist in db
- surveys = [SpectrumSurvey.model_validate_json(x) for x in SURVEYS_JSON]
+ surveys = [SpectrumSurvey.model_validate_json(x) for x in spectrum_api_surveys_json]
for s in surveys:
- s.modified_api = datetime.now(tz=timezone.utc)
+ s.modified_api = datetime.now(tz=UTC)
+
spectrum_survey_manager.create_or_update(surveys)
- spectrum_criteria_manager.update(CONDITIONS)
+ spectrum_criteria_manager.update(spectrum_conditions)
# and make sure they have allocation for 687
spectrum_rw.execute_sql_query(
@@ -66,10 +80,10 @@ def setup_spectrum_surveys(
["687", "GRL", "x", "x", "x", "x"],
commit=True,
)
- supplier687_pk = spectrum_rw.execute_sql_query(
- f"""
- select id from `{spectrum_rw.db}`.spectrum_supplier where supplier_id = '687'"""
- )[0]["id"]
+ supplier687_pk = spectrum_rw.execute_sql_query(f"""
+ select id from `{spectrum_rw.db}`.spectrum_supplier where supplier_id = '687'""")[
+ 0
+ ]["id"]
conn = spectrum_rw.make_connection()
c = conn.cursor()
c.executemany(
@@ -83,3 +97,206 @@ def setup_spectrum_surveys(
conn.commit()
# Wait a second to make sure the spectrum-grpc pulls these from the db into global-vars
time.sleep(1)
+
+
+@pytest.fixture(scope="session")
+def spectrum_api_surveys_json() -> list[str]:
+ return [
+ (
+ '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":["1235","212"],"survey_id":"111111","survey_name":"Exciting New Survey #14472374",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
+ '"include_psids":null,"exclude_psids":null'
+ ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ (
+ '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":["1235","212"],"survey_id":"14472374","survey_name":"Exciting New Survey #14472374",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
+ '"include_psids":null,"exclude_psids":"0408319875e9dbffdc09e86671ad5636,23c4c66ecbc465906d0b0fd798740e64,'
+ '861df4603df3b7f754b8d4b89cbdb313","qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ (
+ '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":["1235","212"],"survey_id":"12345","survey_name":"Exciting New Survey #14472374",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
+ '"include_psids":"7d043991b1494dbbb57786b11c88239c","exclude_psids":null'
+ ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ (
+ '{"cpi":"1.40","country_isos":["us"],"language_isos":["eng"],"buyer_id":"233","bid_loi":null,"source":"s",'
+ '"used_question_ids":["245","244","212","211","225"],"survey_id":"14970164","survey_name":"Exciting New Survey '
+ '#14970164","status":22,"field_end_date":"2024-05-07T16:18:33.000000Z","category_code":"232",'
+ '"calculation_type":"COMPLETES","requires_pii":false,"survey_exclusions":"14970164,29690277",'
+ '"exclusion_period":30,"bid_ir":null,"overall_loi":900,"overall_ir":0.56,"last_block_loi":600,'
+ '"last_block_ir":0.01,"project_last_complete_date":"2024-05-28T04:12:56.297000Z","country_iso":"us",'
+ '"language_iso":"eng","include_psids":null,"exclude_psids":"01c7156fd9639737effbbdebd7fd66f6,'
+ "0508b88f4991bac8b10e9de74ce80194,0a51c627d77cef41f802e51a00126697,15b888176ac4781c2c978a9a05c396f8,"
+ "17bc146b4f7fb05c7058d25da70c6a44,29935289c1f86a4144aab2e12652f305,2fe9d1d451efca10eba4fa4e5e2b74c9,"
+ "c3527b7ef570a1571ea19870f3c25600,cdf2771d57cda9f1bf334382b2b7afd8,cebf3ec50395d973310ea526457dd5a0,"
+ "cf3877cfc15e2e6ef2a56a7a7a37f3d3,dfa691e6d060e3643d5731df30be9f69,e0cb49537182660826aa351e1187809f,"
+ 'edb6d280113ca49561f25fdcb500fde6,fbfba66cfad602f1c26e61e6174eb1f7,fd4307b16fd15e8534a4551c9b6872fc",'
+ '"qualifications":["1ab337d","a01aa68","437774f","dc6065b","82b6ad6"],"quotas":[{"remaining_count":242,'
+ '"condition_hashes":["c23c0b9"]},{"remaining_count":0,"condition_hashes":["5b8c6cf"]},{"remaining_count":126,'
+ '"condition_hashes":["ac35a6e"]},{"remaining_count":110,"condition_hashes":["5e7e5aa"]},{"remaining_count":108,'
+ '"condition_hashes":["9a7aef3"]},{"remaining_count":127,"condition_hashes":["4f75127"]},{"remaining_count":0,'
+ '"condition_hashes":["95437ed"]},{"remaining_count":17,"condition_hashes":["b4b7b95"]},{"remaining_count":16,'
+ '"condition_hashes":["0ab0ae6"]},{"remaining_count":8,"condition_hashes":["6e86fb5"]},{"remaining_count":12,'
+ '"condition_hashes":["24de31e"]},{"remaining_count":69,"condition_hashes":["6bdf350"]},{"remaining_count":411,'
+ '"condition_hashes":["c94d422"]}],"conditions":null,"created_api":"2023-03-30T22:47:36.324000Z",'
+ '"modified_api":"2024-05-30T13:07:16.489000Z","updated":"2024-05-30T21:52:37.493282Z","is_live":true,'
+ '"all_hashes":["c94d422","b4b7b95","6bdf350","6e86fb5","82b6ad6","24de31e","1ab337d","c23c0b9","9a7aef3",'
+ '"ac35a6e","95437ed","5b8c6cf","437774f","a01aa68","5e7e5aa","4f75127","0ab0ae6","dc6065b"]}'
+ ),
+ (
+ '{"cpi":"1.23","country_isos":["au"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":[],"survey_id":"69420","survey_name":"Everyone is eligible AU",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"au","language_iso":"eng",'
+ '"include_psids":null,"exclude_psids":null'
+ ',"qualifications":[],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ (
+ '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":[],"survey_id":"69421","survey_name":"Everyone is eligible US",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
+ '"include_psids":null,"exclude_psids":null'
+ ',"qualifications":[],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ # For partial eligibility
+ (
+ '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
+ '"used_question_ids":["1031", "212"],"survey_id":"999000","survey_name":"Pet owners",'
+ '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
+ '"requires_pii":false,"survey_exclusions":"13947261",'
+ '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
+ '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
+ '"include_psids":null,"exclude_psids":null'
+ ',"qualifications":["0039b0c", "00f60a8"],"quotas":[{"remaining_count":100,'
+ '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
+ '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
+ "}"
+ ),
+ ]
+
+
+def spectrum_conditions() -> list[SpectrumCondition]:
+ # make sure hashes for 111111 are in db
+ c1 = SpectrumCondition(
+ question_id="1001",
+ value_type=ConditionValueType.LIST,
+ values=["a", "b", "c"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c2 = SpectrumCondition(
+ question_id="1001",
+ value_type=ConditionValueType.LIST,
+ values=["a"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c3 = SpectrumCondition(
+ question_id="1002",
+ value_type=ConditionValueType.RANGE,
+ values=["18-24", "30-32"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c4 = SpectrumCondition(
+ question_id="212",
+ value_type=ConditionValueType.LIST,
+ values=["23", "24"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c5 = SpectrumCondition(
+ question_id="1031",
+ value_type=ConditionValueType.LIST,
+ values=["113", "114", "121"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ return [c1, c2, c3, c4, c5]
+
+
+@pytest.fixture(scope="session")
+def spectrum_api_survey_json() -> dict[str, Any]:
+ return {
+ "survey_id": 29333264,
+ "survey_name": "#29333264",
+ "survey_status": 22,
+ "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC),
+ "category": "Exciting New",
+ "category_code": 232,
+ "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC),
+ "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC),
+ "soft_launch": False,
+ "click_balancing": 0,
+ "price_type": 1,
+ "pii": False,
+ "buyer_message": "",
+ "buyer_id": 4726,
+ "incl_excl": 0,
+ "cpi": Decimal("1.20"),
+ "last_complete_date": None,
+ "project_last_complete_date": None,
+ "quotas": [
+ {
+ "quota_id": "c2bc961e-4f26-4223-b409-ebe9165cfdf5",
+ "quantities": {"currently_open": 491, "remaining": 495, "achieved": 0},
+ "criteria": [
+ {
+ "qualification_code": 214,
+ "range_sets": [{"units": 311, "to": 64, "from": 18}],
+ }
+ ],
+ }
+ ],
+ "qualifications": [
+ {
+ "range_sets": [{"units": 311, "to": 64, "from": 18}],
+ "qualification_code": 212,
+ },
+ {"condition_codes": ["111", "117", "112"], "qualification_code": 1202},
+ ],
+ "country_iso": "fr",
+ "language_iso": "fre",
+ "bid_ir": 0.4,
+ "bid_loi": 600,
+ "overall_ir": None,
+ "overall_loi": None,
+ "last_block_ir": None,
+ "last_block_loi": None,
+ "survey_exclusions": set(),
+ "exclusion_period": 0,
+ }
diff --git a/test_utils/spectrum/surveys_json.py b/test_utils/spectrum/surveys_json.py
deleted file mode 100644
index eb747a5..0000000
--- a/test_utils/spectrum/surveys_json.py
+++ /dev/null
@@ -1,140 +0,0 @@
-from generalresearch.models import LogicalOperator
-from generalresearch.models.spectrum.survey import (
- SpectrumCondition,
- SpectrumSurvey,
-)
-from generalresearch.models.thl.survey.condition import ConditionValueType
-
-SURVEYS_JSON = [
- '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":["1235","212"],"survey_id":"111111","survey_name":"Exciting New Survey #14472374",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
- '"include_psids":null,"exclude_psids":null'
- ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
- '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
- '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":["1235","212"],"survey_id":"14472374","survey_name":"Exciting New Survey #14472374",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
- '"include_psids":null,"exclude_psids":"0408319875e9dbffdc09e86671ad5636,23c4c66ecbc465906d0b0fd798740e64,'
- '861df4603df3b7f754b8d4b89cbdb313","qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
- '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
- '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":["1235","212"],"survey_id":"12345","survey_name":"Exciting New Survey #14472374",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
- '"include_psids":"7d043991b1494dbbb57786b11c88239c","exclude_psids":null'
- ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,'
- '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
- '{"cpi":"1.40","country_isos":["us"],"language_isos":["eng"],"buyer_id":"233","bid_loi":null,"source":"s",'
- '"used_question_ids":["245","244","212","211","225"],"survey_id":"14970164","survey_name":"Exciting New Survey '
- '#14970164","status":22,"field_end_date":"2024-05-07T16:18:33.000000Z","category_code":"232",'
- '"calculation_type":"COMPLETES","requires_pii":false,"survey_exclusions":"14970164,29690277",'
- '"exclusion_period":30,"bid_ir":null,"overall_loi":900,"overall_ir":0.56,"last_block_loi":600,'
- '"last_block_ir":0.01,"project_last_complete_date":"2024-05-28T04:12:56.297000Z","country_iso":"us",'
- '"language_iso":"eng","include_psids":null,"exclude_psids":"01c7156fd9639737effbbdebd7fd66f6,'
- "0508b88f4991bac8b10e9de74ce80194,0a51c627d77cef41f802e51a00126697,15b888176ac4781c2c978a9a05c396f8,"
- "17bc146b4f7fb05c7058d25da70c6a44,29935289c1f86a4144aab2e12652f305,2fe9d1d451efca10eba4fa4e5e2b74c9,"
- "c3527b7ef570a1571ea19870f3c25600,cdf2771d57cda9f1bf334382b2b7afd8,cebf3ec50395d973310ea526457dd5a0,"
- "cf3877cfc15e2e6ef2a56a7a7a37f3d3,dfa691e6d060e3643d5731df30be9f69,e0cb49537182660826aa351e1187809f,"
- 'edb6d280113ca49561f25fdcb500fde6,fbfba66cfad602f1c26e61e6174eb1f7,fd4307b16fd15e8534a4551c9b6872fc",'
- '"qualifications":["1ab337d","a01aa68","437774f","dc6065b","82b6ad6"],"quotas":[{"remaining_count":242,'
- '"condition_hashes":["c23c0b9"]},{"remaining_count":0,"condition_hashes":["5b8c6cf"]},{"remaining_count":126,'
- '"condition_hashes":["ac35a6e"]},{"remaining_count":110,"condition_hashes":["5e7e5aa"]},{"remaining_count":108,'
- '"condition_hashes":["9a7aef3"]},{"remaining_count":127,"condition_hashes":["4f75127"]},{"remaining_count":0,'
- '"condition_hashes":["95437ed"]},{"remaining_count":17,"condition_hashes":["b4b7b95"]},{"remaining_count":16,'
- '"condition_hashes":["0ab0ae6"]},{"remaining_count":8,"condition_hashes":["6e86fb5"]},{"remaining_count":12,'
- '"condition_hashes":["24de31e"]},{"remaining_count":69,"condition_hashes":["6bdf350"]},{"remaining_count":411,'
- '"condition_hashes":["c94d422"]}],"conditions":null,"created_api":"2023-03-30T22:47:36.324000Z",'
- '"modified_api":"2024-05-30T13:07:16.489000Z","updated":"2024-05-30T21:52:37.493282Z","is_live":true,'
- '"all_hashes":["c94d422","b4b7b95","6bdf350","6e86fb5","82b6ad6","24de31e","1ab337d","c23c0b9","9a7aef3",'
- '"ac35a6e","95437ed","5b8c6cf","437774f","a01aa68","5e7e5aa","4f75127","0ab0ae6","dc6065b"]}',
- '{"cpi":"1.23","country_isos":["au"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":[],"survey_id":"69420","survey_name":"Everyone is eligible AU",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"au","language_iso":"eng",'
- '"include_psids":null,"exclude_psids":null'
- ',"qualifications":[],"quotas":[{"remaining_count":100,'
- '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
- '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":[],"survey_id":"69421","survey_name":"Everyone is eligible US",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
- '"include_psids":null,"exclude_psids":null'
- ',"qualifications":[],"quotas":[{"remaining_count":100,'
- '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
- # For partial eligibility
- '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",'
- '"used_question_ids":["1031", "212"],"survey_id":"999000","survey_name":"Pet owners",'
- '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",'
- '"requires_pii":false,"survey_exclusions":"13947261",'
- '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,'
- '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",'
- '"include_psids":null,"exclude_psids":null'
- ',"qualifications":["0039b0c", "00f60a8"],"quotas":[{"remaining_count":100,'
- '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",'
- '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true'
- "}",
-]
-
-# make sure hashes for 111111 are in db
-c1 = SpectrumCondition(
- question_id="1001",
- value_type=ConditionValueType.LIST,
- values=["a", "b", "c"],
- negate=False,
- logical_operator=LogicalOperator.OR,
-)
-c2 = SpectrumCondition(
- question_id="1001",
- value_type=ConditionValueType.LIST,
- values=["a"],
- negate=False,
- logical_operator=LogicalOperator.OR,
-)
-c3 = SpectrumCondition(
- question_id="1002",
- value_type=ConditionValueType.RANGE,
- values=["18-24", "30-32"],
- negate=False,
- logical_operator=LogicalOperator.OR,
-)
-c4 = SpectrumCondition(
- question_id="212",
- value_type=ConditionValueType.LIST,
- values=["23", "24"],
- negate=False,
- logical_operator=LogicalOperator.OR,
-)
-c5 = SpectrumCondition(
- question_id="1031",
- value_type=ConditionValueType.LIST,
- values=["113", "114", "121"],
- negate=False,
- logical_operator=LogicalOperator.OR,
-)
-CONDITIONS = [c1, c2, c3, c4, c5]
-survey = SpectrumSurvey.model_validate_json(SURVEYS_JSON[0])
-assert c1.criterion_hash in survey.qualifications
-assert c3.criterion_hash in survey.qualifications
diff --git a/tests/conftest.py b/tests/conftest.py
index 6748592..b69d7ea 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -12,7 +12,6 @@ pytest_plugins = [
"test_utils.managers.contest.conftest",
"test_utils.managers.gr.conftest",
"test_utils.managers.ledger.conftest",
- "test_utils.managers.network.conftest",
"test_utils.managers.thl.conftest",
"test_utils.managers.upk.conftest",
# -- Models
@@ -20,7 +19,9 @@ pytest_plugins = [
"test_utils.models.contest.conftest",
"test_utils.models.gr.conftest",
"test_utils.models.ledger.conftest",
- "test_utils.models.network.conftest",
"test_utils.models.thl.conftest",
"test_utils.models.upk.conftest",
+ # -- Marketplaces
+ "test_utils.precision.conftest",
+ "test_utils.spectrum.conftest",
]
diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py
index e4854e8..2254829 100644
--- a/tests/grliq/managers/test_forensic_data.py
+++ b/tests/grliq/managers/test_forensic_data.py
@@ -1,5 +1,6 @@
from __future__ import annotations
+from collections.abc import Callable
from datetime import timedelta
from typing import TYPE_CHECKING
from uuid import uuid4
@@ -16,6 +17,8 @@ from generalresearch.grliq.models.forensic_result import (
if TYPE_CHECKING:
from generalresearch.grliq.managers.forensic_data import (
GrlIqDataManager,
+ )
+ from generalresearch.grliq.managers.forensic_events import (
GrlIqEventManager,
)
from generalresearch.models.thl.product import Product
@@ -28,10 +31,13 @@ except ImportError:
class TestGrlIqDataManager:
- def test_create_dummy(self, grliq_dm: GrlIqDataManager):
+ def test_factory(
+ self,
+ grliq_data_factory: Callable[..., GrlIqData],
+ ):
from generalresearch.grliq.models.forensic_data import GrlIqData
- gd1: GrlIqData = grliq_dm.create_dummy(is_attempt_allowed=True)
+ gd1: GrlIqData = grliq_data_factory(is_attempt_allowed=True)
assert isinstance(gd1, GrlIqData)
assert isinstance(gd1.results, GrlIqCheckerResults)
@@ -119,7 +125,9 @@ class TestGrlIqDataManager:
class TestForensicDataGetAndFilter:
- def test_events(self, grliq_dm: GrlIqDataManager):
+ def test_events(
+ self, grliq_dm: GrlIqDataManager, grliq_data_factory: Callable[..., GrlIqData]
+ ):
"""If load_events=True, the events and mouse_events attributes should
be an array no matter what. An empty array means that the events were
loaded, but there were no events available.
@@ -129,7 +137,7 @@ class TestForensicDataGetAndFilter:
"""
# Load Events == False
forensic_uuid = uuid4().hex
- grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid)
+ grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid)
instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0]
assert isinstance(instance, GrlIqData)
@@ -144,41 +152,53 @@ class TestForensicDataGetAndFilter:
assert len(instance.events) == 0
assert len(instance.mouse_events) == 0
- def test_timing(self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager):
+ def test_timing(
+ self,
+ grliq_data_factory: Callable[..., GrlIqData],
+ grliq_data_manager: GrlIqDataManager,
+ grliq_event_manager: GrlIqEventManager,
+ ):
forensic_uuid = uuid4().hex
- grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid)
+ grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid)
- instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0]
+ instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0]
- grliq_em.update_or_create_timing(
+ grliq_event_manager.update_or_create_timing(
session_uuid=instance.mid,
timing_data=TimingData(
client_rtts=[100, 200, 150], server_rtts=[150, 120, 120]
),
)
- instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True)
+ instance = grliq_data_manager.get_data(
+ forensic_uuid=forensic_uuid, load_events=True
+ )
assert isinstance(instance, GrlIqData)
assert isinstance(instance.events, list)
assert isinstance(instance.mouse_events, list)
assert isinstance(instance.timing_data, TimingData)
def test_events_events(
- self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager
+ self,
+ grliq_data_factory: Callable[..., GrlIqData],
+ grliq_data_manager: GrlIqDataManager,
+ grliq_event_manager: GrlIqEventManager,
):
forensic_uuid = uuid4().hex
- grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid)
+ grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid)
- instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0]
+ instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0]
- grliq_em.update_or_create_events(
+ grliq_event_manager.update_or_create_events(
session_uuid=instance.mid,
events=[{"a": "b"}],
mouse_events=[],
event_start=instance.created_at,
event_end=instance.created_at + timedelta(minutes=1),
)
- instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True)
+ instance = grliq_data_manager.get_data(
+ forensic_uuid=forensic_uuid, load_events=True
+ )
assert isinstance(instance, GrlIqData)
assert isinstance(instance.events, list)
assert isinstance(instance.mouse_events, list)
@@ -189,11 +209,16 @@ class TestForensicDataGetAndFilter:
assert len(instance.keyboard_events) == 0
def test_events_click(
- self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager
+ self,
+ grliq_data_factory: Callable[..., GrlIqData],
+ grliq_data_manager: GrlIqDataManager,
+ grliq_event_manager: GrlIqEventManager,
):
forensic_uuid = uuid4().hex
- grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid)
- instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True)
+ grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid)
+ instance = grliq_data_manager.get_data(
+ forensic_uuid=forensic_uuid, load_events=True
+ )
click_event = {
"type": "click",
@@ -203,14 +228,16 @@ class TestForensicDataGetAndFilter:
"pointerType": "mouse",
}
me = MouseEvent.from_dict(click_event)
- grliq_em.update_or_create_events(
+ grliq_event_manager.update_or_create_events(
session_uuid=instance.mid,
events=[click_event],
mouse_events=[],
event_start=instance.created_at,
event_end=instance.created_at + timedelta(minutes=1),
)
- instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True)
+ instance = grliq_data_manager.get_data(
+ forensic_uuid=forensic_uuid, load_events=True
+ )
assert isinstance(instance, GrlIqData)
assert isinstance(instance.events, list)
assert isinstance(instance.mouse_events, list)
diff --git a/tests/grliq/managers/test_forensic_results.py b/tests/grliq/managers/test_forensic_results.py
index a030451..86834d0 100644
--- a/tests/grliq/managers/test_forensic_results.py
+++ b/tests/grliq/managers/test_forensic_results.py
@@ -1,18 +1,21 @@
from __future__ import annotations
+from collections.abc import Callable
from typing import TYPE_CHECKING
if TYPE_CHECKING:
- from generalresearch.grliq.managers.forensic_data import GrlIqDataManager
from generalresearch.grliq.managers.forensic_results import (
GrlIqCategoryResultsReader,
)
+ from generalresearch.grliq.models.forensic_data import GrlIqData
class TestGrlIqCategoryResultsReader:
def test_filter_category_results(
- self, grliq_dm: GrlIqDataManager, grliq_crr: GrlIqCategoryResultsReader
+ self,
+ grliq_data_factory: Callable[..., GrlIqData],
+ grliq_crr: GrlIqCategoryResultsReader,
):
from generalresearch.grliq.models.forensic_result import (
GrlIqForensicCategoryResult,
@@ -20,8 +23,8 @@ class TestGrlIqCategoryResultsReader:
)
# this is just testing that it doesn't fail
- grliq_dm.create_dummy(is_attempt_allowed=True)
- grliq_dm.create_dummy(is_attempt_allowed=True)
+ grliq_data_factory(is_attempt_allowed=True)
+ grliq_data_factory(is_attempt_allowed=True)
res = grliq_crr.filter_category_results(limit=2, phase=Phase.OFFERWALL_ENTER)[0]
assert res.get("category_result")
diff --git a/tests/grliq/models/test_forensic_data.py b/tests/grliq/models/test_forensic_data.py
index 4fbf962..a901dc3 100644
--- a/tests/grliq/models/test_forensic_data.py
+++ b/tests/grliq/models/test_forensic_data.py
@@ -9,16 +9,16 @@ if TYPE_CHECKING:
class TestGrlIqData:
- def test_supported_fonts(self, grliq_data: "GrlIqData"):
+ def test_supported_fonts(self, grliq_data: GrlIqData):
s = grliq_data.supported_fonts_binary
assert len(s) == 1043
assert "Ubuntu" in grliq_data.supported_fonts
- def test_battery(self, grliq_data: "GrlIqData"):
+ def test_battery(self, grliq_data: GrlIqData):
assert not grliq_data.battery_charging
assert grliq_data.battery_level == 0.41
- def test_base(self, grliq_data: "GrlIqData"):
+ def test_base(self, grliq_data: GrlIqData):
from generalresearch.grliq.models.forensic_data import Platform
assert grliq_data.timezone == "America/Los_Angeles"
@@ -41,7 +41,7 @@ class TestGrlIqData:
# Testing things that will cause a validation error, should only be
# because something is "corrupt", not b/c the user is a baddie
- def test_corrupt(self, grliq_data: "GrlIqData"):
+ def test_corrupt(self, grliq_data: GrlIqData):
"""Test for timestamp and timezone offset mismatch validation."""
from generalresearch.grliq.models.forensic_data import GrlIqData
diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py
index 31d1720..6d715fa 100644
--- a/tests/incite/collections/test_df_collection_base.py
+++ b/tests/incite/collections/test_df_collection_base.py
@@ -1,18 +1,18 @@
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
import pandas as pd
import pytest
from pandera.pandas import DataFrameSchema
-from generalresearch.incite.collections import (
+from generalresearch.incite.collections.base import (
DFCollection,
DFCollectionType,
)
-from test_utils.incite.conftest import mnt_filepath
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.pg_helper import PostgresConfig
df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST]
@@ -24,7 +24,7 @@ class TestDFCollectionBase:
"""
- def test_init(self, mnt_filepath: "GRLDatasets", df_coll_type: DFCollectionType):
+ def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType):
"""Try to initialize the DFCollection with various invalid parameters"""
with pytest.raises(expected_exception=ValueError) as cm:
DFCollection(archive_path=mnt_filepath.data_src)
@@ -46,24 +46,28 @@ class TestDFCollectionBase:
class TestDFCollectionBaseProperties:
@pytest.mark.skip
- def test_df_collection_items(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_df_collection_items(
+ self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType
+ ):
instance = DFCollection(
data_type=df_coll_type,
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- offset="100d",
+ start=datetime(year=1800, month=1, day=1, tzinfo=UTC),
+ finished=datetime(year=1900, month=1, day=1, tzinfo=UTC),
+ offset="100D",
archive_path=mnt_filepath.archive_path(enum_type=df_coll_type),
)
assert len(instance.interval_range) == len(instance.items)
assert len(instance.items) == 366
- def test_df_collection_progress(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_df_collection_progress(
+ self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType
+ ):
instance = DFCollection(
data_type=df_coll_type,
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- offset="100d",
+ start=datetime(year=1800, month=1, day=1, tzinfo=UTC),
+ finished=datetime(year=1900, month=1, day=1, tzinfo=UTC),
+ offset="100D",
archive_path=mnt_filepath.archive_path(enum_type=df_coll_type),
)
@@ -71,7 +75,9 @@ class TestDFCollectionBaseProperties:
assert isinstance(instance.progress, pd.DataFrame)
assert instance.progress.shape == (366, 6)
- def test_df_collection_schema(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_df_collection_schema(
+ self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType
+ ):
instance1 = DFCollection(
data_type=DFCollectionType.WALL, archive_path=mnt_filepath.data_src
)
@@ -88,12 +94,12 @@ class TestDFCollectionBaseProperties:
class TestDFCollectionBaseMethods:
@pytest.mark.skip
- def test_initial_load(self, mnt_filepath: "GRLDatasets", thl_web_rr):
+ def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr: PostgresConfig):
instance = DFCollection(
pg_config=thl_web_rr,
data_type=DFCollectionType.USER,
- start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=timezone.utc),
- finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=timezone.utc),
+ start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC),
+ finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC),
offset="2min",
archive_path=mnt_filepath.data_src,
)
diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py
index 136d234..7a8793d 100644
--- a/tests/incite/collections/test_df_collection_item_base.py
+++ b/tests/incite/collections/test_df_collection_item_base.py
@@ -1,30 +1,30 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
import pytest
-from generalresearch.incite.collections import (
+from generalresearch.incite.collections.base import (
+ MYSQL_ALLOWED_COLL_TYPES,
DFCollection,
DFCollectionItem,
DFCollectionType,
)
-from generalresearch.pg_helper import PostgresConfig
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
-
-df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST]
+ from generalresearch.pg_helper import PostgresConfig
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemBase:
-
- def test_init(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType):
collection = DFCollection(
data_type=df_coll_type,
- offset="100d",
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
+ offset="100D",
+ start=datetime(year=1800, month=1, day=1, tzinfo=UTC),
+ finished=datetime(year=1900, month=1, day=1, tzinfo=UTC),
archive_path=mnt_filepath.archive_path(enum_type=df_coll_type),
)
@@ -34,23 +34,23 @@ class TestDFCollectionItemBase:
assert isinstance(item, DFCollectionItem)
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemProperties:
-
@pytest.mark.skip
- def test_filename(self, df_coll_type):
+ def test_filename(self, df_coll_type: DFCollectionType):
pass
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemMethods:
-
- def test_has_mysql_false(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_has_mysql_false(
+ self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType
+ ):
collection = DFCollection(
data_type=df_coll_type,
- offset="100d",
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
+ offset="100D",
+ start=datetime(year=1800, month=1, day=1, tzinfo=UTC),
+ finished=datetime(year=1900, month=1, day=1, tzinfo=UTC),
archive_path=mnt_filepath.archive_path(enum_type=df_coll_type),
)
@@ -58,13 +58,16 @@ class TestDFCollectionItemMethods:
assert not instance1.has_mysql()
def test_has_mysql_true(
- self, thl_web_rr: PostgresConfig, mnt_filepath: "GRLDatasets", df_coll_type
+ self,
+ thl_web_rr: PostgresConfig,
+ mnt_filepath: GRLDatasets,
+ df_coll_type: DFCollectionType,
):
collection = DFCollection(
data_type=df_coll_type,
- offset="100d",
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
+ offset="100D",
+ start=datetime(year=1800, month=1, day=1, tzinfo=UTC),
+ finished=datetime(year=1900, month=1, day=1, tzinfo=UTC),
archive_path=mnt_filepath.archive_path(enum_type=df_coll_type),
pg_config=thl_web_rr,
)
@@ -74,5 +77,5 @@ class TestDFCollectionItemMethods:
assert instance2.has_mysql()
@pytest.mark.skip
- def test_update_partial_archive(self, df_coll_type):
+ def test_update_partial_archive(self, df_coll_type: DFCollectionType):
pass
diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py
index 8b8bcbe..eeabb41 100644
--- a/tests/incite/collections/test_df_collection_item_thl_web.py
+++ b/tests/incite/collections/test_df_collection_item_thl_web.py
@@ -1,17 +1,19 @@
from __future__ import annotations
-from collections.abc import Generator
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable, Generator
+from datetime import UTC, datetime, timedelta
from itertools import product as iter_product
from os.path import join as pjoin
from pathlib import Path, PurePath
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import dask.dataframe as dd
import pandas as pd
import pytest
-from distributed import Client, Scheduler, Worker
+from dask.distributed import Client as DaskClient
+from dask.distributed import Scheduler as DaskScheduler
+from dask.distributed import Worker as DaskWorker
# noinspection PyUnresolvedReferences
from distributed.utils_test import (
@@ -22,18 +24,21 @@ from pandera.pandas import DataFrameSchema
from pydantic import FilePath
from generalresearch.incite.base import CollectionItemBase
-from generalresearch.incite.collections import (
- DFCollectionItem,
+from generalresearch.incite.collections.base import (
DFCollectionType,
)
from generalresearch.incite.schemas import ARCHIVE_AFTER
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
from generalresearch.sql_helper import PostgresDsn
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.base import (
+ DFCollection,
+ DFCollectionItem,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
fake = Faker()
@@ -52,12 +57,11 @@ unsupported_mock_types = {
}
-def combo_object() -> Generator[str, None, None]:
- for x in iter_product(
+def combo_object() -> Generator[tuple[DFCollectionType, str]]:
+ yield from iter_product(
df_collections,
- ["15min", "45min", "1H"],
- ):
- yield from x
+ ["15min", "45min", "1h"],
+ )
class TestDFCollectionItemBase:
@@ -71,8 +75,12 @@ class TestDFCollectionItemBase:
argnames="df_collection_data_type, offset", argvalues=combo_object()
)
class TestDFCollectionItemProperties:
-
- def test_filename(self, df_collection_data_type, df_collection, offset: str):
+ def test_filename(
+ self,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
+ offset: str,
+ ):
for i in df_collection.items:
assert isinstance(i.filename, str)
@@ -88,38 +96,59 @@ class TestDFCollectionItemProperties:
argnames="df_collection_data_type, offset", argvalues=combo_object()
)
class TestDFCollectionItemPropertiesBase:
-
- def test_name(self, df_collection_data_type, offset: str, df_collection):
+ def test_name(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.name, str)
- def test_finish(self, df_collection_data_type, offset: str, df_collection):
+ def test_finish(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.finish, datetime)
- def test_interval(self, df_collection_data_type, offset: str, df_collection):
+ def test_interval(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.interval, pd.Interval)
def test_partial_filename(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection: DFCollection,
):
for i in df_collection.items:
assert isinstance(i.partial_filename, str)
- def test_empty_filename(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_filename(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_filename, str)
- def test_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.path, FilePath)
- def test_partial_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_partial_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.partial_path, FilePath)
- def test_empty_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_path, FilePath)
@@ -135,26 +164,25 @@ class TestDFCollectionItemPropertiesBase:
),
)
class TestDFCollectionItemMethod:
-
- def test_has_mysql(
+ def test_has_postgres(
self,
- df_collection,
- thl_web_rr: PostgresConfig,
+ df_collection_data_type: DFCollectionType,
offset: str,
duration: timedelta,
- df_collection_data_type,
- delete_df_collection,
+ delete_df_collection: Callable[..., None],
+ df_collection: DFCollection,
+ thl_web_rr: PostgresConfig,
):
delete_df_collection(coll=df_collection)
df_collection.pg_config = None
for i in df_collection.items:
- assert not i.has_mysql()
+ assert not i.has_postgres()
# Confirm that the regular connection should work as expected
df_collection.pg_config = thl_web_rr
for i in df_collection.items:
- assert i.has_mysql()
+ assert i.has_postgres()
# Make a fake connection and confirm it does NOT work
df_collection.pg_config = PostgresConfig(
@@ -163,17 +191,14 @@ class TestDFCollectionItemMethod:
statement_timeout=1,
)
for i in df_collection.items:
- assert not i.has_mysql()
+ assert not i.has_postgres()
@pytest.mark.skip
def test_update_partial_archive(
self,
- df_collection,
+ df_collection_data_type: DFCollectionType,
offset: str,
duration: timedelta,
- thl_web_rw: PostgresConfig,
- df_collection_data_type,
- delete_df_collection,
):
# for i in collection.items:
# assert i.update_partial_archive()
@@ -183,29 +208,16 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_create_partial_archive(
self,
- df_collection,
+ df_collection_data_type: DFCollectionType,
offset: str,
- duration: str,
- create_main_accounts,
- thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
- user_factory: Callable[..., User],
- product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ duration: timedelta,
):
- assert 1 + 1 == 2
+ pass
def test_dict(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- df_collection,
- delete_df_collection,
+ df_collection: DFCollection,
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
@@ -225,18 +237,17 @@ class TestDFCollectionItemMethod:
def test_from_mysql(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
- create_main_accounts,
+ create_main_accounts: Callable[..., None],
thl_web_rw: PostgresConfig,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -249,38 +260,32 @@ class TestDFCollectionItemMethod:
for item in df_collection.items:
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
- delete_df_collection(coll=df_collection)
- df = item.from_mysql()
if df_collection.data_type == DFCollectionType.LEDGER:
- assert df is None
- else:
- assert df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ continue
+ delete_df_collection(coll=df_collection)
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
+ assert df.empty
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
incite_item_factory(user=u1, item=item)
- df = item.from_mysql()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
- if df_collection.data_type == DFCollectionType.LEDGER:
- # The number of rows in this dataframe will change depending
- # on the mocking of data. It's because if the account has
- # user wallet on, then there will be more transactions for
- # example.
- assert df.shape[0] > 0
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
- def test_from_mysql_standard(
+ def test_from_postgres_standard(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -292,48 +297,31 @@ class TestDFCollectionItemMethod:
item: DFCollectionItem
if df_collection.data_type == DFCollectionType.LEDGER:
- # We're using parametrize, so this If statement is just to
- # confirm other Item Types will always raise an assertion
- with pytest.raises(expected_exception=AssertionError) as cm:
- res = item.from_mysql_standard()
- assert (
- "Can't call from_mysql_standard for Ledger DFCollectionItem"
- in str(cm.value)
- )
-
continue
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
- df = item.from_mysql_standard()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
incite_item_factory(user=u1, item=item)
- df = item.from_mysql_standard()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
assert df.shape[0] > 0
- def test_from_mysql_ledger(
+ def test_from_postgres_ledger(
self,
- df_collection,
- user: User,
- create_main_accounts,
- offset: str,
- duration: timedelta,
- thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type != DFCollectionType.LEDGER:
return
@@ -348,14 +336,14 @@ class TestDFCollectionItemMethod:
# Okay, now continue with the actual Ledger Item tests... we need
# to ensure that this item.start - item.finish range hasn't had
# any prior transactions created within that range.
- assert item.from_mysql_ledger() is None
+ assert item.from_postgres_ledger() is None
# Create main accounts doesn't matter because it doesn't
# add any transactions to the db
- assert item.from_mysql_ledger() is None
+ assert item.from_postgres_ledger() is None
incite_item_factory(user=u1, item=item)
- df = item.from_mysql_ledger()
+ df = item.from_postgres_ledger()
assert isinstance(df, pd.DataFrame)
# Not only is this a np.int64 to int comparison, but I also know it
@@ -373,19 +361,12 @@ class TestDFCollectionItemMethod:
def test_to_archive(
self,
- df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -400,7 +381,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods.
- df = item.from_mysql()
+ df = item.from_postgres_standard()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Write the basic archive, the issue is that because it's
@@ -411,17 +392,12 @@ class TestDFCollectionItemMethod:
def test__to_archive(
self,
- df_collection_data_type,
- df_collection,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- client_no_amm,
- user: User,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ mnt_filepath: GRLDatasets,
):
"""We already have a test for the "non-private" version of this,
which primarily just uses the respective Client to determine if
@@ -443,7 +419,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Confirm a missing ddf (shouldn't bc of type hint) should
@@ -484,19 +460,28 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_to_archive_numbered_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_initial_load(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_clear_corrupt_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@@ -506,37 +491,36 @@ class TestDFCollectionItemMethod:
argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])),
)
class TestDFCollectionItemMethodBase:
-
- @pytest.mark.skip
- def test_path_exists(
- self, df_collection_data_type, offset: str, duration: timedelta
- ):
- pass
-
- @pytest.mark.skip
- def test_next_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
- ):
- pass
-
@pytest.mark.skip
def test_search_highest_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_tmp_filename(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
- def test_tmp_path(self, df_collection_data_type, offset: str, duration: timedelta):
+ def test_tmp_path(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
+ ):
pass
def test_is_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
"""
test_has_empty was merged into this because item.has_empty is
@@ -553,7 +537,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_empty()
def test_has_partial_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
assert not item.has_partial_archive()
@@ -561,7 +546,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_partial_archive()
def test_has_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
# (1) Originally, nothing exists... so let's just make a file and
@@ -598,7 +584,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_archive(include_empty=True)
def test_delete_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
item: DFCollectionItem
@@ -621,9 +608,11 @@ class TestDFCollectionItemMethodBase:
assert not item.partial_path.exists()
def test_should_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
- schema: DataFrameSchema = df_collection._schema
+ schema: DataFrameSchema = df_collection.type_schema
+ assert schema.metadata
aa = schema.metadata[ARCHIVE_AFTER]
# It shouldn't be None, it can be timedelta(seconds=0)
@@ -632,19 +621,23 @@ class TestDFCollectionItemMethodBase:
for item in df_collection.items:
item: DFCollectionItem
- if datetime.now(tz=timezone.utc) > item.finish + aa:
+ if datetime.now(tz=UTC) > item.finish + aa:
assert item.should_archive()
else:
assert not item.should_archive()
@pytest.mark.skip
def test_set_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
def test_valid_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
# Originally, nothing has been saved or anything.. so confirm it
# always comes back as None
@@ -668,18 +661,28 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_validate_df(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_from_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
def test__to_dict(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
+ df_collection: DFCollection,
):
for item in df_collection.items:
@@ -698,30 +701,39 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_delete_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_cleanup_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_delete_dangling_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@gen_cluster(client=True, nthreads=[("127.0.0.1", 1)])
-async def test_client(client, s, worker):
+async def test_client(client: DaskClient, s: DaskScheduler, worker: DaskWorker):
"""c,s,a are all required - the secondary Worker (b) is not required"""
- assert isinstance(client, Client)
- assert isinstance(s, Scheduler)
- assert isinstance(worker, Worker)
+ assert isinstance(client, DaskClient)
+ assert isinstance(s, DaskScheduler)
+ assert isinstance(worker, DaskWorker)
@pytest.mark.parametrize(
@@ -730,12 +742,18 @@ async def test_client(client, s, worker):
)
@gen_cluster(client=True, nthreads=[("127.0.0.1", 1)])
@pytest.mark.anyio
-async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str):
+async def test_client_parametrize(
+ c: DaskClient,
+ s: DaskScheduler,
+ w: DaskWorker,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+):
"""c,s,a are all required - the secondary Worker (b) is not required"""
- assert isinstance(c, Client), f"c is not Client, it's {type(c)}"
- assert isinstance(s, Scheduler), f"s is not Scheduler, it's {type(s)}"
- assert isinstance(w, Worker), f"w is not Worker, it's {type(w)}"
+ assert isinstance(c, DaskClient), f"c is not Client, it's {type(c)}"
+ assert isinstance(s, DaskScheduler), f"s is not Scheduler, it's {type(s)}"
+ assert isinstance(w, DaskWorker), f"w is not Worker, it's {type(w)}"
assert df_collection_data_type is not None
assert isinstance(offset, str)
@@ -751,22 +769,15 @@ async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str)
argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])),
)
class TestDFCollectionItemFunctionalTest:
-
def test_to_archive_and_ddf(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- client_no_amm,
- df_collection,
- user: User,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -804,17 +815,11 @@ class TestDFCollectionItemFunctionalTest:
def test_filesize_estimate(
self,
- df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- client_no_amm,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- df_collection_data_type,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""A functional test to write some Parquet files for the
DFCollection and then confirm that the files get written
@@ -828,8 +833,6 @@ class TestDFCollectionItemFunctionalTest:
import pyarrow.parquet as pq
- from generalresearch.models.thl.user import User
-
if df_collection.data_type in unsupported_mock_types:
return
delete_df_collection(coll=df_collection)
@@ -853,18 +856,13 @@ class TestDFCollectionItemFunctionalTest:
def test_to_archive_client(
self,
- client_no_amm,
- df_collection,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=df_collection)
df_collection._client = client_no_amm
@@ -880,7 +878,7 @@ class TestDFCollectionItemFunctionalTest:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
assert isinstance(ddf, dd.DataFrame)
@@ -893,7 +891,8 @@ class TestDFCollectionItemFunctionalTest:
@pytest.mark.skip
def test_get_items(
- self, df_collection, product: Product, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
with pytest.warns(expected_warning=ResourceWarning) as cm:
df_collection.get_items_last365()
@@ -906,27 +905,22 @@ class TestDFCollectionItemFunctionalTest:
def test_saving_protections(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
):
"""Don't allow creating an archive for data that will likely be
overwritten or updated
"""
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
u1: User = user_factory(product=product)
- schema: DataFrameSchema = df_collection._schema
+ schema: DataFrameSchema = df_collection.type_schema
+ assert schema.metadata
aa = schema.metadata[ARCHIVE_AFTER]
assert isinstance(aa, timedelta)
@@ -948,21 +942,14 @@ class TestDFCollectionItemFunctionalTest:
def test_empty_item(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
+ df_collection: DFCollection,
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
for item in df_collection.items:
assert not item.has_empty()
- df: pd.DataFrame = item.from_mysql()
+ df: pd.DataFrame = item.from_db()
# We do this check b/c the Ledger returns back None and
# I don't want it to fail when we go to make a ddf
@@ -976,18 +963,13 @@ class TestDFCollectionItemFunctionalTest:
def test_file_touching(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- mnt_filepath,
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=df_collection)
df_collection._client = client_no_amm
diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py
index 981f62e..0f79b81 100644
--- a/tests/incite/collections/test_df_collection_thl_marketplaces.py
+++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py
@@ -1,25 +1,26 @@
-from datetime import datetime, timezone
+from collections.abc import Generator
+from datetime import UTC, datetime
from itertools import product
from typing import TYPE_CHECKING
import pytest
from pandera.pandas import Column, DataFrameSchema, Index
-from generalresearch.incite.collections import DFCollection, DFCollectionType
+from generalresearch.incite.collections.base import DFCollection, DFCollectionType
from generalresearch.incite.collections.thl_marketplaces import (
InnovateSurveyHistoryCollection,
MorningSurveyTimeseriesCollection,
SagoSurveyHistoryCollection,
SpectrumSurveyTimeseriesCollection,
)
-from test_utils.incite.conftest import mnt_filepath
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.pg_helper import PostgresConfig
-def combo_object():
- for x in product(
+def combo_object() -> Generator[tuple[type, str]]:
+ yield from product(
[
InnovateSurveyHistoryCollection,
MorningSurveyTimeseriesCollection,
@@ -27,14 +28,19 @@ def combo_object():
SpectrumSurveyTimeseriesCollection,
],
["5min", "6H", "30D"],
- ):
- yield from x
+ )
@pytest.mark.parametrize("df_coll, offset", combo_object())
class TestDFCollection_thl_marketplaces:
- def test_init(self, mnt_filepath, df_coll, offset, spectrum_rw):
+ def test_init(
+ self,
+ mnt_filepath: GRLDatasets,
+ df_coll: DFCollection,
+ offset: str,
+ spectrum_rw: PostgresConfig,
+ ):
assert issubclass(df_coll, DFCollection)
# This is stupid, but we need to pull the default from the
@@ -43,7 +49,7 @@ class TestDFCollection_thl_marketplaces:
assert isinstance(data_type, DFCollectionType)
# (1) Can't be totally empty, needs a path...
- with pytest.raises(expected_exception=Exception) as cm:
+ with pytest.raises(expected_exception=ValueError):
instance = df_coll()
# (2) Confirm it only needs the archive_path
@@ -57,8 +63,8 @@ class TestDFCollection_thl_marketplaces:
archive_path=mnt_filepath.archive_path(enum_type=data_type),
sql_helper=spectrum_rw,
offset=offset,
- start=datetime(year=2023, month=6, day=1, minute=0, tzinfo=timezone.utc),
- finished=datetime(year=2023, month=6, day=1, minute=5, tzinfo=timezone.utc),
+ start=datetime(year=2023, month=6, day=1, minute=0, tzinfo=UTC),
+ finished=datetime(year=2023, month=6, day=1, minute=5, tzinfo=UTC),
)
assert isinstance(instance, DFCollection)
@@ -66,7 +72,7 @@ class TestDFCollection_thl_marketplaces:
assert isinstance(instance._schema, DataFrameSchema)
assert isinstance(instance._schema.index, Index)
- for c in instance._schema.columns.keys():
+ for c in instance._schema.columns:
assert isinstance(c, str)
col = instance._schema.columns[c]
assert isinstance(col, Column)
diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py
index b09d44c..7253dd0 100644
--- a/tests/incite/collections/test_df_collection_thl_web.py
+++ b/tests/incite/collections/test_df_collection_thl_web.py
@@ -3,25 +3,20 @@ from __future__ import annotations
from collections.abc import Generator
from datetime import datetime
from itertools import product
-from typing import TYPE_CHECKING
import dask.dataframe as dd
import pandas as pd
import pytest
from pandera.pandas import DataFrameSchema
-from generalresearch.incite.collections import DFCollection, DFCollectionType
-
-if TYPE_CHECKING:
- from generalresearch.incite.base import GRLDatasets
- from generalresearch.incite.collections import (
- DFCollectionItem,
- DFCollectionType,
- )
+from generalresearch.incite.collections.base import (
+ DFCollection,
+ DFCollectionType,
+)
-def combo_object() -> Generator[tuple, None, None]:
- for x in product(
+def combo_object() -> Generator[tuple[DFCollectionType, str]]:
+ yield from product(
[
DFCollectionType.USER,
DFCollectionType.WALL,
@@ -30,9 +25,8 @@ def combo_object() -> Generator[tuple, None, None]:
DFCollectionType.AUDIT_LOG,
DFCollectionType.LEDGER,
],
- ["30min", "1H"],
- ):
- yield from x
+ ["30min", "1h"],
+ )
@pytest.mark.parametrize(
@@ -41,7 +35,10 @@ def combo_object() -> Generator[tuple, None, None]:
class TestDFCollection_thl_web:
def test_init(
- self, df_collection_data_type: DFCollectionType, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
assert isinstance(df_collection_data_type, DFCollectionType)
assert isinstance(df_collection, DFCollection)
@@ -52,12 +49,12 @@ class TestDFCollection_thl_web:
)
class TestDFCollection_thl_web_Properties:
- def test_items(self, df_collection):
+ def test_items(self, df_collection: DFCollection):
assert isinstance(df_collection.items, list)
for i in df_collection.items:
assert i._collection == df_collection
- def test__schema(self, df_collection):
+ def test__schema(self, df_collection: DFCollection):
assert isinstance(df_collection._schema, DataFrameSchema)
@@ -67,16 +64,16 @@ class TestDFCollection_thl_web_Properties:
class TestDFCollection_thl_web_BaseProperties:
@pytest.mark.skip
- def test__interval_range(self, df_collection):
+ def test__interval_range(self, df_collection: DFCollection):
pass
- def test_interval_start(self, df_collection):
+ def test_interval_start(self, df_collection: DFCollection):
assert isinstance(df_collection.interval_start, datetime)
- def test_interval_range(self, df_collection):
+ def test_interval_range(self, df_collection: DFCollection):
assert isinstance(df_collection.interval_range, list)
- def test_progress(self, df_collection):
+ def test_progress(self, df_collection: DFCollection):
assert isinstance(df_collection.progress, pd.DataFrame)
@@ -86,17 +83,21 @@ class TestDFCollection_thl_web_BaseProperties:
class TestDFCollection_thl_web_Methods:
@pytest.mark.skip
- def test_initial_loads(self, df_collection_data_type, df_collection, offset):
+ def test_initial_loads(
+ self, df_collection_data_type, df_collection: DFCollection, offset: str
+ ):
pass
@pytest.mark.skip
def test_fetch_force_rr_latest(
- self, df_collection_data_type, df_collection, offset: str
+ self, df_collection_data_type, df_collection: DFCollection, offset: str
):
pass
@pytest.mark.skip
- def test_force_rr_latest(self, df_collection_data_type, df_collection, offset):
+ def test_force_rr_latest(
+ self, df_collection_data_type, df_collection: DFCollection, offset: str
+ ):
pass
@@ -105,63 +106,108 @@ class TestDFCollection_thl_web_Methods:
)
class TestDFCollection_thl_web_BaseMethods:
- def test_fetch_all_paths(self, df_collection_data_type, offset: str, df_collection):
+ def test_fetch_all_paths(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
res = df_collection.fetch_all_paths(
items=None, force_rr_latest=False, include_partial=False
)
assert isinstance(res, list)
@pytest.mark.skip
- def test_ddf(self, df_collection_data_type, offset: str, df_collection):
+ def test_ddf(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
res = df_collection.ddf()
assert isinstance(res, dd.DataFrame)
# -- cleanup --
@pytest.mark.skip
def test_schedule_cleanup(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
- def test_cleanup(self, df_collection_data_type, offset: str, df_collection):
+ def test_cleanup(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
pass
@pytest.mark.skip
def test_cleanup_partials(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
def test_clear_tmp_archives(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
def test_clear_corrupt_archives(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
def test_rebuild_symlinks(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
# -- Source timing --
@pytest.mark.skip
- def test_get_item(self, df_collection_data_type, offset: str, df_collection):
+ def test_get_item(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
pass
@pytest.mark.skip
- def test_get_item_start(self, df_collection_data_type, offset: str, df_collection):
+ def test_get_item_start(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
pass
@pytest.mark.skip
- def test_get_items(self, df_collection_data_type, offset: str, df_collection):
+ def test_get_items(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
+ ):
# If we get all the items from the start of the collection, it
# should include all the items!
res1 = df_collection.items
@@ -170,18 +216,27 @@ class TestDFCollection_thl_web_BaseMethods:
@pytest.mark.skip
def test_get_items_from_year(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
def test_get_items_last90(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
@pytest.mark.skip
def test_get_items_last365(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ df_collection: DFCollection,
):
pass
diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py
index 47f243e..71b2442 100644
--- a/tests/incite/mergers/foundations/test_enriched_session.py
+++ b/tests/incite/mergers/foundations/test_enriched_session.py
@@ -1,20 +1,35 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from itertools import product
-from typing import Optional
+from typing import TYPE_CHECKING
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
from generalresearch.incite.schemas.admin_responses import (
AdminPOPSessionSchema,
)
-from generalresearch.pg_helper import PostgresConfig
-from test_utils.incite.collections.conftest import (
- session_collection,
- wall_collection,
-)
+
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
+ from generalresearch.models.admin.request import (
+ ReportRequest,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
@pytest.mark.parametrize(
@@ -30,17 +45,16 @@ class TestEnrichedSession:
def test_base(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- session_collection,
- enriched_session_merge,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
thl_web_rr: PostgresConfig,
- delete_df_collection,
- incite_item_factory,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
@@ -77,31 +91,31 @@ class TestEnrichedSession:
class TestEnrichedSessionAdmin:
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2020, month=3, day=14, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
return "1d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return timedelta(days=5)
def test_to_admin_response(
self,
- event_report_request,
- enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
+ event_report_request: ReportRequest,
+ enriched_session_merge: EnrichedSessionMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
+ session_report_request: ReportRequest,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
@@ -112,7 +126,7 @@ class TestEnrichedSessionAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
index 96c214f..877d22f 100644
--- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py
+++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
@@ -1,16 +1,30 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import timedelta
from itertools import product as iter_product
+from typing import TYPE_CHECKING
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
-from test_utils.incite.collections.conftest import (
- wall_collection,
- task_adj_collection,
- session_collection,
-)
-from test_utils.incite.mergers.conftest import enriched_wall_merge
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ TaskAdjustmentDFCollection,
+ WallDFCollection,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_task_adjust import (
+ EnrichedTaskAdjustMerge,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
@pytest.mark.parametrize(
@@ -27,19 +41,18 @@ class TestEnrichedTaskAdjust:
@pytest.mark.skip
def test_base(
self,
- client_no_amm,
- user_factory,
- product,
- task_adj_collection,
- wall_collection,
- session_collection,
- enriched_wall_merge,
- enriched_task_adjust_merge,
- incite_item_factory,
- delete_df_collection,
- thl_web_rr,
+ client_no_amm: DaskClient,
+ user_factory: Callable[..., User],
+ product: Product,
+ task_adj_collection: TaskAdjustmentDFCollection,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ enriched_task_adjust_merge: EnrichedTaskAdjustMerge,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py
index 8f4995b..2b9afb8 100644
--- a/tests/incite/mergers/foundations/test_enriched_wall.py
+++ b/tests/incite/mergers/foundations/test_enriched_wall.py
@@ -1,34 +1,33 @@
-from datetime import timedelta, timezone, datetime
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from itertools import product as iter_product
-from typing import Optional
+from typing import TYPE_CHECKING
import dask.dataframe as dd
import pandas as pd
import pytest
-
-# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
+from dask.distributed import Client as DaskClient
from generalresearch.incite.mergers.foundations.enriched_wall import (
EnrichedWallMergeItem,
)
-from test_utils.incite.collections.conftest import (
- session_collection,
- wall_collection,
-)
-from test_utils.incite.conftest import incite_item_factory
-from test_utils.incite.mergers.conftest import (
- enriched_wall_merge,
-)
+
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+ )
+
+ # noinspection PyUnresolvedReferences
+ from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+ )
+ from generalresearch.models.admin.request import ReportRequest
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
@pytest.mark.parametrize(
@@ -39,17 +38,16 @@ class TestEnrichedWall:
def test_base(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- thl_web_rr,
- session_collection,
- enriched_wall_merge,
- delete_df_collection,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ thl_web_rr: PostgresConfig,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
@@ -82,15 +80,15 @@ class TestEnrichedWall:
def test_base_item(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- session_collection,
- enriched_wall_merge,
- delete_df_collection,
- thl_web_rr,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ delete_df_collection: Callable[..., None],
+ thl_web_rr: PostgresConfig,
+ incite_item_factory: Callable[..., None],
):
# -- Build & Setup
delete_df_collection(coll=session_collection)
@@ -118,7 +116,7 @@ class TestEnrichedWall:
try:
modified_time1 = path.stat().st_mtime
- except (Exception,):
+ except OSError:
modified_time1 = 0
item.build(
@@ -158,18 +156,23 @@ class TestEnrichedWall:
class TestEnrichedWallToAdmin:
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2020, month=3, day=14, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
return "1d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return timedelta(days=5)
- def test_empty(self, enriched_wall_merge, client_no_amm, start):
+ def test_empty(
+ self,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ start: datetime,
+ ):
from generalresearch.models.admin.request import ReportRequest
rr = ReportRequest.model_validate({"interval": "5min", "start": start})
@@ -186,18 +189,18 @@ class TestEnrichedWallToAdmin:
def test_to_admin_response(
self,
- event_report_request,
- enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- user,
- session_factory,
- delete_df_collection,
- product_factory,
- user_factory,
- start,
+ event_report_request: ReportRequest,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ thl_web_rr: PostgresConfig,
+ user: User,
+ session_factory: Callable[..., Session],
+ delete_df_collection: Callable[..., None],
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ start: datetime,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
@@ -208,7 +211,7 @@ class TestEnrichedWallToAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=2,
wall_req_cpi=Decimal("1.00"),
diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py
index f96bfb4..8c4b2f7 100644
--- a/tests/incite/mergers/foundations/test_user_id_product.py
+++ b/tests/incite/mergers/foundations/test_user_id_product.py
@@ -1,24 +1,22 @@
-from datetime import timedelta, datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from itertools import product
+from typing import TYPE_CHECKING
import pandas as pd
import pytest
-
-# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
+from dask.distributed import Client as DaskClient
from generalresearch.incite.mergers.foundations.user_id_product import (
UserIdProductMergeItem,
)
-from test_utils.incite.mergers.conftest import user_id_product_merge
+
+if TYPE_CHECKING:
+ # noinspection PyUnresolvedReferences
+ from generalresearch.incite.mergers.foundations.user_id_product import (
+ UserIdProductMerge,
+ )
@pytest.mark.parametrize(
@@ -27,25 +25,28 @@ from test_utils.incite.mergers.conftest import user_id_product_merge
product(
["12h", "3D"],
[timedelta(days=5)],
- [
- (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace(
- microsecond=0
- )
- ],
+ [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)],
)
),
)
class TestUserIDProduct:
@pytest.mark.skip
- def test_base(self, client_no_amm, user_id_product_merge):
+ def test_base(
+ self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge
+ ):
ddf = user_id_product_merge.ddf()
df = client_no_amm.compute(collections=ddf, sync=True)
assert isinstance(df, pd.DataFrame)
assert not df.empty
@pytest.mark.skip
- def test_base_item(self, client_no_amm, user_id_product_merge, user_collection):
+ def test_base_item(
+ self,
+ client_no_amm: DaskClient,
+ user_id_product_merge: UserIdProductMerge,
+ user_collection,
+ ):
assert len(user_id_product_merge.items) == 1
for item in user_id_product_merge.items:
@@ -55,7 +56,7 @@ class TestUserIDProduct:
try:
modified_time1 = path.stat().st_mtime
- except (Exception,):
+ except OSError:
modified_time1 = 0
user_id_product_merge.build(client=client_no_amm, user_coll=user_collection)
@@ -64,7 +65,9 @@ class TestUserIDProduct:
assert modified_time2 > modified_time1
@pytest.mark.skip
- def test_read(self, client_no_amm, user_id_product_merge):
+ def test_read(
+ self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge
+ ):
users_ddf = user_id_product_merge.ddf()
df = client_no_amm.compute(collections=users_ddf, sync=True)
diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py
index ec507bc..7ed3996 100644
--- a/tests/incite/mergers/test_merge_collection.py
+++ b/tests/incite/mergers/test_merge_collection.py
@@ -1,17 +1,22 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from itertools import product
+from typing import TYPE_CHECKING
import pandas as pd
import pytest
from pandera.pandas import DataFrameSchema
-from generalresearch.incite.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeType,
)
-from test_utils.incite.conftest import mnt_filepath
-merge_types = list(e for e in MergeType if e != MergeType.TEST)
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+
+merge_types = [e for e in MergeType if e != MergeType.TEST]
@pytest.mark.parametrize(
@@ -21,17 +26,20 @@ merge_types = list(e for e in MergeType if e != MergeType.TEST)
merge_types,
["5min", "6h", "14D"],
[timedelta(days=30)],
- [
- (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace(
- microsecond=0
- )
- ],
+ [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)],
)
),
)
class TestMergeCollection:
- def test_init(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_init(
+ self,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ mnt_filepath: GRLDatasets,
+ ):
with pytest.raises(expected_exception=ValueError) as cm:
MergeCollection(archive_path=mnt_filepath.data_src)
assert "Must explicitly provide a merge_type" in str(cm.value)
@@ -42,7 +50,14 @@ class TestMergeCollection:
)
assert instance.merge_type == merge_type
- def test_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -53,7 +68,14 @@ class TestMergeCollection:
assert len(instance.interval_range) == len(instance.items)
- def test_progress(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_progress(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -67,7 +89,14 @@ class TestMergeCollection:
assert instance.progress.shape[1] == 7
assert instance.progress["group_by"].isnull().all()
- def test_schema(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_schema(
+ self,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ ):
instance = MergeCollection(
merge_type=merge_type,
archive_path=mnt_filepath.archive_path(enum_type=merge_type),
@@ -75,7 +104,14 @@ class TestMergeCollection:
assert isinstance(instance._schema, DataFrameSchema)
- def test_load(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_load(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
start=start,
@@ -87,7 +123,14 @@ class TestMergeCollection:
# Confirm that there are no archives available yet
assert instance.progress.has_archive.eq(False).all()
- def test_get_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_get_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
start=start,
finished=start + duration,
diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py
index 96f8789..baf1bc4 100644
--- a/tests/incite/mergers/test_merge_collection_item.py
+++ b/tests/incite/mergers/test_merge_collection_item.py
@@ -1,17 +1,19 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import timedelta
from itertools import product
from pathlib import PurePath
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.incite.mergers import MergeCollectionItem, MergeType
-from generalresearch.incite.mergers.foundations.enriched_session import (
- EnrichedSessionMerge,
-)
-from generalresearch.incite.mergers.foundations.enriched_wall import (
- EnrichedWallMerge,
-)
-from test_utils.incite.mergers.conftest import merge_collection
+from generalresearch.incite.mergers.base import MergeType
+
+if TYPE_CHECKING:
+ from generalresearch.incite.mergers.base import (
+ MergeCollection,
+ MergeCollectionItem,
+ )
@pytest.mark.parametrize(
@@ -26,7 +28,10 @@ from test_utils.incite.mergers.conftest import merge_collection
)
class TestMergeCollectionItem:
- def test_file_naming(self, merge_collection, offset, duration, start):
+ def test_file_naming(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
items: list[MergeCollectionItem] = merge_collection.items
@@ -41,7 +46,10 @@ class TestMergeCollectionItem:
assert i._collection.offset in i.filename
assert i.start.strftime("%Y-%m-%d-%H-%M-%S") in i.filename
- def test_archives(self, merge_collection, offset, duration, start):
+ def test_archives(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
for i in merge_collection.items:
@@ -51,10 +59,13 @@ class TestMergeCollectionItem:
assert not i.has_partial_archive()
assert i.has_archive() == i.path_exists(generic_path=i.path)
- res = set([i.should_archive() for i in merge_collection.items])
+ res = {i.should_archive() for i in merge_collection.items}
assert len(res) == 1
- def test_item_to_archive(self, merge_collection, offset, duration, start):
+ def test_item_to_archive(
+ self,
+ merge_collection: MergeCollection,
+ ):
for item in merge_collection.items:
item: MergeCollectionItem
assert not item.has_archive()
diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py
index 6f96108..9ec188b 100644
--- a/tests/incite/mergers/test_pop_ledger.py
+++ b/tests/incite/mergers/test_pop_ledger.py
@@ -1,18 +1,28 @@
-from datetime import timedelta, datetime, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from itertools import product as iter_product
-from typing import Optional
+from typing import TYPE_CHECKING
import pandas as pd
import pytest
-from distributed.utils_test import client_no_amm
+from dask.distributed import Client as DaskClient
from generalresearch.incite.schemas.mergers.pop_ledger import (
numerical_col_names,
)
-from test_utils.incite.collections.conftest import ledger_collection
-from test_utils.incite.conftest import mnt_filepath, incite_item_factory
-from test_utils.incite.mergers.conftest import pop_ledger_merge
-from test_utils.managers.ledger.conftest import create_main_accounts
+
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.thl_web import (
+ LedgerDFCollection,
+ SessionDFCollection,
+ )
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
@pytest.mark.parametrize(
@@ -27,25 +37,25 @@ from test_utils.managers.ledger.conftest import create_main_accounts
class TestMergePOPLedger:
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2020, month=3, day=14, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2020, month=3, day=14, tzinfo=UTC)
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return timedelta(days=5)
def test_base(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- product,
- user_factory,
- create_main_accounts,
- thl_lm,
- delete_df_collection,
- incite_item_factory,
- delete_ledger_db,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ product: Product,
+ user_factory: Callable[..., User],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
):
from generalresearch.models.thl.ledger import LedgerAccount
@@ -79,19 +89,19 @@ class TestMergePOPLedger:
# --
- user_wallet_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet(
- user=u
+ thl_ledger_manager.get_account_or_create_user_wallet(user=u)
+ cash_account: LedgerAccount = thl_ledger_manager.get_account_cash()
+ rev_account: LedgerAccount = (
+ thl_ledger_manager.get_account_task_complete_revenue()
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
item_finishes = [i.finish for i in ledger_collection.items]
item_finishes.sort(reverse=True)
last_item_finish = item_finishes[0]
# Pure SQL based lookups
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert cash_balance > rev_balance
# (1) Test Cash Account
@@ -129,39 +139,42 @@ class TestMergePOPLedger:
def test_pydantic_init(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
- product,
- user_factory,
- create_main_accounts,
- offset,
- duration,
- start,
- thl_lm,
- incite_item_factory,
- delete_df_collection,
- delete_ledger_db,
- session_collection,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
+ product: Product,
+ user_factory: Callable[..., User],
+ create_main_accounts: Callable[..., None],
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ session_collection: SessionDFCollection,
):
+ from generalresearch.models.thl.finance import ProductBalances
from generalresearch.models.thl.ledger import LedgerAccount
from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.finance import ProductBalances
u = user_factory(product=product, created=session_collection.start)
assert ledger_collection.finished is not None
assert isinstance(u.product, Product)
delete_ledger_db()
- create_main_accounts(),
+ create_main_accounts()
+
delete_df_collection(coll=ledger_collection)
- bp_account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(
+ bp_account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet(
product=u.product
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
+ cash_account: LedgerAccount = thl_ledger_manager.get_account_cash()
+ rev_account: LedgerAccount = (
+ thl_ledger_manager.get_account_task_complete_revenue()
+ )
for item in ledger_collection.items:
incite_item_factory(item=item, user=u)
@@ -191,8 +204,10 @@ class TestMergePOPLedger:
assert instance.payout == instance.net == instance.bp_payment_credit
assert instance.available_balance < instance.net
assert instance.available_balance + instance.retainer == instance.net
- assert instance.balance == thl_lm.get_account_balance(bp_account)
- assert df["bp_payment.CREDIT"].sum() == thl_lm.get_account_balance(bp_account)
+ assert instance.balance == thl_ledger_manager.get_account_balance(bp_account)
+ assert df["bp_payment.CREDIT"].sum() == thl_ledger_manager.get_account_balance(
+ bp_account
+ )
# (2) Filter by the Cash Account
ddf = pop_ledger_merge.ddf(
@@ -205,7 +220,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
assert df["bp_payment.CREDIT"].sum() == 0
assert cash_balance > 0
assert df["mp_payment.CREDIT"].sum() == 0
@@ -222,7 +237,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert rev_balance == 0
assert df["bp_payment.CREDIT"].sum() == 0
assert df["mp_payment.DEBIT"].sum() == 0
@@ -230,27 +245,28 @@ class TestMergePOPLedger:
def test_resample(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
- user_factory,
- product,
- create_main_accounts,
- offset,
- duration,
- start,
- thl_lm,
- delete_df_collection,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
+ user_factory: Callable[..., User],
+ product: Product,
+ create_main_accounts: Callable[..., None],
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
assert ledger_collection.finished is not None
delete_df_collection(coll=ledger_collection)
u1: User = user_factory(product=product)
- bp_account = thl_lm.get_account_or_create_bp_wallet(product=u1.product)
+ bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=u1.product
+ )
for item in ledger_collection.items:
incite_item_factory(user=u1, item=item)
@@ -280,7 +296,7 @@ class TestMergePOPLedger:
assert isinstance(df.index, pd.Index)
assert isinstance(df.index, pd.DatetimeIndex)
- bp_account_balance = thl_lm.get_account_balance(account=bp_account)
+ thl_ledger_manager.get_account_balance(account=bp_account)
# Initial sum
initial_sum = df.sum().sum()
diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py
index 4c2df6b..d83a98c 100644
--- a/tests/incite/mergers/test_ym_survey_merge.py
+++ b/tests/incite/mergers/test_ym_survey_merge.py
@@ -1,25 +1,28 @@
-from datetime import timedelta, timezone, datetime
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from itertools import product
+from typing import TYPE_CHECKING
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
+ from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
-
-from test_utils.incite.collections.conftest import wall_collection, session_collection
-from test_utils.incite.mergers.conftest import (
- enriched_session_merge,
- ym_survey_wall_merge,
-)
@pytest.mark.parametrize(
@@ -28,11 +31,7 @@ from test_utils.incite.mergers.conftest import (
product(
["12h", "3D"],
[timedelta(days=30)],
- [
- (datetime.now(tz=timezone.utc) - timedelta(days=35)).replace(
- microsecond=0
- )
- ],
+ [(datetime.now(tz=UTC) - timedelta(days=35)).replace(microsecond=0)],
)
),
)
@@ -46,18 +45,17 @@ class TestYMSurveyMerge:
def test_base(
self,
- client_no_amm,
- user_factory,
- product,
- ym_survey_wall_merge,
- wall_collection,
- session_collection,
- enriched_session_merge,
- delete_df_collection,
- incite_item_factory,
- thl_web_rr,
+ client_no_amm: DaskClient,
+ user_factory: Callable[..., User],
+ product: Product,
+ ym_survey_wall_merge: YMSurveyWallMerge,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
+ thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
user: User = user_factory(product=product, created=session_collection.start)
@@ -85,10 +83,10 @@ class TestYMSurveyMerge:
assert enriched_session_merge.progress.has_archive.eq(True).all()
ddf = enriched_session_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df1: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df1, pd.DataFrame)
+ assert not df1.empty
# --
@@ -102,18 +100,18 @@ class TestYMSurveyMerge:
# --
ddf = ym_survey_wall_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df2: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df2, pd.DataFrame)
+ assert not df2.empty
# --
- assert df.product_id.nunique() == 1
- assert df.team_id.nunique() == 1
- assert df.source.nunique() > 1
+ assert df2.product_id.nunique() == 1
+ assert df2.team_id.nunique() == 1
+ assert df2.source.nunique() > 1
- started_min_ts = df.started.min()
- started_max_ts = df.started.max()
+ started_min_ts = df2.started.min()
+ started_max_ts = df2.started.max()
assert type(started_min_ts) is pd.Timestamp
assert type(started_max_ts) is pd.Timestamp
diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py
index 43aa399..d2658ea 100644
--- a/tests/incite/schemas/test_admin_responses.py
+++ b/tests/incite/schemas/test_admin_responses.py
@@ -1,15 +1,17 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from random import sample
-from typing import List
import numpy as np
import pandas as pd
+import pandera as pa
import pytest
from generalresearch.incite.schemas import empty_dataframe_from_schema
from generalresearch.incite.schemas.admin_responses import (
- AdminPOPSchema,
SIX_HOUR_SECONDS,
+ AdminPOPSchema,
)
from generalresearch.locales import Localelator
@@ -17,12 +19,14 @@ from generalresearch.locales import Localelator
class TestAdminPOPSchema:
schema_df = empty_dataframe_from_schema(AdminPOPSchema)
countries = list(Localelator().get_all_countries())[:5]
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10) # noqa
+ ]
@classmethod
def assign_valid_vals(cls, df: pd.DataFrame) -> pd.DataFrame:
for c in df.columns:
- check_attrs: dict = AdminPOPSchema.columns[c].checks[0].statistics
+ check_attrs = AdminPOPSchema.columns[c].checks[0].statistics
df[c] = np.random.randint(
check_attrs["min_value"], check_attrs["max_value"], df.shape[0]
)
@@ -30,7 +34,7 @@ class TestAdminPOPSchema:
return df
def test_empty(self):
- with pytest.raises(Exception):
+ with pytest.raises(pa.errors.SchemaError):
AdminPOPSchema.validate(pd.DataFrame())
def test_new_empty_df(self):
@@ -43,7 +47,7 @@ class TestAdminPOPSchema:
def test_valid(self):
# (1) Works with raw naive datetime
dates = [
- datetime(year=2024, month=1, day=i, tzinfo=None).isoformat()
+ datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() # noqa
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -58,7 +62,10 @@ class TestAdminPOPSchema:
assert isinstance(df, pd.DataFrame)
# (2) Works with isoformat naive datetime
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) # noqa
+ for i in range(1, 10)
+ ]
df = pd.DataFrame(
index=pd.MultiIndex.from_product(
iterables=[dates, self.countries], names=["index0", "index1"]
@@ -72,8 +79,7 @@ class TestAdminPOPSchema:
def test_index_tz_parser(self):
tz_dates = [
- datetime(year=2024, month=1, day=i, tzinfo=timezone.utc)
- for i in range(1, 10)
+ datetime(year=2024, month=1, day=i, tzinfo=UTC) for i in range(1, 10)
]
df = pd.DataFrame(
@@ -85,16 +91,16 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
# Initially, they're all set with a timezone
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz == timezone.utc for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all(ts.tz == UTC for ts in timestmaps)
# After validation, the timezone is removed
df = AdminPOPSchema.validate(df)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all(ts.tz is None for ts in timestmaps)
def test_index_tz_no_future_beyond_one_year(self):
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
tz_dates = [now + timedelta(days=i * 365) for i in range(1, 10)]
df = pd.DataFrame(
@@ -125,12 +131,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, float) for v in vals])
+ assert all(isinstance(v, float) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# --- int to str ---
@@ -144,12 +150,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, int) for v in vals])
+ assert all(isinstance(v, int) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# a = 1
assert isinstance(df, pd.DataFrame)
@@ -157,9 +163,7 @@ class TestAdminPOPSchema:
def test_invalid_parsing(self):
# (1) Timezones AND as strings will still parse correctly
tz_str_dates = [
- datetime(
- year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc
- ).isoformat()
+ datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC).isoformat()
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -173,12 +177,12 @@ class TestAdminPOPSchema:
df = AdminPOPSchema.validate(df, lazy=True)
assert isinstance(df, pd.DataFrame)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all(ts.tz is None for ts in timestmaps)
# (2) Timezones are removed
dates = [
- datetime(year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc)
+ datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC)
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -190,13 +194,13 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
# Has tz before validation, and none after
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is timezone.utc for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all(ts.tz is UTC for ts in timestmaps)
df = AdminPOPSchema.validate(df, lazy=True)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all(ts.tz is None for ts in timestmaps)
def test_clipping(self):
df = pd.DataFrame(
diff --git a/tests/incite/schemas/test_thl_web.py b/tests/incite/schemas/test_thl_web.py
index 7f4434b..9b34ce0 100644
--- a/tests/incite/schemas/test_thl_web.py
+++ b/tests/incite/schemas/test_thl_web.py
@@ -16,7 +16,7 @@ class TestWallSchema:
df = pd.DataFrame(columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_no_rows(self):
@@ -24,7 +24,7 @@ class TestWallSchema:
df = pd.DataFrame(index=["uuid"], columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_new_empty_df(self):
@@ -50,7 +50,7 @@ class TestSessionSchema:
df = pd.DataFrame(columns=THLSessionSchema.columns.keys())
df.set_index("uuid", inplace=True)
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_no_rows(self):
@@ -58,7 +58,7 @@ class TestSessionSchema:
df = pd.DataFrame(index=["id"], columns=THLSessionSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_new_empty_df(self):
diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py
index 7e6605f..1a664a2 100644
--- a/tests/incite/test_collection_base.py
+++ b/tests/incite/test_collection_base.py
@@ -1,7 +1,10 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta, timezone
from os.path import exists as pexists
from os.path import join as pjoin
from pathlib import Path
+from typing import TYPE_CHECKING
from uuid import uuid4
import numpy as np
@@ -10,21 +13,21 @@ import pytest
from _pytest._code.code import ExceptionInfo
from generalresearch.incite.base import CollectionBase
-from test_utils.incite.conftest import mnt_filepath
-AGO_15min = (datetime.now(tz=timezone.utc) - timedelta(minutes=15)).replace(
- microsecond=0
-)
-AGO_1HR = (datetime.now(tz=timezone.utc) - timedelta(hours=1)).replace(microsecond=0)
-AGO_2HR = (datetime.now(tz=timezone.utc) - timedelta(hours=2)).replace(microsecond=0)
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+
+AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0)
+AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0)
+AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0)
class TestCollectionBase:
- def test_init(self, mnt_filepath):
+ def test_init(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.df.empty is True
- def test_init_df(self, mnt_filepath):
+ def test_init_df(self, mnt_filepath: GRLDatasets):
# Only an empty pd.DataFrame can ever be provided
instance = CollectionBase(
df=pd.DataFrame({}), archive_path=mnt_filepath.data_src
@@ -46,11 +49,11 @@ class TestCollectionBase:
)
assert "Do not provide a pd.DataFrame" in str(cm.value)
- def test_init_start(self, mnt_filepath):
+ def test_init_start(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(
- start=datetime.now(tz=timezone.utc) - timedelta(days=10),
+ start=datetime.now(tz=UTC) - timedelta(days=10),
archive_path=mnt_filepath.data_src,
)
assert "Collection.start must not have microseconds" in str(cm.value)
@@ -66,9 +69,7 @@ class TestCollectionBase:
assert "Timezone is not UTC" in str(cm.value)
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- assert instance.start == datetime(
- year=2018, month=1, day=1, tzinfo=timezone.utc
- )
+ assert instance.start == datetime(year=2018, month=1, day=1, tzinfo=UTC)
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
@@ -79,7 +80,7 @@ class TestCollectionBase:
cm.value
)
- def test_init_archive_path(self, mnt_filepath):
+ def test_init_archive_path(self, mnt_filepath: GRLDatasets):
"""DirectoryPath is apparently smart enough to confirm that the
directory path exists.
"""
@@ -104,7 +105,7 @@ class TestCollectionBase:
CollectionBase(archive_path=new_path)
assert "Path does not point to a directory" in str(cm.value)
- def test_init_offset(self, mnt_filepath):
+ def test_init_offset(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(offset="1:X", archive_path=mnt_filepath.data_src)
@@ -112,7 +113,7 @@ class TestCollectionBase:
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
- CollectionBase(offset=f"59sec", archive_path=mnt_filepath.data_src)
+ CollectionBase(offset="59sec", archive_path=mnt_filepath.data_src)
assert "Must be equal to, or longer than 1 min" in str(cm.value)
with pytest.raises(expected_exception=ValueError) as cm:
@@ -123,14 +124,14 @@ class TestCollectionBase:
class TestCollectionBaseProperties:
- def test_items(self, mnt_filepath):
+ def test_items(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- x = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
- def test_interval_range(self, mnt_filepath):
+ def test_interval_range(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
# Private method requires the end parameter
with pytest.raises(expected_exception=AssertionError) as cm:
@@ -145,14 +146,14 @@ class TestCollectionBaseProperties:
instance._interval_range(end=datetime.now(tz=tz))
assert "Timezones must match" in str(cm.value)
- res = instance._interval_range(end=datetime.now(tz=timezone.utc))
+ res = instance._interval_range(end=datetime.now(tz=UTC))
assert isinstance(res, pd.IntervalIndex)
assert res.closed_left
assert res.is_non_overlapping_monotonic
assert res.is_monotonic_increasing
assert res.is_unique
- def test_interval_range2(self, mnt_filepath):
+ def test_interval_range2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert isinstance(instance.interval_range, list)
@@ -171,16 +172,16 @@ class TestCollectionBaseProperties:
)
assert len(instance.interval_range) == 2
- def test_progress(self, mnt_filepath):
+ def test_progress(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(
start=AGO_15min, offset="3min", archive_path=mnt_filepath.data_src
)
- x = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_progress2(self, mnt_filepath):
+ def test_progress2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
@@ -189,10 +190,10 @@ class TestCollectionBaseProperties:
assert instance.df.empty
with pytest.raises(expected_exception=NotImplementedError) as cm:
- df = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_items2(self, mnt_filepath):
+ def test_items2(self, mnt_filepath: GRLDatasets):
"""There can't be a test for this because the Items need a path whic
isn't possible in the generic form
"""
@@ -202,7 +203,7 @@ class TestCollectionBaseProperties:
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
- items = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
# item = items[-3]
@@ -213,19 +214,19 @@ class TestCollectionBaseProperties:
# assert str(df.product_id.dtype) == "object"
# assert str(ddf.product_id.dtype) == "string"
- def test_items3(self, mnt_filepath):
+ def test_items3(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
archive_path=mnt_filepath.data_src,
)
with pytest.raises(expected_exception=NotImplementedError) as cm:
- item = instance.items[0]
+ _ = instance.items[0]
assert "Must override" in str(cm.value)
class TestCollectionBaseMethodsCleanup:
- def test_fetch_force_rr_latest(self, mnt_filepath):
+ def test_fetch_force_rr_latest(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=Exception) as cm:
@@ -233,7 +234,7 @@ class TestCollectionBaseMethodsCleanup:
coll.fetch_force_rr_latest(sources=[])
assert "Must override" in str(cm.value)
- def test_fetch_all_paths(self, mnt_filepath):
+ def test_fetch_all_paths(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
@@ -244,19 +245,19 @@ class TestCollectionBaseMethodsCleanup:
assert "Must override" in str(cm.value)
-class TestCollectionBaseMethodsCleanup:
+class TestCollectionBaseMethodsCleanup2:
@pytest.mark.skip
- def test_cleanup_partials(self, mnt_filepath):
+ def test_cleanup_partials(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.cleanup_partials() is None # it doesn't return anything
- def test_clear_tmp_archives(self, mnt_filepath):
+ def test_clear_tmp_archives(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.clear_tmp_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_clear_corrupt_archives(self, mnt_filepath):
+ def test_clear_corrupt_archives(self, mnt_filepath: GRLDatasets):
"""TODO: expand this so it actually has corrupt archives that we
check to see if they're removed
"""
@@ -264,14 +265,14 @@ class TestCollectionBaseMethodsCleanup:
assert instance.clear_corrupt_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_rebuild_symlinks(self, mnt_filepath):
+ def test_rebuild_symlinks(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.rebuild_symlinks() is None
class TestCollectionBaseMethodsSourceTiming:
- def test_get_item(self, mnt_filepath):
+ def test_get_item(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
i = pd.Interval(left=1, right=2, closed="left")
@@ -279,40 +280,40 @@ class TestCollectionBaseMethodsSourceTiming:
instance.get_item(interval=i)
assert "Must override" in str(cm.value)
- def test_get_item_start(self, mnt_filepath):
+ def test_get_item_start(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
start = pd.Timestamp(dt)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_item_start(start=start)
assert "Must override" in str(cm.value)
- def test_get_items(self, mnt_filepath):
+ def test_get_items(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items(since=dt)
assert "Must override" in str(cm.value)
- def test_get_items_from_year(self, mnt_filepath):
+ def test_get_items_from_year(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_from_year(year=2020)
assert "Must override" in str(cm.value)
- def test_get_items_last90(self, mnt_filepath):
+ def test_get_items_last90(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_last90()
assert "Must override" in str(cm.value)
- def test_get_items_last365(self, mnt_filepath):
+ def test_get_items_last365(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py
index e5d1d02..b9f1c26 100644
--- a/tests/incite/test_collection_base_item.py
+++ b/tests/incite/test_collection_base_item.py
@@ -1,6 +1,9 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from os.path import join as pjoin
from pathlib import Path
+from typing import TYPE_CHECKING
from uuid import uuid4
import dask.dataframe as dd
@@ -10,10 +13,13 @@ from pydantic import ValidationError
from generalresearch.incite.base import CollectionItemBase
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+
class TestCollectionItemBase:
def test_init(self):
- dt = datetime.now(tz=timezone.utc).replace(microsecond=0)
+ dt = datetime.now(tz=UTC).replace(microsecond=0)
instance = CollectionItemBase()
instance2 = CollectionItemBase(start=dt)
@@ -25,7 +31,7 @@ class TestCollectionItemBase:
assert 0 == instance.start.microsecond == instance2.start.microsecond
def test_init_start(self):
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
with pytest.raises(expected_exception=ValidationError) as cm:
CollectionItemBase(start=dt)
@@ -40,20 +46,20 @@ class TestCollectionItemBaseProperties:
def test_finish(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.finish
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.finish
def test_interval(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.interval
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.interval
def test_filename(self):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -61,7 +67,7 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -69,27 +75,27 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.path
def test_partial_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.partial_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.partial_path
def test_empty_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.empty_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.empty_path
class TestCollectionItemBaseMethods:
@@ -106,41 +112,41 @@ class TestCollectionItemBaseMethods:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.tmp_filename()
+ _ = instance.tmp_filename()
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_tmp_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.tmp_path()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.tmp_path()
def test_is_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.is_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.is_empty()
def test_has_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_empty()
def test_has_partial_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_partial_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_partial_archive()
@pytest.mark.parametrize("include_empty", [True, False])
- def test_has_archive(self, include_empty):
+ def test_has_archive(self, include_empty: bool):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_archive(include_empty=include_empty)
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_archive(include_empty=include_empty)
- def test_delete_archive_file(self, mnt_filepath):
+ def test_delete_archive_file(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}.zip"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -155,7 +161,7 @@ class TestCollectionItemBaseMethods:
CollectionItemBase.delete_archive(generic_path=path1)
assert not path1.exists()
- def test_delete_archive_dir(self, mnt_filepath):
+ def test_delete_archive_dir(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -174,20 +180,20 @@ class TestCollectionItemBaseMethods:
def test_should_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.should_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.should_archive()
def test_set_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.set_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.set_empty()
def test_valid_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.valid_archive(generic_path=None, sample=None)
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.valid_archive(generic_path=None, sample=None)
class TestCollectionItemBaseMethodsORM:
@@ -197,11 +203,11 @@ class TestCollectionItemBaseMethodsORM:
pass
@pytest.mark.parametrize("is_partial", [True, False])
- def test_to_archive(self, is_partial):
+ def test_to_archive(self, is_partial: bool):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.to_archive(
+ _ = instance.to_archive(
ddf=dd.from_pandas(data=pd.DataFrame()), is_partial=is_partial
)
assert "Must override" in str(cm.value)
diff --git a/tests/incite/test_grl_flow.py b/tests/incite/test_grl_flow.py
index c632f9a..6aea182 100644
--- a/tests/incite/test_grl_flow.py
+++ b/tests/incite/test_grl_flow.py
@@ -1,15 +1,16 @@
class TestGRLFlow:
def test_init(self, mnt_filepath, thl_web_rr):
+ from generalresearch.incite.collections.thl_web import (
+ LedgerDFCollection,
+ TaskAdjustmentDFCollection,
+ )
from generalresearch.incite.defaults import (
ledger_df_collection,
task_df_collection,
- pop_ledger as plm,
)
-
- from generalresearch.incite.collections.thl_web import (
- LedgerDFCollection,
- TaskAdjustmentDFCollection,
+ from generalresearch.incite.defaults import (
+ pop_ledger as plm,
)
from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py
index ea2bced..04d0bb2 100644
--- a/tests/incite/test_interval_idx.py
+++ b/tests/incite/test_interval_idx.py
@@ -1,12 +1,13 @@
+from datetime import UTC, datetime
+
import pandas as pd
-from datetime import datetime, timezone, timedelta
class TestIntervalIndex:
def test_init(self):
- start = datetime(year=2000, month=1, day=1)
- end = datetime(year=2000, month=1, day=10)
+ start = datetime(year=2000, month=1, day=1, tzinfo=UTC)
+ end = datetime(year=2000, month=1, day=10, tzinfo=UTC)
iv_r: pd.IntervalIndex = pd.interval_range(
start=start, end=end, freq="1d", closed="left"
@@ -17,7 +18,7 @@ class TestIntervalIndex:
# If the offset is longer than the end - start it will not
# error. It will simply have 0 rows.
iv_r: pd.IntervalIndex = pd.interval_range(
- start=start, end=end, freq="30d", closed="left"
+ start=start, end=end, freq="30D", closed="left"
)
assert isinstance(iv_r, pd.IntervalIndex)
assert len(iv_r.to_list()) == 0
diff --git a/tests/managers/gr/test_authentication.py b/tests/managers/gr/test_authentication.py
index 53b6931..1310c79 100644
--- a/tests/managers/gr/test_authentication.py
+++ b/tests/managers/gr/test_authentication.py
@@ -1,120 +1,150 @@
import logging
-from random import randint
+from collections.abc import Callable
from uuid import uuid4
import pytest
-from generalresearch.models.gr.authentication import GRUser
-from test_utils.models.conftest import gr_user
+from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
+from generalresearch.managers.gr.team import TeamManager
+from generalresearch.models.gr.authentication import GRToken, GRUser
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
SSO_ISSUER = ""
class TestGRUserManager:
- def test_create(self, gr_um):
- from generalresearch.models.gr.authentication import GRUser
-
- user: GRUser = gr_um.create_dummy()
- instance = gr_um.get_by_id(user.id)
- assert user.id == instance.id
+ def test_create(self, gr_user: GRUser, gr_user_manager: GRUserManager):
+ instance = gr_user_manager.get_by_id(gr_user.id)
+ assert isinstance(instance, GRUser)
+ assert gr_user.id == instance.id
- instance2 = gr_um.get_by_id(user.id)
- assert user.model_dump_json() == instance2.model_dump_json()
+ instance2 = gr_user_manager.get_by_id(gr_user.id)
+ assert isinstance(instance2, GRUser)
+ assert gr_user.model_dump_json() == instance2.model_dump_json()
- def test_get_by_id(self, gr_user, gr_um):
+ def test_get_by_id(self, gr_user: GRUser, gr_user_manager: GRUserManager):
with pytest.raises(expected_exception=ValueError) as cm:
- gr_um.get_by_id(gr_user_id=999_999_999)
+ gr_user_manager.get_by_id(gr_user_id=999_999_999)
assert "GRUser not found" in str(cm.value)
- instance = gr_um.get_by_id(gr_user_id=gr_user.id)
+ instance = gr_user_manager.get_by_id(gr_user_id=gr_user.id)
+ assert isinstance(instance, GRUser)
assert instance.sub == gr_user.sub
- def test_get_by_sub(self, gr_user, gr_um):
+ def test_get_by_sub(self, gr_user: GRUser, gr_user_manager: GRUserManager):
with pytest.raises(expected_exception=ValueError) as cm:
- gr_um.get_by_sub(sub=uuid4().hex)
+ gr_user_manager.get_by_sub(sub=uuid4().hex)
assert "GRUser not found" in str(cm.value)
- instance = gr_um.get_by_sub(sub=gr_user.sub)
+ instance = gr_user_manager.get_by_sub(sub=gr_user.sub)
+ assert isinstance(instance, GRUser)
assert instance.id == gr_user.id
- def test_get_by_sub_or_create(self, gr_user, gr_um):
+ def test_get_by_sub_or_create(
+ self, gr_user: GRUser, gr_user_manager: GRUserManager
+ ):
sub = f"{uuid4().hex}-{uuid4().hex}"
with pytest.raises(expected_exception=ValueError) as cm:
- gr_um.get_by_sub(sub=sub)
+ gr_user_manager.get_by_sub(sub=sub)
assert "GRUser not found" in str(cm.value)
- instance = gr_um.get_by_sub_or_create(sub=sub)
+ instance = gr_user_manager.get_by_sub_or_create(sub=sub)
assert isinstance(instance, GRUser)
assert instance.sub == sub
- def test_get_all(self, gr_um):
- res1 = gr_um.get_all()
+ def test_get_all(
+ self, gr_user_factory: Callable[..., GRUser], gr_user_manager: GRUserManager
+ ):
+ res1 = gr_user_manager.get_all()
assert isinstance(res1, list)
- gr_um.create_dummy()
- res2 = gr_um.get_all()
+ gr_user_factory(save=True)
+ res2 = gr_user_manager.get_all()
assert len(res1) == len(res2) - 1
- def test_get_by_team(self, gr_um):
- res = gr_um.get_by_team(team_id=999_999_999)
+ def test_get_by_team(self, gr_user_manager: GRUserManager):
+ res = gr_user_manager.get_by_team(team_id=999_999_999)
assert isinstance(res, list)
assert res == []
- def test_list_product_uuids(self, caplog, gr_user, gr_um, thl_web_rr):
+ def test_list_product_uuids(
+ self,
+ caplog,
+ gr_user: GRUser,
+ gr_user_manager: GRUserManager,
+ thl_web_rr: PostgresConfig,
+ ):
with caplog.at_level(logging.WARNING):
- gr_um.list_product_uuids(user=gr_user, thl_pg_config=thl_web_rr)
+ gr_user_manager.list_product_uuids(user=gr_user, thl_pg_config=thl_web_rr)
assert "prefetch not run" in caplog.text
class TestGRTokenManager:
- def test_create(self, gr_user, gr_tm):
- assert gr_tm.create(user_id=gr_user.id) is None
+ def test_create(self, gr_user: GRUser, gr_token_manager: GRTokenManager):
+ assert gr_token_manager.create(user_id=gr_user.id) is None
- token = gr_tm.get_by_user_id(user_id=gr_user.id)
+ token = gr_token_manager.get_by_user_id(user_id=gr_user.id)
+ assert isinstance(token, GRToken)
assert gr_user.id == token.user_id
- def test_get_by_user_id(self, gr_user, gr_tm):
- assert gr_tm.create(user_id=gr_user.id) is None
+ def test_get_by_user_id(self, gr_user: GRUser, gr_token_manager: GRTokenManager):
+ assert gr_token_manager.create(user_id=gr_user.id) is None
- token = gr_tm.get_by_user_id(user_id=gr_user.id)
+ token = gr_token_manager.get_by_user_id(user_id=gr_user.id)
+ assert isinstance(token, GRToken)
assert gr_user.id == token.user_id
- def test_prefetch_user(self, gr_user, gr_tm, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRToken
+ def test_prefetch_user(
+ self,
+ gr_user: GRUser,
+ gr_token_manager: GRTokenManager,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
- gr_tm.create(user_id=gr_user.id)
+ gr_token_manager.create(user_id=gr_user.id)
- token: GRToken = gr_tm.get_by_user_id(user_id=gr_user.id)
+ token: GRToken | None = gr_token_manager.get_by_user_id(user_id=gr_user.id)
+ assert isinstance(token, GRToken)
assert token.user is None
token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config)
assert token.user.id == gr_user.id
- def test_get_by_key(self, gr_user, gr_um, gr_tm):
- gr_tm.create(user_id=gr_user.id)
- token = gr_tm.get_by_user_id(user_id=gr_user.id)
+ def test_get_by_key(
+ self,
+ gr_user: GRUser,
+ gr_token_manager: GRTokenManager,
+ ):
+ gr_token_manager.create(user_id=gr_user.id)
+ token = gr_token_manager.get_by_user_id(user_id=gr_user.id)
+ assert isinstance(token, GRToken)
- instance = gr_tm.get_by_key(api_key=token.key)
+ instance = gr_token_manager.get_by_key(api_key=token.key)
assert token.created == instance.created
# Search for non-existent key
with pytest.raises(expected_exception=Exception) as cm:
- gr_tm.get_by_key(api_key=uuid4().hex)
+ gr_token_manager.get_by_key(api_key=uuid4().hex)
assert "No GRUser with token of " in str(cm.value)
@pytest.mark.skip(reason="no idea how to actually test this...")
- def test_get_by_sso_key(self, gr_user, gr_um, gr_tm, gr_redis_config):
- from generalresearch.models.gr.authentication import GRToken
+ def test_get_by_sso_key(
+ self,
+ gr_team_manager: TeamManager,
+ gr_redis_config: RedisConfig,
+ ):
api_key = "..."
jwks = {
# ...
}
- instance = gr_tm.get_by_key(
+ instance = gr_team_manager.get_by_key(
api_key=api_key,
jwks=jwks,
audience="...",
diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py
index 7eb77f8..022086a 100644
--- a/tests/managers/gr/test_business.py
+++ b/tests/managers/gr/test_business.py
@@ -1,30 +1,52 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from test_utils.models.conftest import business
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+ BusinessBankAccount,
+)
+from generalresearch.models.gr.definitions import TransferMethod
+from generalresearch.models.gr.team import Team
+
+if TYPE_CHECKING:
+ from generalresearch.managers.gr.business import (
+ BusinessAddressManager,
+ BusinessBankAccountManager,
+ BusinessManager,
+ )
+ from generalresearch.managers.gr.team import MembershipManager, TeamManager
+ from generalresearch.models.gr.authentication import GRUser
+ from generalresearch.pg_helper import PostgresConfig
class TestBusinessBankAccountManager:
- def test_init(self, business_bank_account_manager, gr_db):
- assert business_bank_account_manager.pg_config == gr_db
+ def test_init(
+ self,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ gr_db: PostgresConfig,
+ ):
+ assert gr_business_bank_account_manager.pg_config == gr_db
- def test_create(self, business, business_bank_account_manager):
- from generalresearch.models.gr.business import (
- TransferMethod,
- BusinessBankAccount,
- )
+ def test_create(
+ self,
+ gr_business: Business,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ ):
- instance = business_bank_account_manager.create(
- business_id=business.id,
+ instance = gr_business_bank_account_manager.create(
+ business_id=gr_business.id,
uuid=uuid4().hex,
transfer_method=TransferMethod.ACH,
)
assert isinstance(instance, BusinessBankAccount)
assert isinstance(instance.id, int)
- res = business_bank_account_manager.get_by_business_id(
+ res = gr_business_bank_account_manager.get_by_business_id(
business_id=instance.business_id
)
assert isinstance(res, list)
@@ -35,42 +57,49 @@ class TestBusinessBankAccountManager:
class TestBusinessAddressManager:
- def test_create(self, business, business_address_manager):
- from generalresearch.models.gr.business import BusinessAddress
-
- res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id)
+ def test_create(
+ self, gr_business: Business, gr_business_address_manager: BusinessAddressManager
+ ):
+ assert gr_business.id
+ res = gr_business_address_manager.create(
+ uuid=uuid4().hex, business_id=gr_business.id
+ )
assert isinstance(res, BusinessAddress)
assert isinstance(res.id, int)
class TestBusinessManager:
- def test_create(self, business_manager):
- from generalresearch.models.gr.business import Business
+ def test_create(self, gr_business_factory: Callable[..., Business]):
- instance = business_manager.create_dummy()
+ instance = gr_business_factory()
assert isinstance(instance, Business)
assert isinstance(instance.id, int)
- def test_get_or_create(self, business_manager):
+ def test_get_or_create(self, gr_business_manager: BusinessManager):
uuid_key = uuid4().hex
- assert business_manager.get_by_uuid(business_uuid=uuid_key) is None
+ assert gr_business_manager.get_by_uuid(business_uuid=uuid_key) is None
- instance = business_manager.get_or_create(
+ instance = gr_business_manager.get_or_create(
uuid=uuid_key,
name=f"name-{uuid4().hex[:6]}",
)
- res = business_manager.get_by_uuid(business_uuid=uuid_key)
+ res = gr_business_manager.get_by_uuid(business_uuid=uuid_key)
+ assert isinstance(res, Business)
assert res.id == instance.id
- def test_get_all(self, business_manager):
- res1 = business_manager.get_all()
+ def test_get_all(
+ self,
+ gr_business_manager: BusinessManager,
+ gr_business_factory: Callable[..., Business],
+ ):
+ res1 = gr_business_manager.get_all()
assert isinstance(res1, list)
- business_manager.create_dummy()
- res2 = business_manager.get_all()
+ gr_business_factory()
+ res2 = gr_business_manager.get_all()
assert len(res1) == len(res2) - 1
@pytest.mark.skip(reason="TODO")
@@ -78,53 +107,65 @@ class TestBusinessManager:
pass
def test_get_by_user_id(
- self, business_manager, gr_user, team_manager, membership_manager
+ self,
+ gr_business_manager: BusinessManager,
+ gr_user: GRUser,
+ gr_team_manager: TeamManager,
+ gr_membership_manager: MembershipManager,
+ gr_business_factory: Callable[..., Business],
+ gr_team_factory: Callable[..., Team],
):
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
- # Create a Business, but don't add it to anything
- b1 = business_manager.create_dummy()
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ # Create a business: Business, but don't add it to anything
+ b1 = gr_business_factory()
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
# Create a Team, but don't create any Memberships
- t1 = team_manager.create_dummy()
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ t1 = gr_team_factory()
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
# Create a Membership for the gr_user to the Team... but it doesn't
# matter because the Team doesn't have any Business yet
- m1 = membership_manager.create(team=t1, gr_user=gr_user)
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ _ = gr_membership_manager.create(team=t1, gr_user=gr_user)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
# Add the Business to the Team... now the Business should be available
# to the gr_user
- team_manager.add_business(team=t1, business=b1)
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ gr_team_manager.add_business(team=t1, business=b1)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 1
# Add another Business to the Team!
- b2 = business_manager.create_dummy()
- team_manager.add_business(team=t1, business=b2)
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ b2 = gr_business_factory()
+ gr_team_manager.add_business(team=t1, business=b2)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 2
@pytest.mark.skip(reason="TODO")
def test_get_uuids_by_user_id(self):
pass
- def test_get_by_uuid(self, business, business_manager):
- instance = business_manager.get_by_uuid(business_uuid=business.uuid)
- assert business.id == instance.id
+ def test_get_by_uuid(
+ self, gr_business: Business, gr_business_manager: BusinessManager
+ ):
+ instance = gr_business_manager.get_by_uuid(business_uuid=gr_business.uuid)
+ assert isinstance(instance, Business)
+ assert gr_business.id == instance.id
- def test_get_by_id(self, business, business_manager):
- instance = business_manager.get_by_id(business_id=business.id)
- assert business.uuid == instance.uuid
+ def test_get_by_id(
+ self, gr_business: Business, gr_business_manager: BusinessManager
+ ):
+ instance = gr_business_manager.get_by_id(business_id=gr_business.id)
+ assert isinstance(instance, Business)
+ assert gr_business.uuid == instance.uuid
- def test_cache_key(self, business):
- assert "business:" in business.cache_key
+ def test_cache_key(self, gr_business: Business):
+ assert "business:" in gr_business.cache_key
# def test_create_raise_on_duplicate(self):
# b_uuid = uuid4().hex
@@ -133,7 +174,7 @@ class TestBusinessManager:
# business = BusinessManager.create(
# uuid=b_uuid,
# name=f"test-{b_uuid[:6]}")
- # assert isinstance(business, Business)
+ # assert isinstance(gr_business: Business, Business)
#
# # Try to make it again
# with pytest.raises(expected_exception=psycopg.errors.UniqueViolation):
diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py
index 9215da4..878a9ca 100644
--- a/tests/managers/gr/test_team.py
+++ b/tests/managers/gr/test_team.py
@@ -1,105 +1,135 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
-from test_utils.models.conftest import team
+from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.team import Membership, Team
+
+if TYPE_CHECKING:
+ from generalresearch.managers.gr.authentication import GRUserManager
+ from generalresearch.managers.gr.team import MembershipManager, TeamManager
+ from generalresearch.models.gr.authentication import GRUser
+ from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class TestMembershipManager:
- def test_init(self, membership_manager, gr_db):
- assert membership_manager.pg_config == gr_db
+ def test_init(
+ self, gr_membership_manager: MembershipManager, gr_db: PostgresConfig
+ ):
+ assert gr_membership_manager.pg_config == gr_db
class TestTeamManager:
- def test_init(self, team_manager, gr_db):
- assert team_manager.pg_config == gr_db
+ def test_init(self, gr_team_manager: TeamManager, gr_db: PostgresConfig):
+ assert gr_team_manager.pg_config == gr_db
- def test_get_or_create(self, team_manager):
+ def test_get_or_create(self, gr_team_manager: TeamManager):
from generalresearch.models.gr.team import Team
new_uuid = uuid4().hex
- team: Team = team_manager.get_or_create(uuid=new_uuid)
+ team: Team = gr_team_manager.get_or_create(uuid=new_uuid)
assert isinstance(team, Team)
assert isinstance(team.id, int)
assert team.uuid == new_uuid
assert team.name == "< Unknown >"
- def test_get_all(self, team_manager):
- res1 = team_manager.get_all()
+ def test_get_all(
+ self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager
+ ):
+ res1 = gr_team_manager.get_all()
assert isinstance(res1, list)
- team_manager.create_dummy()
- res2 = team_manager.get_all()
+ gr_team_factory()
+ res2 = gr_team_manager.get_all()
assert len(res1) == len(res2) - 1
- def test_create(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_create(
+ self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager
+ ):
- team: Team = team_manager.create_dummy()
+ team: Team = gr_team_factory()
assert isinstance(team, Team)
assert isinstance(team.id, int)
- def test_add_user(self, team, team_manager, gr_um, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Membership
+ def test_add_user(
+ self,
+ gr_team: Team,
+ gr_team_manager: TeamManager,
+ gr_user_manager: GRUserManager,
+ gr_user_factory: Callable[..., GRUser],
+ ):
- user: GRUser = gr_um.create_dummy()
+ user: GRUser = gr_user_factory()
- instance = team_manager.add_user(team=team, gr_user=user)
+ instance = gr_team_manager.add_user(
+ gr_user_manager=gr_user_manager, team=gr_team, gr_user=user
+ )
assert isinstance(instance, Membership)
# assert team.gr_users is None
- team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config)
- assert isinstance(team.gr_users, list)
- assert len(team.gr_users)
- assert team.gr_users == [user]
+ gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager)
+ assert isinstance(gr_team.gr_users, list)
+ assert len(gr_team.gr_users)
+ assert gr_team.gr_users == [user]
- def test_get_by_uuid(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_get_by_uuid(
+ self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager
+ ):
- team: Team = team_manager.create_dummy()
+ team: Team = gr_team_factory()
- instance = team_manager.get_by_uuid(team_uuid=team.uuid)
+ instance = gr_team_manager.get_by_uuid(team_uuid=team.uuid)
+ assert isinstance(instance, Team)
assert team.id == instance.id
- def test_get_by_id(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_get_by_id(
+ self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager
+ ):
- team: Team = team_manager.create_dummy()
+ team: Team = gr_team_factory()
- instance = team_manager.get_by_id(team_id=team.id)
+ instance = gr_team_manager.get_by_id(team_id=team.id)
+ assert isinstance(instance, Team)
assert team.uuid == instance.uuid
- def test_get_by_user(self, team, team_manager, gr_um):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Team
-
- user: GRUser = gr_um.create_dummy()
- team_manager.add_user(team=team, gr_user=user)
+ def test_get_by_user(
+ self,
+ gr_team: Team,
+ gr_user_factory: Callable[..., GRUser],
+ gr_team_manager: TeamManager,
+ gr_user_manager: GRUserManager,
+ ):
+ user: GRUser = gr_user_factory()
+ gr_team_manager.add_user(
+ gr_user_manager=gr_user_manager, team=gr_team, gr_user=user
+ )
- res = team_manager.get_by_user(gr_user=user)
+ res = gr_team_manager.get_by_user(gr_user=user)
assert isinstance(res, list)
assert len(res) == 1
instance = res[0]
assert isinstance(instance, Team)
- assert instance.uuid == team.uuid
+ assert instance.uuid == gr_team.uuid
def test_get_by_user_duplicates(
self,
- gr_user_token,
- gr_user,
- membership,
- product_factory,
- membership_factory,
- team,
- thl_web_rr,
- gr_redis_config,
- gr_db,
+ gr_user: GRUser,
+ product_factory: Callable[..., Product],
+ gr_membership_factory: Callable[..., Membership],
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
):
- product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ product_factory(team=gr_team)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
gr_user.prefetch_teams(
pg_config=gr_db,
diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py
index 4d32dd0..fad0b6b 100644
--- a/tests/managers/leaderboard.py
+++ b/tests/managers/leaderboard.py
@@ -1,8 +1,12 @@
+from __future__ import annotations
+
import os
import time
import zoneinfo
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -10,9 +14,6 @@ import pytest
from generalresearch.managers.leaderboard.manager import LeaderboardManager
from generalresearch.managers.leaderboard.tasks import hit_leaderboards
from generalresearch.models.thl.definitions import Status
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.session import Session
from generalresearch.models.thl.leaderboard import (
LeaderboardCode,
LeaderboardFrequency,
@@ -22,7 +23,13 @@ from generalresearch.models.thl.product import (
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ Product,
)
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.redis_helper import RedisConfig
# random uuid for leaderboard tests
product_id = uuid4().hex
@@ -44,7 +51,9 @@ def session_factory():
def _create_session(
- product_user_id="aaa", country_iso="us", user_payout=Decimal("1.00")
+ product_user_id: str = "aaa",
+ country_iso: str = "us",
+ user_payout: Decimal = Decimal("1.00"),
):
user = User(
product_id=product_id,
@@ -63,7 +72,7 @@ def _create_session(
)
session = Session(
user=user,
- started=datetime(2025, 2, 5, 6, tzinfo=timezone.utc),
+ started=datetime(2025, 2, 5, 6, tzinfo=UTC),
id=1,
country_iso=country_iso,
status=Status.COMPLETE,
@@ -74,59 +83,68 @@ def _create_session(
@pytest.fixture(scope="function")
-def setup_leaderboards(thl_redis):
- complete_count = {
- "aaa": 10,
- "bbb": 6,
- "ccc": 6,
- "ddd": 6,
- "eee": 2,
- "fff": 1,
- "ggg": 1,
- }
- sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100}
- max_payout = sum_payout
- country_iso = "us"
- for freq in [
- LeaderboardFrequency.DAILY,
- LeaderboardFrequency.WEEKLY,
- LeaderboardFrequency.MONTHLY,
- ]:
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.COMPLETE_COUNT,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, complete_count)
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.SUM_PAYOUTS,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, sum_payout)
- m = LeaderboardManager(
- redis_client=thl_redis,
- board_code=LeaderboardCode.LARGEST_PAYOUT,
- freq=freq,
- product_id=product_id,
- country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
- )
- thl_redis.delete(m.key)
- thl_redis.zadd(m.key, max_payout)
+def setup_leaderboards(thl_redis_config: RedisConfig) -> Callable[..., None]:
+ thl_redis = thl_redis_config.create_redis_client()
+
+ def _inner():
+ complete_count = {
+ "aaa": 10,
+ "bbb": 6,
+ "ccc": 6,
+ "ddd": 6,
+ "eee": 2,
+ "fff": 1,
+ "ggg": 1,
+ }
+ sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100}
+ max_payout = sum_payout
+ country_iso = "us"
+ for freq in [
+ LeaderboardFrequency.DAILY,
+ LeaderboardFrequency.WEEKLY,
+ LeaderboardFrequency.MONTHLY,
+ ]:
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.COMPLETE_COUNT,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, complete_count)
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.SUM_PAYOUTS,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, sum_payout)
+ m = LeaderboardManager(
+ redis_client=thl_redis,
+ board_code=LeaderboardCode.LARGEST_PAYOUT,
+ freq=freq,
+ product_id=product_id,
+ country_iso=country_iso,
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
+ )
+ thl_redis.delete(m.key)
+ thl_redis.zadd(m.key, max_payout)
+
+ return _inner
class TestLeaderboards:
+ def test_leaderboard_manager(
+ self, setup_leaderboards: Callable[..., None], thl_redis_config: RedisConfig
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
- def test_leaderboard_manager(self, setup_leaderboards, thl_redis):
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -136,6 +154,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso=country_iso,
+ # This is supposed to not have a timezone. @max don't change it
within_time=datetime(2025, 2, 5, 0, 0, 0),
)
lb = m.get_leaderboard()
@@ -152,7 +171,7 @@ class TestLeaderboards:
999999,
tzinfo=zoneinfo.ZoneInfo(key="America/New_York"),
)
- assert lb.period_start_utc == datetime(2025, 2, 5, 5, tzinfo=timezone.utc)
+ assert lb.period_start_utc == datetime(2025, 2, 5, 5, tzinfo=UTC)
assert lb.row_count == 7
assert lb.rows == [
LeaderboardRow(bpuid="aaa", rank=1, value=10),
@@ -164,7 +183,12 @@ class TestLeaderboards:
LeaderboardRow(bpuid="ggg", rank=6, value=1),
]
- def test_leaderboard_manager_bpuid(self, setup_leaderboards, thl_redis):
+ def test_leaderboard_manager_bpuid(
+ self, setup_leaderboards: Callable[..., None], thl_redis_config: RedisConfig
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -174,7 +198,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso=country_iso,
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(bp_user_id="fff", limit=1)
@@ -191,7 +215,15 @@ class TestLeaderboards:
lb.censor()
assert lb.rows[0].bpuid == "ee*"
- def test_leaderboard_hit(self, setup_leaderboards, session_factory, thl_redis):
+ def test_leaderboard_hit(
+ self,
+ setup_leaderboards: Callable[..., None],
+ session_factory: Callable[..., Session],
+ thl_redis_config: RedisConfig,
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
hit_leaderboards(redis_client=thl_redis, session=session_factory())
for freq in [
@@ -205,7 +237,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 7
@@ -216,7 +248,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 3
@@ -227,15 +259,21 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard(limit=1)
assert lb.row_count == 3
assert lb.rows == [LeaderboardRow(bpuid="aaa", rank=1, value=345 + 100)]
def test_leaderboard_hit_new_row(
- self, setup_leaderboards, session_factory, thl_redis
+ self,
+ setup_leaderboards: Callable[..., None],
+ session_factory: Callable[..., None],
+ thl_redis_config: RedisConfig,
):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
session = session_factory(product_user_id="zzz")
hit_leaderboards(redis_client=thl_redis, session=session)
m = LeaderboardManager(
@@ -244,24 +282,21 @@ class TestLeaderboards:
freq=LeaderboardFrequency.DAILY,
product_id=product_id,
country_iso="us",
- within_time=datetime(2025, 2, 5, 12, 12, 12),
+ within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC),
)
lb = m.get_leaderboard()
assert lb.row_count == 8
assert LeaderboardRow(bpuid="zzz", value=1, rank=6) in lb.rows
- def test_leaderboard_country(self, thl_redis):
+ def test_leaderboard_country(self, thl_redis_config: RedisConfig):
+ thl_redis = thl_redis_config.create_redis_client()
m = LeaderboardManager(
redis_client=thl_redis,
board_code=LeaderboardCode.COMPLETE_COUNT,
freq=LeaderboardFrequency.DAILY,
product_id=product_id,
country_iso="jp",
- within_time=datetime(
- 2025,
- 2,
- 1,
- ),
+ within_time=datetime(2025, 2, 1, tzinfo=UTC),
)
lb = m.get_leaderboard()
assert lb.row_count == 0
@@ -270,5 +305,5 @@ class TestLeaderboards:
)
assert lb.local_start_time == "2025-02-01T00:00:00+09:00"
assert lb.local_end_time == "2025-02-01T23:59:59.999999+09:00"
- assert lb.period_start_utc == datetime(2025, 1, 31, 15, tzinfo=timezone.utc)
+ assert lb.period_start_utc == datetime(2025, 1, 31, 15, tzinfo=UTC)
print(lb.model_dump(mode="json"))
diff --git a/tests/managers/network/__init__.py b/tests/managers/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/tests/managers/network/__init__.py
+++ /dev/null
diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py
deleted file mode 100644
index 5b9a790..0000000
--- a/tests/managers/network/test_label.py
+++ /dev/null
@@ -1,202 +0,0 @@
-import ipaddress
-
-import faker
-import pytest
-from psycopg.errors import UniqueViolation
-from pydantic import ValidationError
-
-from generalresearch.managers.network.label import IPLabelManager
-from generalresearch.models.network.label import (
- IPLabel,
- IPLabelKind,
- IPLabelSource,
- IPLabelMetadata,
-)
-from generalresearch.models.thl.ipinfo import normalize_ip
-
-fake = faker.Faker()
-
-
-@pytest.fixture
-def ip_label(utc_now) -> IPLabel:
- ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False)
- return IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip,
- metadata=IPLabelMetadata(services=["RDP"])
- )
-
-
-def test_model(utc_now):
- ip = fake.ipv4_public()
- lbl = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip,
- )
- assert lbl.ip.prefixlen == 32
- print(f"{lbl.ip=}")
-
- ip = ipaddress.IPv4Network((ip, 24), strict=False)
- lbl = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip,
- )
- print(f"{lbl.ip=}")
-
- with pytest.raises(ValidationError, match="IPv6 network must be /64 or larger"):
- IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=fake.ipv6(),
- )
-
- ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False)
- lbl = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip,
- )
- print(f"{lbl.ip=}")
-
- ip = ipaddress.IPv6Network((ip.network_address, 48), strict=False)
- lbl = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip,
- )
- print(f"{lbl.ip=}")
-
-
-def test_create(iplabel_manager: IPLabelManager, ip_label: IPLabel):
- iplabel_manager.create(ip_label)
-
- with pytest.raises(
- UniqueViolation, match="duplicate key value violates unique constraint"
- ):
- iplabel_manager.create(ip_label)
-
-
-def test_filter(iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago):
- res = iplabel_manager.filter(ips=[ip_label.ip])
- assert len(res) == 0
-
- iplabel_manager.create(ip_label)
- res = iplabel_manager.filter(ips=[ip_label.ip])
- assert len(res) == 1
-
- out = res[0]
- assert out == ip_label
-
- res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago)
- assert len(res) == 1
-
- ip_label2 = ip_label.model_copy()
- ip_label2.ip = fake.ipv4_public()
- iplabel_manager.create(ip_label2)
- res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip])
- assert len(res) == 2
-
-
-def test_filter_network(
- iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago
-):
- print(ip_label)
- ip_label = ip_label.model_copy()
- ip_label.ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False)
-
- iplabel_manager.create(ip_label)
- res = iplabel_manager.filter(ips=[ip_label.ip])
- assert len(res) == 1
-
- out = res[0]
- assert out == ip_label
-
- res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago)
- assert len(res) == 1
-
- ip_label2 = ip_label.model_copy()
- ip_label2.ip = fake.ipv4_public()
- iplabel_manager.create(ip_label2)
- res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip])
- assert len(res) == 2
-
-
-def test_network(iplabel_manager: IPLabelManager, utc_now):
- # This is a fully-specific /128 ipv6 address.
- # e.g. '51b7:b38d:8717:6c5b:cd3e:f5c3:3aba:17d'
- ip = fake.ipv6()
- # Generally, we'd want to annotate the /64 network
- # e.g. '51b7:b38d:8717:6c5b::/64'
- ip_64 = ipaddress.IPv6Network((ip, 64), strict=False)
-
- label = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip_64,
- )
- iplabel_manager.create(label)
-
- # If I query for the /128 directly, I won't find it
- res = iplabel_manager.filter(ips=[ip])
- assert len(res) == 0
-
- # If I query for the /64 network I will
- res = iplabel_manager.filter(ips=[ip_64])
- assert len(res) == 1
-
- # Or, I can query for the /128 ip IN a network
- res = iplabel_manager.filter(ip_in_network=ip)
- assert len(res) == 1
-
-
-def test_label_cidr_and_ipinfo(
- iplabel_manager: IPLabelManager, ip_information_factory, ip_geoname, utc_now
-):
- # We have network_iplabel.ip as a cidr col and
- # thl_ipinformation.ip as a inet col. Make sure we can join appropriately
- ip = fake.ipv6()
- ip_information_factory(ip=ip, geoname=ip_geoname)
- # We normalize for storage into ipinfo table
- ip_norm, prefix = normalize_ip(ip)
-
- # Test with a larger network
- ip_48 = ipaddress.IPv6Network((ip, 48), strict=False)
- print(f"{ip=}")
- print(f"{ip_norm=}")
- print(f"{ip_48=}")
- label = IPLabel(
- label_kind=IPLabelKind.VPN,
- labeled_at=utc_now,
- source=IPLabelSource.INTERNAL_USE,
- provider="GeoNodE",
- created_at=utc_now,
- ip=ip_48,
- )
- iplabel_manager.create(label)
-
- res = iplabel_manager.test_join(ip_norm)
- print(res)
diff --git a/tests/managers/network/test_tool_run.py b/tests/managers/network/test_tool_run.py
deleted file mode 100644
index a815809..0000000
--- a/tests/managers/network/test_tool_run.py
+++ /dev/null
@@ -1,25 +0,0 @@
-def test_create_tool_run_from_nmap_run(nmap_run, toolrun_manager):
-
- toolrun_manager.create_nmap_run(nmap_run)
-
- run_out = toolrun_manager.get_nmap_run(nmap_run.id)
-
- assert nmap_run == run_out
-
-
-def test_create_tool_run_from_rdns_run(rdns_run, toolrun_manager):
-
- toolrun_manager.create_rdns_run(rdns_run)
-
- run_out = toolrun_manager.get_rdns_run(rdns_run.id)
-
- assert rdns_run == run_out
-
-
-def test_create_tool_run_from_mtr_run(mtr_run, toolrun_manager):
-
- toolrun_manager.create_mtr_run(mtr_run)
-
- run_out = toolrun_manager.get_mtr_run(mtr_run.id)
-
- assert mtr_run == run_out
diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py
index a0fab38..5ebd015 100644
--- a/tests/managers/test_events.py
+++ b/tests/managers/test_events.py
@@ -1,60 +1,48 @@
-import random
+from __future__ import annotations
+
+import math
import time
-from datetime import timedelta, datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from functools import partial
-from typing import Optional
+from typing import TYPE_CHECKING
from uuid import uuid4
-import math
import pytest
-from math import floor
from generalresearch.managers.events import EventSubscriber
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.events import (
- MessageKind,
- EventType,
AggregateBySource,
+ EventType,
MaxGaugeBySource,
+ MessageKind,
)
from generalresearch.models.legacy.bucket import Bucket
+from generalresearch.models.thl import Product
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.session import Session, Wall
from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.managers.events import EventManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.redis_helper import RedisConfig
+
# We don't need anything in the db, so not using the db fixtures
@pytest.fixture(scope="function")
-def product_id(product_manager):
+def product_id(product_manager: ProductManager) -> str:
return uuid4().hex
@pytest.fixture(scope="function")
-def user_factory(product_id):
- return partial(create_dummy, product_id=product_id)
-
-
-@pytest.fixture(scope="function")
-def event_subscriber(thl_redis_config, product_id):
+def event_subscriber(thl_redis_config: RedisConfig, product_id: str) -> EventSubscriber:
return EventSubscriber(redis_config=thl_redis_config, product_id=product_id)
-def create_dummy(
- product_id: Optional[str] = None, product_user_id: Optional[str] = None
-) -> User:
- return User(
- product_id=product_id,
- product_user_id=product_user_id or uuid4().hex,
- uuid=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- user_id=random.randint(0, floor(2**32 / 2)),
- )
-
-
class TestActiveUsers:
-
- def test_run_empty(self, event_manager, product_id):
+ def test_run_empty(self, event_manager: EventManager, product_id: str):
res = event_manager.get_user_stats(product_id)
assert res == {
"active_users_last_1h": 0,
@@ -63,7 +51,12 @@ class TestActiveUsers:
"in_progress_users": 0,
}
- def test_run(self, event_manager, product_id, user_factory):
+ def test_run(
+ self,
+ event_manager: EventManager,
+ product_factory,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
@@ -72,7 +65,7 @@ class TestActiveUsers:
event_manager.handle_user(user1)
event_manager.handle_user(user1)
- res = event_manager.get_user_stats(product_id)
+ res = event_manager.get_user_stats(user1.product_id)
assert res == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
@@ -88,21 +81,23 @@ class TestActiveUsers:
}
# Create a 2nd user in another product
- product_id2 = uuid4().hex
- user2: User = user_factory(product_id=product_id2)
+ product2 = product_factory()
+ user2: User = user_factory(product=product2)
+ assert isinstance(user2, User)
+ assert isinstance(user2.created, datetime)
# Change to say user was created >24 hrs ago
user2.created = user2.created - timedelta(hours=25)
event_manager.handle_user(user2)
# And now each have 1 active user
- assert event_manager.get_user_stats(product_id) == {
+ assert event_manager.get_user_stats(user1.product_id) == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
"signups_last_24h": 1,
"in_progress_users": 0,
}
# user2 was created older than 24hrs ago
- assert event_manager.get_user_stats(product_id2) == {
+ assert event_manager.get_user_stats(user2.product_id) == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
"signups_last_24h": 0,
@@ -116,10 +111,16 @@ class TestActiveUsers:
"in_progress_users": 0,
}
- def test_inprogress(self, event_manager, product_id, user_factory):
+ def test_inprogress(
+ self,
+ event_manager: EventManager,
+ user_factory: Callable[..., User],
+ product
+ ):
event_manager.clear_global_user_stats()
- user1: User = user_factory()
- user2: User = user_factory()
+ user1: User = user_factory(product=product)
+ user2: User = user_factory(product=product)
+ product_id = product.id
# No matter how many times we do this, they're only active once
event_manager.mark_user_inprogress(user1)
@@ -139,9 +140,14 @@ class TestActiveUsers:
res = event_manager.get_user_stats(product_id)
assert res["in_progress_users"] == 1
- def test_expiry(self, event_manager, product_id, user_factory):
+ def test_expiry(
+ self,
+ event_manager: EventManager,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
+ product_id = user1.product_id
event_manager.handle_user(user1)
event_manager.mark_user_inprogress(user1)
sec_24hr = timedelta(hours=24).total_seconds()
@@ -166,8 +172,7 @@ class TestActiveUsers:
class TestSessionStats:
-
- def test_run_empty(self, event_manager, product_id):
+ def test_run_empty(self, event_manager: EventManager, product_id: str):
res = event_manager.get_session_stats(product_id)
assert res == {
"session_enters_last_1h": 0,
@@ -186,10 +191,18 @@ class TestSessionStats:
"session_fail_avg_loi_last_24h": None,
}
- def test_run(self, event_manager, product_id, user_factory, utc_now, utc_hour_ago):
+ def test_run(
+ self,
+ event_manager: EventManager,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ utc_now: datetime,
+ utc_hour_ago: datetime,
+ ):
event_manager.clear_global_session_stats()
-
- user: User = user_factory()
+ product = product_factory()
+ product_id = product.id
+ user: User = user_factory(product=product)
session = Session(
country_iso="us",
started=utc_hour_ago + timedelta(minutes=10),
@@ -266,29 +279,29 @@ class TestSessionStats:
field_name = str(field)
assert res == {field_name: "1"}
assert (
- 3600 - 60 < event_manager.redis_client.httl(name, field_name)[0] < 3600 + 60
+ 3600 - 61 < event_manager.redis_client.httl(name, field_name)[0] < 3600 + 60
)
# Second BP, fail
- product_id2 = uuid4().hex
- user2: User = user_factory(product_id=product_id2)
+ product2 = product_factory()
+ user2: User = user_factory(product=product2)
session3 = Session(
country_iso="us",
started=utc_now - timedelta(minutes=1),
user=user2,
)
- event_manager.session_on_enter(session=session3, user=user)
+ event_manager.session_on_enter(session=session3, user=user2)
session3.update(
finished=utc_now,
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
)
- event_manager.session_on_finish(session=session3, user=user)
+ event_manager.session_on_finish(session=session3, user=user2)
avg_loi_complete = (
round(session.elapsed.total_seconds())
+ round(session2.elapsed.total_seconds())
) / 2
- assert event_manager.get_session_stats(product_id) == {
+ assert event_manager.get_global_session_stats() == {
"session_enters_last_1h": 2,
"session_enters_last_24h": 3,
"session_fails_last_1h": 1,
@@ -307,7 +320,7 @@ class TestSessionStats:
class TestTaskStatsManager:
- def test_empty(self, event_manager):
+ def test_empty(self, event_manager: EventManager):
event_manager.clear_task_stats()
assert event_manager.get_task_stats_raw() == {
"live_task_count": AggregateBySource(total=0),
@@ -321,7 +334,7 @@ class TestTaskStatsManager:
assert sm.data.task_created_count_last_24h.total == 0
assert sm.data.live_tasks_max_payout.value is None
- def test(self, event_manager):
+ def test(self, event_manager: EventManager):
event_manager.clear_task_stats()
event_manager.set_source_task_stats(
source=Source.TESTING,
@@ -384,7 +397,7 @@ class TestTaskStatsManager:
"task_created_count_last_24h": AggregateBySource(total=0),
}
event_manager.set_source_task_stats(
- source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal("0")
+ source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal(0)
)
assert event_manager.get_task_stats_raw() == {
"live_task_count": AggregateBySource(
@@ -400,7 +413,7 @@ class TestTaskStatsManager:
event_manager.set_source_task_stats(
source=Source.TESTING,
live_task_count=0,
- live_tasks_max_payout=Decimal("0"),
+ live_tasks_max_payout=Decimal(0),
created_count=10,
)
res = event_manager.get_task_stats_raw()
@@ -414,7 +427,7 @@ class TestTaskStatsManager:
event_manager.set_source_task_stats(
source=Source.TESTING,
live_task_count=0,
- live_tasks_max_payout=Decimal("0"),
+ live_tasks_max_payout=Decimal(0),
created_count=10,
)
res = event_manager.get_task_stats_raw()
@@ -428,7 +441,7 @@ class TestTaskStatsManager:
event_manager.set_source_task_stats(
source=Source.TESTING2,
live_task_count=0,
- live_tasks_max_payout=Decimal("0"),
+ live_tasks_max_payout=Decimal(0),
created_count=1,
)
res = event_manager.get_task_stats_raw()
@@ -444,14 +457,15 @@ class TestTaskStatsManager:
class TestChannelsSubscriptions:
+ @pytest.mark.skip("sits there doing nothing forever? todo")
def test_stats_worker(
self,
- event_manager,
- event_subscriber,
- product_id,
- user_factory,
- utc_hour_ago,
- utc_now,
+ event_manager: EventManager,
+ event_subscriber: EventSubscriber,
+ product_id: str,
+ user_factory: Callable[..., User],
+ utc_hour_ago: datetime,
+ utc_now: datetime,
):
event_manager.clear_stats()
assert event_subscriber.pubsub
@@ -481,7 +495,7 @@ class TestChannelsSubscriptions:
wall = Wall(
req_survey_id="a",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
source=Source.TESTING,
session_id=session.id,
user_id=user.user_id,
@@ -496,8 +510,8 @@ class TestChannelsSubscriptions:
wall.update(
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- finished=datetime.now(tz=timezone.utc),
- cpi=Decimal("1"),
+ finished=datetime.now(tz=UTC),
+ cpi=Decimal(1),
)
event_manager.handle_task_finish(wall, session, user)
msg = event_subscriber.get_next_message()
diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py
index 1a1bae7..6771a0c 100644
--- a/tests/managers/test_lucid.py
+++ b/tests/managers/test_lucid.py
@@ -1,14 +1,21 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.lucid.profiling import get_profiling_library
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
+
qids = ["42", "43", "45", "97", "120", "639", "15297"]
class TestLucidProfiling:
@pytest.mark.skip
- def test_get_library(self, thl_web_rr):
+ def test_get_library(self, thl_web_rr: PostgresConfig):
pks = [(qid, "us", "eng") for qid in qids]
qs = get_profiling_library(thl_web_rr, pks=pks)
assert len(qids) == len(qs)
diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py
index 4a3f699..e74e40b 100644
--- a/tests/managers/test_userpid.py
+++ b/tests/managers/test_userpid.py
@@ -1,11 +1,12 @@
+from __future__ import annotations
+
import pytest
from pydantic import MySQLDsn
-from generalresearch.managers.marketplace.user_pid import UserPidMultiManager
-from generalresearch.sql_helper import SqlHelper
from generalresearch.managers.cint.user_pid import CintUserPidManager
from generalresearch.managers.dynata.user_pid import DynataUserPidManager
from generalresearch.managers.innovate.user_pid import InnovateUserPidManager
+from generalresearch.managers.marketplace.user_pid import UserPidMultiManager
from generalresearch.managers.morning.user_pid import MorningUserPidManager
# from generalresearch.managers.precision import PrecisionUserPidManager
@@ -13,6 +14,7 @@ from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager
from generalresearch.managers.repdata.user_pid import RepdataUserPidManager
from generalresearch.managers.sago.user_pid import SagoUserPidManager
from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager
+from generalresearch.sql_helper import SqlHelper
dsn = ""
diff --git a/tests/managers/thl/test_buyer.py b/tests/managers/thl/test_buyer.py
index 69ea105..0ab2d52 100644
--- a/tests/managers/thl/test_buyer.py
+++ b/tests/managers/thl/test_buyer.py
@@ -1,14 +1,24 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
+from generalresearch.models.definitions import Source
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.buyer import BuyerManager
class TestBuyer:
def test(
self,
- delete_buyers_surveys,
- buyer_manager,
+ delete_buyers_surveys: Callable[..., None],
+ buyer_manager: BuyerManager,
):
+ delete_buyers_surveys()
+
bs = buyer_manager.bulk_get_or_create(source=Source.TESTING, codes=["a", "b"])
assert len(bs) == 2
buyer_a = bs[0]
diff --git a/tests/managers/thl/test_cashout_method.py b/tests/managers/thl/test_cashout_method.py
index ee52188..fc364f2 100644
--- a/tests/managers/thl/test_cashout_method.py
+++ b/tests/managers/thl/test_cashout_method.py
@@ -1,62 +1,75 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import pytest
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailCashoutMethodData,
PaypalCashoutMethodData,
USDeliveryAddress,
)
-from test_utils.managers.cashout_methods import (
- EXAMPLE_TANGO_CASHOUT_METHODS,
-)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.managers.thl.cashout_method import (
+ CashoutMethodManager,
+ )
+ from generalresearch.models.thl.user import User
+ from generalresearch.models.thl.wallet.cashout_method import (
+ CashoutMethod,
+ )
class TestTangoCashoutMethods:
- def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db):
+ def test_create_and_get(
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ setup_cashoutmethod_db: Callable[..., None],
+ example_tango_cashout_methods: list[CashoutMethod],
+ ):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.filter(payout_types=[PayoutType.TANGO])
assert len(res) == 2
- cm = [x for x in res if x.ext_id == "U025035"][0]
- assert EXAMPLE_TANGO_CASHOUT_METHODS[0] == cm
+ cm = next(x for x in res if x.ext_id == "U025035")
+ assert example_tango_cashout_methods[0] == cm
def test_user(
- self, cashout_method_manager, user_with_wallet, setup_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ setup_cashoutmethod_db: Callable[..., None],
):
+ setup_cashoutmethod_db()
+
res = cashout_method_manager.get_cashout_methods(user_with_wallet)
# This user ONLY has the two tango cashout methods, no AMT
assert len(res) == 2
-class TestAMTCashoutMethods:
-
- def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db):
- res = cashout_method_manager.filter(payout_types=[PayoutType.AMT])
- assert len(res) == 2
-
- cm = [x for x in res if x.name == "AMT Assignment"][0]
- assert AMT_ASSIGNMENT_CASHOUT_METHOD == cm
-
- cm = [x for x in res if x.name == "AMT Bonus"][0]
- assert AMT_BONUS_CASHOUT_METHOD == cm
-
- def test_user(
- self, cashout_method_manager, user_with_wallet_amt, setup_cashoutmethod_db
- ):
- res = cashout_method_manager.get_cashout_methods(user_with_wallet_amt)
- # This user has the 2 tango, plus amt bonus & assignment
- assert len(res) == 4
-
class TestUserCashoutMethods:
- def test(self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db):
+ def test(
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
+ ):
delete_cashoutmethod_db()
res = cashout_method_manager.get_cashout_methods(user_with_wallet)
assert len(res) == 0
def test_cash_in_mail(
- self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
):
delete_cashoutmethod_db()
@@ -95,7 +108,10 @@ class TestUserCashoutMethods:
assert len(res) == 2
def test_paypal(
- self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db
+ self,
+ cashout_method_manager: CashoutMethodManager,
+ user_with_wallet: User,
+ delete_cashoutmethod_db: Callable[..., None],
):
delete_cashoutmethod_db()
diff --git a/tests/managers/thl/test_category.py b/tests/managers/thl/test_category.py
index ad0f07b..a2805bc 100644
--- a/tests/managers/thl/test_category.py
+++ b/tests/managers/thl/test_category.py
@@ -1,12 +1,21 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.models.thl.category import Category
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.category import CategoryManager
+ from generalresearch.pg_helper import PostgresConfig
+
class TestCategory:
@pytest.fixture
- def beauty_fitness(self, thl_web_rw):
+ def beauty_fitness(self) -> Category:
return Category(
uuid="12c1e96be82c4642a07a12a90ce6f59e",
@@ -16,72 +25,83 @@ class TestCategory:
)
@pytest.fixture
- def hair_care(self, beauty_fitness):
+ def hair_care(self, beauty_fitness: Category) -> Category:
return Category(
uuid="dd76c4b565d34f198dad3687326503d6",
adwords_vertical_id="146",
label="Hair Care",
- path="/Beauty & Fitness/Hair Care",
+ path=f"{beauty_fitness.path}/Hair Care",
)
@pytest.fixture
- def hair_loss(self, hair_care):
+ def hair_loss(self, hair_care: Category) -> Category:
return Category(
uuid="aacff523c8e246888215611ec3b823c0",
adwords_vertical_id="235",
label="Hair Loss",
- path="/Beauty & Fitness/Hair Care/Hair Loss",
+ path=f"{hair_care.path}/Hair Loss",
)
@pytest.fixture
def category_data(
- self, category_manager, thl_web_rw, beauty_fitness, hair_care, hair_loss
- ):
- cats = [beauty_fitness, hair_care, hair_loss]
- data = [x.model_dump(mode="json") for x in cats]
- # We need the parent pk's to set the parent_id. So insert all without a parent,
- # then pull back all pks and map to the parents as parsed by the parent_path
- query = """
- INSERT INTO marketplace_category
- (uuid, adwords_vertical_id, label, path)
- VALUES
- (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s)
- ON CONFLICT (uuid) DO NOTHING;
- """
- with thl_web_rw.make_connection() as conn:
- with conn.cursor() as c:
- c.executemany(query=query, params_seq=data)
- conn.commit()
-
- res = thl_web_rw.execute_sql_query("SELECT id, path FROM marketplace_category")
- path_id = {x["path"]: x["id"] for x in res}
- data = [
- {"id": path_id[c.path], "parent_id": path_id[c.parent_path]}
- for c in cats
- if c.parent_path
- ]
- query = """
- UPDATE marketplace_category
- SET parent_id = %(parent_id)s
- WHERE id = %(id)s;
- """
- with thl_web_rw.make_connection() as conn:
- with conn.cursor() as c:
- c.executemany(query=query, params_seq=data)
- conn.commit()
-
- category_manager.populate_caches()
+ self,
+ category_manager: CategoryManager,
+ thl_web_rw: PostgresConfig,
+ beauty_fitness: Category,
+ hair_care: Category,
+ hair_loss: Category,
+ ) -> Callable[..., None]:
+
+ def _inner():
+ cats = [beauty_fitness, hair_care, hair_loss]
+ data = [x.model_dump(mode="json") for x in cats]
+ # We need the parent pk's to set the parent_id. So insert all without a parent,
+ # then pull back all pks and map to the parents as parsed by the parent_path
+ query = """
+ INSERT INTO marketplace_category
+ (uuid, adwords_vertical_id, label, path)
+ VALUES
+ (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s)
+ ON CONFLICT (uuid) DO NOTHING;
+ """
+ with thl_web_rw.make_connection() as conn:
+ with conn.cursor() as c:
+ c.executemany(query=query, params_seq=data)
+ conn.commit()
+
+ res = thl_web_rw.execute_sql_query(
+ "SELECT id, path FROM marketplace_category"
+ )
+ path_id = {x["path"]: x["id"] for x in res}
+ data = [
+ {"id": path_id[c.path], "parent_id": path_id[c.parent_path]}
+ for c in cats
+ if c.parent_path
+ ]
+ query = """
+ UPDATE marketplace_category
+ SET parent_id = %(parent_id)s
+ WHERE id = %(id)s;
+ """
+ with thl_web_rw.make_connection() as conn:
+ with conn.cursor() as c:
+ c.executemany(query=query, params_seq=data)
+ conn.commit()
+
+ category_manager.populate_caches()
+
+ return _inner
def test(
self,
- category_data,
- category_manager,
- beauty_fitness,
- hair_care,
- hair_loss,
+ category_data: Callable[..., None],
+ category_manager: CategoryManager,
+ beauty_fitness: Category,
):
+ category_data()
+
# category_manager on init caches the category info. This rarely/never changes so this is fine,
# but now that tests get run on a new db each time, the category_manager is inited before
# the fixtures run. so category_manager's cache needs to be rerun
diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py
index 80a88a5..8aa0780 100644
--- a/tests/managers/thl/test_contest/test_leaderboard.py
+++ b/tests/managers/thl/test_contest/test_leaderboard.py
@@ -1,34 +1,41 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
+from typing import TYPE_CHECKING
from zoneinfo import ZoneInfo
from generalresearch.currency import USDCent
from generalresearch.models.thl.contest.definitions import (
- ContestStatus,
ContestEndReason,
+ ContestStatus,
)
from generalresearch.models.thl.contest.leaderboard import (
LeaderboardContest,
- LeaderboardContestCreate,
-)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
-from test_utils.managers.contest.conftest import (
- leaderboard_contest_in_db as contest_in_db,
- leaderboard_contest_create as contest_create,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.contest_manager import ContestManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
+ from generalresearch.models.thl.contest.leaderboard import (
+ LeaderboardContestCreate,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.redis_helper import RedisConfig
-class TestLeaderboardContestCRUD:
+class TestLeaderboardContestCRUD:
def test_create(
self,
- contest_create: LeaderboardContestCreate,
+ leaderboard_contest_create: LeaderboardContestCreate,
product_user_wallet_yes: Product,
- thl_lm,
- contest_manager,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
c = contest_manager.create(
- product_id=product_user_wallet_yes.uuid, contest_create=contest_create
+ product_id=product_user_wallet_yes.uuid,
+ contest_create=leaderboard_contest_create,
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -39,18 +46,19 @@ class TestLeaderboardContestCRUD:
# We have it set in the fixture as the daily contest for 2025-01-01
assert c.end_condition.ends_at == datetime(
2025, 1, 1, 23, 59, 59, 999999, tzinfo=ZoneInfo("America/New_York")
- ).astimezone(tz=timezone.utc) + timedelta(minutes=90)
+ ).astimezone(tz=UTC) + timedelta(minutes=90)
def test_enter(
self,
user_with_wallet: User,
- contest_in_db: LeaderboardContest,
- thl_lm,
- contest_manager,
- user_manager,
- thl_redis,
+ leaderboard_contest_in_db: LeaderboardContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
+ user_manager: UserManager,
+ thl_redis_config: RedisConfig,
):
- contest = contest_in_db
+ thl_redis = thl_redis_config.create_redis_client()
+ contest = leaderboard_contest_in_db
user = user_with_wallet
c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid)
@@ -77,14 +85,15 @@ class TestLeaderboardContestCRUD:
def test_contest_ends(
self,
user_with_wallet: User,
- contest_in_db: LeaderboardContest,
- thl_lm,
- contest_manager,
- user_manager,
- thl_redis,
+ leaderboard_contest_in_db: LeaderboardContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
+ user_manager: UserManager,
+ thl_redis_config: RedisConfig,
):
+ thl_redis = thl_redis_config.create_redis_client()
# The contest should be over. We need to trigger it.
- contest = contest_in_db
+ contest = leaderboard_contest_in_db
contest._redis_client = thl_redis
contest._user_manager = user_manager
user = user_with_wallet
@@ -100,18 +109,22 @@ class TestLeaderboardContestCRUD:
)
assert c.user_rank == 1
- bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid(user.product_id)
- bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet)
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
+ user.product_id
+ )
+ bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet)
assert bp_wallet_balance == 0
- user_wallet = thl_lm.get_account_or_create_user_wallet(user=user)
- user_balance = thl_lm.get_account_balance(user_wallet)
+ user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
+ user_balance = thl_ledger_manager.get_account_balance(user_wallet)
assert user_balance == 0
decision, reason = contest.should_end()
assert decision
assert reason == ContestEndReason.ENDS_AT
- contest_manager.end_contest_if_over(contest=contest, ledger_manager=thl_lm)
+ contest_manager.end_contest_if_over(
+ contest=contest, ledger_manager=thl_ledger_manager
+ )
c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid)
assert c.status == ContestStatus.COMPLETED
@@ -129,10 +142,12 @@ class TestLeaderboardContestCRUD:
assert w.prize.cash_amount == USDCent(15_00)
# The prize is $15.00, so the user should get $15, paid by the bp
- assert thl_lm.get_account_balance(account=user_wallet) == 15_00
+ assert thl_ledger_manager.get_account_balance(account=user_wallet) == 15_00
# contest wallet is 0, and the BP gets 20c
- contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid(
- contest_uuid=c.uuid
+ contest_wallet = (
+ thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid(
+ contest_uuid=c.uuid
+ )
)
- assert thl_lm.get_account_balance(account=contest_wallet) == 0
- assert thl_lm.get_account_balance(account=bp_wallet) == -15_00
+ assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 0
+ assert thl_ledger_manager.get_account_balance(account=bp_wallet) == -15_00
diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py
index 7312a64..f26819b 100644
--- a/tests/managers/thl/test_contest/test_milestone.py
+++ b/tests/managers/thl/test_contest/test_milestone.py
@@ -1,34 +1,41 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
from generalresearch.models.thl.contest.definitions import (
- ContestStatus,
ContestEndReason,
+ ContestEntryTrigger,
+ ContestStatus,
)
from generalresearch.models.thl.contest.milestone import (
MilestoneContest,
- MilestoneContestCreate,
MilestoneUserView,
- ContestEntryTrigger,
-)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
-from test_utils.managers.contest.conftest import (
- milestone_contest as contest,
- milestone_contest_in_db as contest_in_db,
- milestone_contest_create as contest_create,
- milestone_contest_factory as contest_factory,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.contest_manager import ContestManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.contest.milestone import (
+ MilestoneContestCreate,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
-class TestMilestoneContest:
- def test_should_end(self, contest: MilestoneContest, thl_lm, contest_manager):
+class TestMilestoneContest:
+ def test_should_end(
+ self,
+ milestone_contest: MilestoneContest,
+ ):
+ contest = milestone_contest
# contest is active and has no entries
should, msg = contest.should_end()
assert not should, msg
# Change so that the contest ends now
- contest.end_condition.ends_at = datetime.now(tz=timezone.utc)
+ contest.end_condition.ends_at = datetime.now(tz=UTC)
should, msg = contest.should_end()
assert should
assert msg == ContestEndReason.ENDS_AT
@@ -43,16 +50,15 @@ class TestMilestoneContest:
class TestMilestoneContestCRUD:
-
def test_create(
self,
- contest_create: MilestoneContestCreate,
+ milestone_contest_create: MilestoneContestCreate,
product_user_wallet_yes: Product,
- thl_lm,
- contest_manager,
+ contest_manager: ContestManager,
):
c = contest_manager.create(
- product_id=product_user_wallet_yes.uuid, contest_create=contest_create
+ product_id=product_user_wallet_yes.uuid,
+ contest_create=milestone_contest_create,
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -68,20 +74,20 @@ class TestMilestoneContestCRUD:
def test_enter(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Users CANNOT directly enter a milestone contest through the api,
# but we'll call this manager method when a trigger is hit.
- contest = contest_in_db
+ contest = milestone_contest_in_db
user = user_with_wallet
contest_manager.enter_milestone_contest(
contest_uuid=contest.uuid,
user=user,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=1,
)
@@ -96,17 +102,19 @@ class TestMilestoneContestCRUD:
assert c.user_amount == 1
# Contest wallet should have 0 bc there is no ledger
- contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid(
- contest_uuid=contest.uuid
+ contest_wallet = (
+ thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid(
+ contest_uuid=contest.uuid
+ )
)
- assert thl_lm.get_account_balance(contest_wallet) == 0
+ assert thl_ledger_manager.get_account_balance(contest_wallet) == 0
# Enter again!
contest_manager.enter_milestone_contest(
contest_uuid=contest.uuid,
user=user,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=1,
)
c: MilestoneUserView = contest_manager.get_milestone_user_view(
@@ -122,21 +130,21 @@ class TestMilestoneContestCRUD:
def test_enter_win(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# User enters contest, which brings the USER'S total amount above the limit,
# and the user reaches the milestone
- contest = contest_in_db
+ contest = milestone_contest_in_db
user = user_with_wallet
- user_wallet = thl_lm.get_account_or_create_user_wallet(user=user)
- user_balance = thl_lm.get_account_balance(account=user_wallet)
- bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid(
+ user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
+ user_balance = thl_ledger_manager.get_account_balance(account=user_wallet)
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
product_uuid=user.product_id
)
- bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet)
+ bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet)
c: MilestoneUserView = contest_manager.get_milestone_user_view(
contest_uuid=contest.uuid, user=user_with_wallet
@@ -151,7 +159,7 @@ class TestMilestoneContestCRUD:
contest_uuid=contest.uuid,
user=user,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=1,
)
@@ -171,9 +179,12 @@ class TestMilestoneContestCRUD:
assert c.win_count == 1
# The prize was awarded! User should have won $1.00
- assert thl_lm.get_account_balance(user_wallet) - user_balance == 100
+ assert thl_ledger_manager.get_account_balance(user_wallet) - user_balance == 100
# Which was paid from the BP's balance
- assert thl_lm.get_account_balance(bp_wallet) - bp_wallet_balance == -100
+ assert (
+ thl_ledger_manager.get_account_balance(bp_wallet) - bp_wallet_balance
+ == -100
+ )
# winnings = cm.get_winnings_by_user(user=user)
# assert len(winnings) == 1
@@ -182,22 +193,22 @@ class TestMilestoneContestCRUD:
def test_enter_ends(
self,
- user_factory,
+ user_factory: Callable[..., User],
product_user_wallet_yes: Product,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Multiple users reach the milestone. Contest ends after 5 wins.
users = [user_factory(product=product_user_wallet_yes) for _ in range(5)]
- contest = contest_in_db
+ contest = milestone_contest_in_db
for u in users:
contest_manager.enter_milestone_contest(
contest_uuid=contest.uuid,
user=u,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=3,
)
@@ -208,29 +219,33 @@ class TestMilestoneContestCRUD:
def test_trigger(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Pretend user just got a complete
cnt = contest_manager.hit_milestone_triggers(
country_iso="us",
user=user_with_wallet,
event=ContestEntryTrigger.TASK_COMPLETE,
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert cnt == 1
# Assert this contest got entered
c: MilestoneUserView = contest_manager.get_milestone_user_view(
- contest_uuid=contest_in_db.uuid, user=user_with_wallet
+ contest_uuid=milestone_contest_in_db.uuid, user=user_with_wallet
)
assert c.user_amount == 1
class TestMilestoneContestUserViews:
def test_list_user_eligible_country(
- self, user_with_wallet: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_wallet: User,
+ milestone_contest_factory: Callable[..., MilestoneContest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# No contests exists
cs = contest_manager.get_many_by_user_eligible(
@@ -239,7 +254,7 @@ class TestMilestoneContestUserViews:
assert len(cs) == 0
# Create a contest. It'll be in the US/CA
- contest_factory(country_isos={"us", "ca"})
+ milestone_contest_factory(country_isos={"us", "ca"})
# Not eligible in mexico
cs = contest_manager.get_many_by_user_eligible(
@@ -252,7 +267,7 @@ class TestMilestoneContestUserViews:
assert len(cs) == 1
# Create another, any country
- contest_factory(country_isos=None)
+ milestone_contest_factory(country_isos=None)
cs = contest_manager.get_many_by_user_eligible(
user=user_with_wallet, country_iso="mx"
)
@@ -263,10 +278,14 @@ class TestMilestoneContestUserViews:
assert len(cs) == 2
def test_list_user_eligible(
- self, user_with_money: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_money: User,
+ milestone_contest_factory: Callable[..., MilestoneContest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# User reaches milestone after 1 complete
- c = contest_factory(target_amount=1)
+ c = milestone_contest_factory(target_amount=1)
user = user_with_money
cs = contest_manager.get_many_by_user_eligible(
@@ -275,7 +294,10 @@ class TestMilestoneContestUserViews:
assert len(cs) == 1
contest_manager.enter_milestone_contest(
- contest_uuid=c.uuid, user=user, country_iso="us", ledger_manager=thl_lm
+ contest_uuid=c.uuid,
+ user=user,
+ country_iso="us",
+ ledger_manager=thl_ledger_manager,
)
# User isn't eligible anymore
diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py
index 060055a..7388991 100644
--- a/tests/managers/thl/test_contest/test_raffle.py
+++ b/tests/managers/thl/test_contest/test_raffle.py
@@ -1,4 +1,8 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
import pytest
from pydantic import ValidationError
@@ -9,44 +13,49 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerTransactionConditionFailedError,
)
from generalresearch.models.thl.contest import (
- ContestPrize,
- ContestEntryRule,
ContestEndCondition,
+ ContestEntryRule,
+ ContestPrize,
)
-from generalresearch.models.thl.contest.definitions import (
- ContestStatus,
- ContestPrizeKind,
- ContestEndReason,
-)
-from generalresearch.models.thl.contest.exceptions import ContestError
-from generalresearch.models.thl.contest.raffle import (
+from generalresearch.models.thl.contest.contest_entry import (
ContestEntry,
ContestEntryType,
)
-from generalresearch.models.thl.contest.raffle import (
- RaffleContest,
- RaffleContestCreate,
- RaffleUserView,
-)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
-from test_utils.managers.contest.conftest import (
- raffle_contest as contest,
- raffle_contest_in_db as contest_in_db,
- raffle_contest_create as contest_create,
- raffle_contest_factory as contest_factory,
+from generalresearch.models.thl.contest.definitions import (
+ ContestEndReason,
+ ContestPrizeKind,
+ ContestStatus,
)
+from generalresearch.models.thl.contest.exceptions import ContestError
+from generalresearch.models.thl.contest.raffle import RaffleContest
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.contest_manager import ContestManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.contest import (
+ Contest,
+ )
+ from generalresearch.models.thl.contest.raffle import (
+ RaffleContestCreate,
+ RaffleUserView,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
class TestRaffleContest:
- def test_should_end(self, contest: RaffleContest, thl_lm, contest_manager):
+ def test_should_end(
+ self,
+ raffle_contest: RaffleContest,
+ ):
+ contest = raffle_contest
# contest is active and has no entries
should, msg = contest.should_end()
assert not should, msg
# Change so that the contest ends now
- contest.end_condition.ends_at = datetime.now(tz=timezone.utc)
+ contest.end_condition.ends_at = datetime.now(tz=UTC)
should, msg = contest.should_end()
assert should
assert msg == ContestEndReason.ENDS_AT
@@ -63,13 +72,12 @@ class TestRaffleContestCRUD:
def test_create(
self,
- contest_create: RaffleContestCreate,
+ raffle_contest_create: RaffleContestCreate,
product_user_wallet_yes: Product,
- thl_lm,
- contest_manager,
+ contest_manager: ContestManager,
):
c = contest_manager.create(
- product_id=product_user_wallet_yes.uuid, contest_create=contest_create
+ product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -85,18 +93,20 @@ class TestRaffleContestCRUD:
def test_enter(
self,
user_with_money: User,
- contest_in_db: RaffleContest,
- thl_lm,
- contest_manager,
+ raffle_contest_in_db: RaffleContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Raffle ends at $1.00. User enters for $0.60
print(user_with_money.product_id)
- print(contest_in_db.product_id)
- print(contest_in_db.uuid)
- contest = contest_in_db
+ print(raffle_contest_in_db.product_id)
+ print(raffle_contest_in_db.uuid)
+ contest = raffle_contest_in_db
- user_wallet = thl_lm.get_account_or_create_user_wallet(user=user_with_money)
- user_balance = thl_lm.get_account_balance(account=user_wallet)
+ user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(
+ user=user_with_money
+ )
+ user_balance = thl_ledger_manager.get_account_balance(account=user_wallet)
entry = ContestEntry(
entry_type=ContestEntryType.CASH, user=user_with_money, amount=USDCent(60)
@@ -105,7 +115,7 @@ class TestRaffleContestCRUD:
contest_uuid=contest.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
c: RaffleContest = contest_manager.get(contest_uuid=contest.uuid)
assert c.current_amount == USDCent(60)
@@ -120,30 +130,35 @@ class TestRaffleContestCRUD:
assert c.projected_win_probability == approx(60 / 100, rel=0.01)
# Contest wallet should have $0.60
- contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid(
- contest_uuid=contest.uuid
+ contest_wallet = (
+ thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid(
+ contest_uuid=contest.uuid
+ )
)
- assert thl_lm.get_account_balance(account=contest_wallet) == 60
+ assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 60
# User spent 60c
- assert user_balance - thl_lm.get_account_balance(account=user_wallet) == 60
+ assert (
+ user_balance - thl_ledger_manager.get_account_balance(account=user_wallet)
+ == 60
+ )
@pytest.mark.parametrize("user_with_money", [{"min_balance": 120}], indirect=True)
def test_enter_ends(
self,
user_with_money: User,
- contest_in_db: RaffleContest,
- thl_lm,
- contest_manager,
+ raffle_contest_in_db: RaffleContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# User enters contest, which brings the total amount above the limit,
# and the contest should end, with a winner selected
- contest = contest_in_db
+ contest = raffle_contest_in_db
- bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid(
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
user_with_money.product_id
)
# I bribed the user, so the balance is not 0
- bp_wallet_balance = thl_lm.get_account_balance(account=bp_wallet)
+ bp_wallet_balance = thl_ledger_manager.get_account_balance(account=bp_wallet)
for _ in range(2):
entry = ContestEntry(
@@ -155,7 +170,7 @@ class TestRaffleContestCRUD:
contest_uuid=contest.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
c: RaffleContest = contest_manager.get(contest_uuid=contest.uuid)
assert c.status == ContestStatus.COMPLETED
@@ -175,25 +190,33 @@ class TestRaffleContestCRUD:
assert win.product_user_id == user_with_money.product_user_id
# Contest wallet should have gotten zeroed out
- contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid(
- contest_uuid=contest.uuid
+ contest_wallet = (
+ thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid(
+ contest_uuid=contest.uuid
+ )
)
- assert thl_lm.get_account_balance(contest_wallet) == 0
+ assert thl_ledger_manager.get_account_balance(contest_wallet) == 0
# Expense wallet gets the $1.00 expense
- expense_wallet = thl_lm.get_account_or_create_bp_expense_by_uuid(
+ expense_wallet = thl_ledger_manager.get_account_or_create_bp_expense_by_uuid(
product_uuid=user_with_money.product_id, expense_name="Prize"
)
- assert thl_lm.get_account_balance(expense_wallet) == -100
+ assert thl_ledger_manager.get_account_balance(expense_wallet) == -100
# And the BP gets 20c
- assert thl_lm.get_account_balance(bp_wallet) - bp_wallet_balance == 20
+ assert (
+ thl_ledger_manager.get_account_balance(bp_wallet) - bp_wallet_balance == 20
+ )
@pytest.mark.parametrize("user_with_money", [{"min_balance": 120}], indirect=True)
def test_enter_ends_cash_prize(
- self, user_with_money: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_money: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Same as test_enter_ends, but the prize is cash. Just
# testing the ledger methods
- c = contest_factory(
+ c = raffle_contest_factory(
prizes=[
ContestPrize(
name="$1.00 bonus",
@@ -205,12 +228,14 @@ class TestRaffleContestCRUD:
)
assert c.prizes[0].kind == ContestPrizeKind.CASH
- user_wallet = thl_lm.get_account_or_create_user_wallet(user=user_with_money)
- user_balance = thl_lm.get_account_balance(user_wallet)
- bp_wallet = thl_lm.get_account_or_create_bp_wallet_by_uuid(
+ user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(
+ user=user_with_money
+ )
+ user_balance = thl_ledger_manager.get_account_balance(user_wallet)
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
user_with_money.product_id
)
- bp_wallet_balance = thl_lm.get_account_balance(bp_wallet)
+ bp_wallet_balance = thl_ledger_manager.get_account_balance(bp_wallet)
## Enter Contest
entry = ContestEntry(
@@ -220,28 +245,35 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# The prize is $1.00, so the user spent $1.20 entering, won, then got $1.00 back
assert (
- thl_lm.get_account_balance(account=user_wallet) == user_balance + 100 - 120
+ thl_ledger_manager.get_account_balance(account=user_wallet)
+ == user_balance + 100 - 120
)
# contest wallet is 0, and the BP gets 20c
- contest_wallet = thl_lm.get_account_or_create_contest_wallet_by_uuid(
- contest_uuid=c.uuid
+ contest_wallet = (
+ thl_ledger_manager.get_account_or_create_contest_wallet_by_uuid(
+ contest_uuid=c.uuid
+ )
+ )
+ assert thl_ledger_manager.get_account_balance(account=contest_wallet) == 0
+ assert (
+ thl_ledger_manager.get_account_balance(account=bp_wallet)
+ - bp_wallet_balance
+ == 20
)
- assert thl_lm.get_account_balance(account=contest_wallet) == 0
- assert thl_lm.get_account_balance(account=bp_wallet) - bp_wallet_balance == 20
def test_enter_failure(
self,
user_with_wallet: User,
- contest_in_db: RaffleContest,
- thl_lm,
- contest_manager,
+ raffle_contest_in_db: RaffleContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_in_db
+ c = raffle_contest_in_db
user = user_with_wallet
# Tries to enter $0
@@ -260,7 +292,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert e.value.args[0] == "insufficient balance"
@@ -271,16 +303,20 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "incompatible entry type" in str(e.value)
@pytest.mark.parametrize("user_with_money", [{"min_balance": 100}], indirect=True)
def test_enter_not_eligible(
- self, user_with_money: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_money: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Max entry amount per user $0.10. Contest still ends at $1.00
- c = contest_factory(
+ c = raffle_contest_factory(
entry_rule=ContestEntryRule(
max_entry_amount_per_user=USDCent(10),
max_daily_entries_per_user=USDCent(8),
@@ -299,7 +335,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user." in str(e.value)
@@ -312,7 +348,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
@@ -324,7 +360,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# Then can't anymore
@@ -336,14 +372,18 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
class TestRaffleContestUserViews:
def test_list_user_eligible_country(
- self, user_with_wallet: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_wallet: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# No contests exists
cs = contest_manager.get_many_by_user_eligible(
@@ -352,7 +392,7 @@ class TestRaffleContestUserViews:
assert len(cs) == 0
# Create a contest. It'll be in the US/CA
- contest_factory(country_isos={"us", "ca"})
+ raffle_contest_factory(country_isos={"us", "ca"})
# Not eligible in mexico
cs = contest_manager.get_many_by_user_eligible(
@@ -365,7 +405,7 @@ class TestRaffleContestUserViews:
assert len(cs) == 1
# Create another, any country
- contest_factory(country_isos=None)
+ raffle_contest_factory(country_isos=None)
cs = contest_manager.get_many_by_user_eligible(
user=user_with_wallet, country_iso="mx"
)
@@ -376,9 +416,13 @@ class TestRaffleContestUserViews:
assert len(cs) == 2
def test_list_user_eligible(
- self, user_with_money: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_money: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(
+ c = raffle_contest_factory(
end_condition=ContestEndCondition(target_entry_amount=USDCent(10)),
entry_rule=ContestEntryRule(
max_entry_amount_per_user=USDCent(1),
@@ -398,7 +442,7 @@ class TestRaffleContestUserViews:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# User isn't eligible anymore
@@ -422,9 +466,13 @@ class TestRaffleContestUserViews:
assert len(contest_manager.get_winnings_by_user(user_with_money)) == 0
def test_list_user_winnings(
- self, user_with_money: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_money: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(
+ c = raffle_contest_factory(
end_condition=ContestEndCondition(target_entry_amount=USDCent(100)),
)
entry = ContestEntry(
@@ -436,7 +484,7 @@ class TestRaffleContestUserViews:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# Contest ends after 100 entry, user enters 100 entry, user wins!
ws = contest_manager.get_winnings_by_user(user_with_money)
@@ -458,9 +506,13 @@ class TestRaffleContestCRUDCount:
# This is a COUNT contest. No cash moves. Not really fleshed out what we'd do with this.
@pytest.mark.skip
def test_enter(
- self, user_with_wallet: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_wallet: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(entry_type=ContestEntryType.COUNT)
+ c = raffle_contest_factory(entry_type=ContestEntryType.COUNT)
entry = ContestEntry(
entry_type=ContestEntryType.COUNT,
user=user_with_wallet,
@@ -470,5 +522,5 @@ class TestRaffleContestCRUDCount:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py
index 6bbbbe1..2fc0ff0 100644
--- a/tests/managers/thl/test_harmonized_uqa.py
+++ b/tests/managers/thl/test_harmonized_uqa.py
@@ -1,13 +1,18 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.managers.thl.profiling.uqa import UQAManager
from generalresearch.models.thl.profiling.user_question_answer import (
- UserQuestionAnswer,
DUMMY_UQA,
+ UserQuestionAnswer,
)
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.uqa import UQAManager
+ from generalresearch.models.thl.user import User
@pytest.mark.usefixtures("uqa_db_index", "upk_data", "uqa_manager_clear_cache")
@@ -18,7 +23,7 @@ class TestUQAManager:
assert len(uqas) == 0
def test_create(self, uqa_manager: UQAManager, user: User):
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
uqas = [
UserQuestionAnswer(
user_id=user.user_id,
@@ -38,7 +43,7 @@ class TestUQAManager:
assert res[0] == uqas[0]
# Same question, so this gets updated
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
uqas_update = [
UserQuestionAnswer(
user_id=user.user_id,
@@ -57,7 +62,7 @@ class TestUQAManager:
assert res[0] == uqas_update[0]
# Add a new answer
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
uqas_new = [
UserQuestionAnswer(
user_id=user.user_id,
@@ -103,7 +108,7 @@ class TestUQAManagerCache:
UserQuestionAnswer(
question_id="5d6d9f3c03bb40bf9d0a24f306387d7c",
answer=("1",),
- timestamp=datetime.now(tz=timezone.utc),
+ timestamp=datetime.now(tz=UTC),
country_iso="us",
language_iso="eng",
property_code="gr:gender",
diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py
index 847b00c..c021eb9 100644
--- a/tests/managers/thl/test_ipinfo.py
+++ b/tests/managers/thl/test_ipinfo.py
@@ -1,51 +1,75 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import faker
from generalresearch.managers.thl.ipinfo import (
+ GeoIpInfoManager,
IPGeonameManager,
IPInformationManager,
- GeoIpInfoManager,
)
-from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation
+from generalresearch.models.thl.ipinfo import (
+ GeoIPInformation,
+ IPGeoname,
+ IPInformation,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
fake = faker.Faker()
class TestIPGeonameManager:
- def test_init(self, thl_web_rr, ip_geoname_manager: IPGeonameManager):
+ def test_init(
+ self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager
+ ):
instance = IPGeonameManager(pg_config=thl_web_rr)
assert isinstance(instance, IPGeonameManager)
assert isinstance(ip_geoname_manager, IPGeonameManager)
- def test_create(self, ip_geoname_manager: IPGeonameManager):
-
- instance = ip_geoname_manager.create_dummy()
+ def test_create(
+ self,
+ ip_geoname_factory: Callable[..., IPGeoname],
+ ip_geoname_manager: IPGeonameManager,
+ ):
+ instance = ip_geoname_factory()
assert isinstance(instance, IPGeoname)
res = ip_geoname_manager.fetch_geoname_ids(filter_ids=[instance.geoname_id])
-
assert res[0].model_dump_json() == instance.model_dump_json()
class TestIPInformationManager:
- def test_init(self, thl_web_rr, ip_information_manager: IPInformationManager):
+ def test_init(
+ self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager
+ ):
instance = IPInformationManager(pg_config=thl_web_rr)
assert isinstance(instance, IPInformationManager)
assert isinstance(ip_information_manager, IPInformationManager)
- def test_create(self, ip_information_manager: IPInformationManager):
- instance = ip_information_manager.create_dummy()
-
+ def test_create(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_information_manager: IPInformationManager,
+ ):
+ instance = ip_information_factory()
assert isinstance(instance, IPInformation)
res = ip_information_manager.fetch_ip_information(filter_ips=[instance.ip])
-
assert res[0].model_dump_json() == instance.model_dump_json()
- def test_prefetch_geoname(self, ip_information, ip_geoname, thl_web_rr):
+ def test_prefetch_geoname(
+ self,
+ ip_information: IPInformation,
+ ip_geoname: IPGeoname,
+ thl_web_rr: PostgresConfig,
+ ):
assert isinstance(ip_information, IPInformation)
assert ip_information.geoname_id == ip_geoname.geoname_id
@@ -57,13 +81,21 @@ class TestIPInformationManager:
class TestGeoIpInfoManager:
def test_init(
- self, thl_web_rr, thl_redis_config, geoipinfo_manager: GeoIpInfoManager
+ self,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
+ geoipinfo_manager: GeoIpInfoManager,
):
instance = GeoIpInfoManager(pg_config=thl_web_rr, redis_config=thl_redis_config)
assert isinstance(instance, GeoIpInfoManager)
assert isinstance(geoipinfo_manager, GeoIpInfoManager)
- def test_multi(self, ip_information_factory, ip_geoname, geoipinfo_manager):
+ def test_multi(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ geoipinfo_manager: GeoIpInfoManager,
+ ):
ip = fake.ipv4_public()
ip_information_factory(ip=ip, geoname=ip_geoname)
ips = [ip]
@@ -90,7 +122,12 @@ class TestGeoIpInfoManager:
assert res[ip] is not None
assert res[ip2] is not None
- def test_multi_ipv6(self, ip_information_factory, ip_geoname, geoipinfo_manager):
+ def test_multi_ipv6(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ geoipinfo_manager: GeoIpInfoManager,
+ ):
ip = fake.ipv6()
# Make another IP that will be in the same /64 block.
ip2 = ip[:-1] + "a" if ip[-1] != "a" else ip[:-1] + "b"
@@ -105,13 +142,19 @@ class TestGeoIpInfoManager:
# Looks up in redis, if not exists, looks in mysql, then sets
# the caches that didn't exist.
res = geoipinfo_manager.get_multi(ip_addresses=ips)
- assert res[ip].ip == ip
- assert res[ip].lookup_prefix == "/64"
- assert res[ip2].ip == ip2
- assert res[ip2].lookup_prefix == "/64"
+
+ res1 = res[ip]
+ assert isinstance(res1, GeoIPInformation)
+ assert res1.ip == ip
+ assert res1.lookup_prefix == "/64"
+
+ res2 = res[ip2]
+ assert isinstance(res2, GeoIPInformation)
+ assert res2.ip == ip2
+ assert res2.lookup_prefix == "/64"
# they should be the same basically, except for the ip
- def test_doesnt_exist(self, geoipinfo_manager):
+ def test_doesnt_exist(self, geoipinfo_manager: GeoIpInfoManager):
ip = fake.ipv4_public()
res = geoipinfo_manager.get_multi(ip_addresses=[ip])
assert res == {ip: None}
diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py
index 5cfaac1..3af10e7 100644
--- a/tests/managers/thl/test_ledger/test_lm_accounts.py
+++ b/tests/managers/thl/test_ledger/test_lm_accounts.py
@@ -1,9 +1,12 @@
+from __future__ import annotations
+
from itertools import product as iproduct
from random import randint
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from pydantic import PositiveInt
from generalresearch.currency import LedgerCurrency
from generalresearch.managers.base import Permission
@@ -11,6 +14,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerAccountDoesntExistError,
)
from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.ledger import (
AccountType,
Direction,
@@ -19,22 +23,10 @@ from generalresearch.models.thl.ledger import (
)
if TYPE_CHECKING:
- from pydantic import PositiveInt
- from generalresearch.config import GRLSettings
- from generalresearch.currency import LedgerCurrency
- from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
- from generalresearch.models.custom_types import AccountType, Direction, UUIDStr
- from generalresearch.models.thl import Direction
from generalresearch.models.thl.ledger import (
- AccountType,
- LedgerAccount,
LedgerTransaction,
)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.session import Session
- from generalresearch.models.thl.user import User
- from generalresearch.models.thl.wallet import PayoutType
@pytest.mark.parametrize(
@@ -51,53 +43,63 @@ class TestLedgerAccountManagerNoResults:
def test_get_account_no_results(
self,
- currency: "LedgerCurrency",
+ currency: LedgerCurrency,
kind: str,
- acct_id: "UUIDStr",
- lm: "LedgerManager",
+ acct_id: UUIDStr,
+ ledger_manager: LedgerManager,
):
"""Try to query for accounts that we know don't exist and confirm that
we either get the expected None result or it raises the correct
exception
"""
- qn = ":".join([currency, kind, acct_id])
+ qn = f"{currency}:{kind}:{acct_id}"
# (1) .get_account is just a wrapper for .get_account_many_ but
# call it either way
- assert lm.get_account(qualified_name=qn, raise_on_error=False) is None
+ assert (
+ ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None
+ )
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- lm.get_account(qualified_name=qn, raise_on_error=True)
+ ledger_manager.get_account(qualified_name=qn, raise_on_error=True)
# (2) .get_account_if_exists is another wrapper
- assert lm.get_account(qualified_name=qn, raise_on_error=False) is None
+ assert (
+ ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None
+ )
def test_get_account_no_results_many(
self,
- currency: "LedgerCurrency",
+ currency: LedgerCurrency,
kind: str,
- acct_id: "UUIDStr",
- lm: "LedgerManager",
+ acct_id: UUIDStr,
+ ledger_manager: LedgerManager,
):
- qn = ":".join([currency, kind, acct_id])
+ qn = f"{currency}:{kind}:{acct_id}"
# (1) .get_many_
- assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == []
+ assert (
+ ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=False)
+ == []
+ )
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- lm.get_account_many_(qualified_names=[qn], raise_on_error=True)
+ ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=True)
# (2) .get_many
- assert lm.get_account_many(qualified_names=[qn], raise_on_error=False) == []
+ assert (
+ ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=False)
+ == []
+ )
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- lm.get_account_many(qualified_names=[qn], raise_on_error=True)
+ ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=True)
# (3) .get_accounts(..)
- assert lm.get_accounts_if_exists(qualified_names=[qn]) == []
+ assert ledger_manager.get_accounts_if_exists(qualified_names=[qn]) == []
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- lm.get_accounts(qualified_names=[qn])
+ ledger_manager.get_accounts(qualified_names=[qn])
@pytest.mark.parametrize(
@@ -114,10 +116,10 @@ class TestLedgerAccountManagerCreate:
def test_create_account_error_permission(
self,
- currency: "LedgerCurrency",
- account_type: "AccountType",
- direction: "Direction",
- lm: "LedgerManager",
+ currency: LedgerCurrency,
+ account_type: AccountType,
+ direction: Direction,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -134,11 +136,11 @@ class TestLedgerAccountManagerCreate:
# (1) With no Permissions defined
test_lm = LedgerManager(
- pg_config=lm.pg_config,
+ pg_config=ledger_manager.pg_config,
permissions=[],
- redis_config=lm.redis_config,
- cache_prefix=lm.cache_prefix,
- testing=lm.testing,
+ redis_config=ledger_manager.redis_config,
+ cache_prefix=ledger_manager.cache_prefix,
+ testing=ledger_manager.testing,
)
with pytest.raises(expected_exception=AssertionError) as excinfo:
@@ -149,11 +151,11 @@ class TestLedgerAccountManagerCreate:
# (2) With Permissions defined, but not CREATE
test_lm = LedgerManager(
- pg_config=lm.pg_config,
+ pg_config=ledger_manager.pg_config,
permissions=[Permission.READ, Permission.UPDATE, Permission.DELETE],
- redis_config=lm.redis_config,
- cache_prefix=lm.cache_prefix,
- testing=lm.testing,
+ redis_config=ledger_manager.redis_config,
+ cache_prefix=ledger_manager.cache_prefix,
+ testing=ledger_manager.testing,
)
with pytest.raises(expected_exception=AssertionError) as excinfo:
@@ -164,10 +166,10 @@ class TestLedgerAccountManagerCreate:
def test_create(
self,
- currency: "LedgerCurrency",
- account_type: "AccountType",
- direction: "Direction",
- lm: "LedgerManager",
+ currency: LedgerCurrency,
+ account_type: AccountType,
+ direction: Direction,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -184,20 +186,20 @@ class TestLedgerAccountManagerCreate:
account_type=account_type,
normal_balance=direction,
)
- account = lm.create_account(account=acct_model)
+ account = ledger_manager.create_account(account=acct_model)
assert isinstance(account, LedgerAccount)
# Query for, and make sure the Account was saved in the DB
- res = lm.get_account(qualified_name=qn, raise_on_error=True)
+ res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True)
assert res is not None
assert account.uuid == res.uuid
def test_get_or_create(
self,
- currency: "LedgerCurrency",
- account_type: "AccountType",
- direction: "Direction",
- lm: "LedgerManager",
+ currency: LedgerCurrency,
+ account_type: AccountType,
+ direction: Direction,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -214,27 +216,31 @@ class TestLedgerAccountManagerCreate:
account_type=account_type,
normal_balance=direction,
)
- account = lm.get_account_or_create(account=acct_model)
+ account = ledger_manager.get_account_or_create(account=acct_model)
assert isinstance(account, LedgerAccount)
# Query for, and make sure the Account was saved in the DB
- res = lm.get_account(qualified_name=qn, raise_on_error=True)
+ res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True)
assert res is not None
assert account.uuid == res.uuid
class TestLedgerAccountManagerGet:
- def test_get(self, ledger_account: "LedgerAccount", lm: "LedgerManager"):
- res = lm.get_account(qualified_name=ledger_account.qualified_name)
+ def test_get(self, ledger_account: LedgerAccount, ledger_manager: LedgerManager):
+ res = ledger_manager.get_account(qualified_name=ledger_account.qualified_name)
assert res is not None
assert res.uuid == ledger_account.uuid
- res = lm.get_account_many(qualified_names=[ledger_account.qualified_name])
+ res = ledger_manager.get_account_many(
+ qualified_names=[ledger_account.qualified_name]
+ )
assert len(res) == 1
assert res[0].uuid == ledger_account.uuid
- res = lm.get_accounts(qualified_names=[ledger_account.qualified_name])
+ res = ledger_manager.get_accounts(
+ qualified_names=[ledger_account.qualified_name]
+ )
assert len(res) == 1
assert res[0].uuid == ledger_account.uuid
@@ -243,30 +249,30 @@ class TestLedgerAccountManagerGet:
def test_get_balance_empty(
self,
- ledger_account: "LedgerAccount",
- ledger_account_credit: "LedgerAccount",
- ledger_account_debit: "LedgerAccount",
- ledger_tx: "LedgerTransaction",
- lm: "LedgerManager",
+ ledger_account: LedgerAccount,
+ ledger_account_credit: LedgerAccount,
+ ledger_account_debit: LedgerAccount,
+ ledger_tx: LedgerTransaction,
+ ledger_manager: LedgerManager,
):
- res = lm.get_account_balance(account=ledger_account)
+ res = ledger_manager.get_account_balance(account=ledger_account)
assert res == 0
- res = lm.get_account_balance(account=ledger_account_credit)
+ res = ledger_manager.get_account_balance(account=ledger_account_credit)
assert res == 100
- res = lm.get_account_balance(account=ledger_account_debit)
+ res = ledger_manager.get_account_balance(account=ledger_account_debit)
assert res == 100
@pytest.mark.parametrize("n_times", range(5))
def test_get_account_filtered_balance(
self,
- ledger_account: "LedgerAccount",
- ledger_account_credit: "LedgerAccount",
- ledger_account_debit: "LedgerAccount",
- ledger_tx: "LedgerTransaction",
- n_times: "PositiveInt",
- lm: "LedgerManager",
+ ledger_account: LedgerAccount,
+ ledger_account_credit: LedgerAccount,
+ ledger_account_debit: LedgerAccount,
+ ledger_tx: LedgerTransaction,
+ n_times: PositiveInt,
+ ledger_manager: LedgerManager,
):
"""Try searching for random metadata and confirm it's always 0 because
Tx can be found.
@@ -275,7 +281,7 @@ class TestLedgerAccountManagerGet:
rand_value = uuid4().hex
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=ledger_account, metadata_key=rand_key, metadata_value=rand_value
)
== 0
@@ -285,7 +291,7 @@ class TestLedgerAccountManagerGet:
# and that we can filter it back
rand_amount = randint(10, 1_000)
- lm.create_tx(
+ ledger_manager.create_tx(
entries=[
LedgerEntry(
direction=Direction.CREDIT,
@@ -302,7 +308,7 @@ class TestLedgerAccountManagerGet:
)
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=ledger_account_credit,
metadata_key=rand_key,
metadata_value=rand_value,
@@ -311,7 +317,7 @@ class TestLedgerAccountManagerGet:
)
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=ledger_account_debit,
metadata_key=rand_key,
metadata_value=rand_value,
@@ -320,7 +326,7 @@ class TestLedgerAccountManagerGet:
)
def test_get_balance_timerange_empty(
- self, ledger_account: "LedgerAccount", lm: "LedgerManager"
+ self, ledger_account: LedgerAccount, ledger_manager: LedgerManager
):
- res = lm.get_account_balance_timerange(account=ledger_account)
+ res = ledger_manager.get_account_balance_timerange(account=ledger_account)
assert res == 0
diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py
index 37b7ba3..025f6ac 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx.py
@@ -1,33 +1,41 @@
+from __future__ import annotations
+
from decimal import Decimal
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.currency import LedgerCurrency
-from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+from generalresearch.managers.thl.ledger_manager.ledger import (
+ LedgerManager,
+)
from generalresearch.models.thl.ledger import (
Direction,
LedgerEntry,
LedgerTransaction,
)
+if TYPE_CHECKING:
+ from generalresearch.models.thl.ledger import (
+ LedgerAccount,
+ )
+
class TestLedgerManagerCreateTx:
- def test_create_account_error_permission(self, lm):
+ def test_create_account_error_permission(self, ledger_manager: LedgerManager):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
"""
- acct_uuid = uuid4().hex
-
# (1) With no Permissions defined
test_lm = LedgerManager(
- pg_config=lm.pg_config,
+ pg_config=ledger_manager.pg_config,
permissions=[],
- redis_config=lm.redis_config,
- cache_prefix=lm.cache_prefix,
- testing=lm.testing,
+ redis_config=ledger_manager.redis_config,
+ cache_prefix=ledger_manager.cache_prefix,
+ testing=ledger_manager.testing,
)
with pytest.raises(expected_exception=AssertionError) as excinfo:
@@ -37,9 +45,12 @@ class TestLedgerManagerCreateTx:
== "LedgerTransactionManager has insufficient Permissions"
)
- def test_create_assertions(self, ledger_account_debit, ledger_account_credit, lm):
+ def test_create_assertions(
+ self,
+ ledger_manager: LedgerManager,
+ ):
with pytest.raises(expected_exception=ValueError) as excinfo:
- lm.create_tx(
+ ledger_manager.create_tx(
entries=[
{
"direction": Direction.CREDIT,
@@ -53,7 +64,12 @@ class TestLedgerManagerCreateTx:
in str(excinfo.value)
)
- def test_create(self, ledger_account_credit, ledger_account_debit, lm):
+ def test_create(
+ self,
+ ledger_account_credit: LedgerAccount,
+ ledger_account_debit: LedgerAccount,
+ ledger_manager: LedgerManager,
+ ):
amount = int(Decimal("1.00") * 100)
entries = [
@@ -70,15 +86,20 @@ class TestLedgerManagerCreateTx:
]
# Create a Transaction and validate the operation was successful
- tx = lm.create_tx(entries=entries)
+ tx = ledger_manager.create_tx(entries=entries)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert isinstance(res, LedgerTransaction)
assert len(res.entries) == 2
assert tx.id == res.id
- def test_create_and_reverse(self, ledger_account_credit, ledger_account_debit, lm):
+ def test_create_and_reverse(
+ self,
+ ledger_account_credit: LedgerAccount,
+ ledger_account_debit: LedgerAccount,
+ ledger_manager: LedgerManager,
+ ):
amount = int(Decimal("1.00") * 100)
entries = [
@@ -94,13 +115,13 @@ class TestLedgerManagerCreateTx:
),
]
- tx = lm.create_tx(entries=entries)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ tx = ledger_manager.create_tx(entries=entries)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.id == tx.id
- assert lm.get_account_balance(account=ledger_account_credit) == 100
- assert lm.get_account_balance(account=ledger_account_debit) == 100
- assert lm.check_ledger_balanced() is True
+ assert ledger_manager.get_account_balance(account=ledger_account_credit) == 100
+ assert ledger_manager.get_account_balance(account=ledger_account_debit) == 100
+ assert ledger_manager.check_ledger_balanced() is True
# Reverse it
entries = [
@@ -116,13 +137,13 @@ class TestLedgerManagerCreateTx:
),
]
- tx = lm.create_tx(entries=entries)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ tx = ledger_manager.create_tx(entries=entries)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.id == tx.id
- assert lm.get_account_balance(ledger_account_credit) == 0
- assert lm.get_account_balance(ledger_account_debit) == 0
- assert lm.check_ledger_balanced()
+ assert ledger_manager.get_account_balance(ledger_account_credit) == 0
+ assert ledger_manager.get_account_balance(ledger_account_debit) == 0
+ assert ledger_manager.check_ledger_balanced()
# subtract again
entries = [
@@ -137,52 +158,60 @@ class TestLedgerManagerCreateTx:
amount=amount,
),
]
- tx = lm.create_tx(entries=entries)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ tx = ledger_manager.create_tx(entries=entries)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.id == tx.id
- assert lm.get_account_balance(ledger_account_credit) == -100
- assert lm.get_account_balance(ledger_account_debit) == -100
- assert lm.check_ledger_balanced()
+ assert ledger_manager.get_account_balance(ledger_account_credit) == -100
+ assert ledger_manager.get_account_balance(ledger_account_debit) == -100
+ assert ledger_manager.check_ledger_balanced()
class TestLedgerManagerGetTx:
# @pytest.mark.parametrize("currency", [LedgerCurrency.TEST], indirect=True)
- def test_get_tx_by_id(self, ledger_tx, lm):
+ def test_get_tx_by_id(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
with pytest.raises(expected_exception=AssertionError):
- lm.get_tx_by_id(transaction_id=ledger_tx)
+ ledger_manager.get_tx_by_id(transaction_id=ledger_tx)
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert res.id == ledger_tx.id
# @pytest.mark.parametrize("currency", [LedgerCurrency.TEST], indirect=True)
- def test_get_tx_by_ids(self, ledger_tx, lm):
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ def test_get_tx_by_ids(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert res.id == ledger_tx.id
@pytest.mark.parametrize(
"tag", [f"{LedgerCurrency.TEST}:{uuid4().hex}"], indirect=True
)
- def test_get_tx_ids_by_tag(self, ledger_tx, tag, lm):
+ def test_get_tx_ids_by_tag(
+ self, ledger_tx: LedgerTransaction, tag: str, ledger_manager: LedgerManager
+ ):
# (1) search for a random tag
- res = lm.get_tx_ids_by_tag(tag="aaa:bbb")
+ res = ledger_manager.get_tx_ids_by_tag(tag="aaa:bbb")
assert isinstance(res, set)
assert len(res) == 0
# (2) search for the tag that was used during ledger_transaction creation
- res = lm.get_tx_ids_by_tag(tag=tag)
+ res = ledger_manager.get_tx_ids_by_tag(tag=tag)
assert isinstance(res, set)
assert len(res) == 1
- def test_get_tx_by_tag(self, ledger_tx, tag, lm):
+ def test_get_tx_by_tag(
+ self, ledger_tx: LedgerTransaction, tag: str, ledger_manager: LedgerManager
+ ):
# (1) search for a random tag
- res = lm.get_tx_by_tag(tag="aaa:bbb")
+ res = ledger_manager.get_tx_by_tag(tag="aaa:bbb")
assert isinstance(res, list)
assert len(res) == 0
# (2) search for the tag that was used during ledger_transaction creation
- res = lm.get_tx_by_tag(tag=tag)
+ res = ledger_manager.get_tx_by_tag(tag=tag)
assert isinstance(res, list)
assert len(res) == 1
@@ -190,42 +219,60 @@ class TestLedgerManagerGetTx:
assert ledger_tx.id == res[0].id
def test_get_tx_filtered_by_account(
- self, ledger_tx, ledger_account, ledger_account_debit, ledger_account_credit, lm
+ self,
+ ledger_tx: LedgerTransaction,
+ ledger_account: LedgerAccount,
+ ledger_account_debit: LedgerAccount,
+ ledger_account_credit: LedgerAccount,
+ ledger_manager: LedgerManager,
):
# (1) Do basic assertion checks first
with pytest.raises(expected_exception=AssertionError) as excinfo:
- lm.get_tx_filtered_by_account(account_uuid=ledger_account)
+ ledger_manager.get_tx_filtered_by_account(account_uuid=ledger_account)
assert str(excinfo.value) == "account_uuid must be a str"
# (2) This search doesn't return anything because this ledger account
# wasn't actually used in the entries for the ledger_transaction
- res = lm.get_tx_filtered_by_account(account_uuid=ledger_account.uuid)
+ res = ledger_manager.get_tx_filtered_by_account(
+ account_uuid=ledger_account.uuid
+ )
assert len(res) == 0
# (3) Either the credit or the debit example ledger_accounts wll work
# to find this transaction because they're both used in the entries
- res = lm.get_tx_filtered_by_account(account_uuid=ledger_account_debit.uuid)
+ res = ledger_manager.get_tx_filtered_by_account(
+ account_uuid=ledger_account_debit.uuid
+ )
assert len(res) == 1
assert res[0].id == ledger_tx.id
- res = lm.get_tx_filtered_by_account(account_uuid=ledger_account_credit.uuid)
+ res = ledger_manager.get_tx_filtered_by_account(
+ account_uuid=ledger_account_credit.uuid
+ )
assert len(res) == 1
assert ledger_tx.id == res[0].id
- res2 = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res2 = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert res2.model_dump_json() == res[0].model_dump_json()
- def test_filter_metadata(self, ledger_tx, tx_metadata, lm):
+ def test_filter_metadata(
+ self,
+ ledger_tx: LedgerTransaction,
+ tx_metadata: dict[str, str] | None,
+ ledger_manager: LedgerManager,
+ ):
key, value = next(iter(tx_metadata.items()))
# (1) Confirm a random key,value pair returns nothing
- res = lm.get_tx_filtered_by_metadata(
+ res = ledger_manager.get_tx_filtered_by_metadata(
metadata_key=f"key-{uuid4().hex[:10]}", metadata_value=uuid4().hex[:12]
)
assert len(res) == 0
# (2) confirm a key,value pair return the correct results
- res = lm.get_tx_filtered_by_metadata(metadata_key=key, metadata_value=value)
+ res = ledger_manager.get_tx_filtered_by_metadata(
+ metadata_key=key, metadata_value=value
+ )
assert len(res) == 1
# assert 0 == THL_lm.get_filtered_account_balance(account2, "thl_wall", "ccc")
diff --git a/tests/managers/thl/test_ledger/test_lm_tx_entries.py b/tests/managers/thl/test_ledger/test_lm_tx_entries.py
index 5bf1c48..03c6e02 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx_entries.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx_entries.py
@@ -1,25 +1,41 @@
-from generalresearch.models.thl.ledger import LedgerEntry
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+from generalresearch.models.thl.ledger import (
+ LedgerEntry,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.models.thl.ledger import (
+ LedgerTransaction,
+ )
class TestLedgerEntryManager:
- def test_get_tx_entries_by_tx(self, ledger_tx, lm):
+ def test_get_tx_entries_by_tx(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
# First confirm the Ledger TX exists with 2 Entries
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert len(res.entries) == 2
- tx_entries = lm.get_tx_entries_by_tx(transaction=ledger_tx)
+ tx_entries = ledger_manager.get_tx_entries_by_tx(transaction=ledger_tx)
assert len(tx_entries) == 2
assert res.entries == tx_entries
assert isinstance(tx_entries[0], LedgerEntry)
- def test_get_tx_entries_by_txs(self, ledger_tx, lm):
+ def test_get_tx_entries_by_txs(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
# First confirm the Ledger TX exists with 2 Entries
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert len(res.entries) == 2
- tx_entries = lm.get_tx_entries_by_txs(transactions=[ledger_tx])
+ tx_entries = ledger_manager.get_tx_entries_by_txs(transactions=[ledger_tx])
assert len(tx_entries) == 2
assert res.entries == tx_entries
diff --git a/tests/managers/thl/test_ledger/test_lm_tx_locks.py b/tests/managers/thl/test_ledger/test_lm_tx_locks.py
index df2611b..166598e 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py
@@ -1,29 +1,38 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable, Generator
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
+from pytest import LogCaptureFixture
from generalresearch.managers.thl.ledger_manager.conditions import (
generate_condition_mp_payment,
)
from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerTransactionCreateError,
LedgerTransactionCreateLockError,
LedgerTransactionFlagAlreadyExistsError,
- LedgerTransactionCreateError,
)
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.ledger import LedgerTransaction
from generalresearch.models.thl.session import (
- Wall,
+ Session,
Status,
StatusCode1,
- Session,
+ Wall,
WallAdjustedStatus,
)
-from generalresearch.models.thl.user import User
-from test_utils.models.conftest import user_factory, session, product_user_wallet_no
+
+if TYPE_CHECKING:
+ from generalresearch.currency import LedgerCurrency
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
logger = logging.getLogger("LedgerManager")
@@ -32,17 +41,17 @@ class TestLedgerLocks:
def test_a(
self,
- user_factory,
- session_factory,
- product_user_wallet_no,
- create_main_accounts,
- caplog,
- thl_lm,
- lm,
- utc_hour_ago,
- currency,
- wall_factory,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ caplog: Generator[LogCaptureFixture],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ utc_hour_ago: datetime,
+ currency: LedgerCurrency,
+ wall_factory: Callable[..., Wall],
+ delete_ledger_db: Callable[..., None],
):
"""
TODO: This whole test is confusing a I don't really understand.
@@ -56,18 +65,22 @@ class TestLedgerLocks:
s1 = session_factory(
user=user,
wall_count=3,
- wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal("4")],
+ wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal(4)],
wall_statuses=[Status.COMPLETE, Status.COMPLETE, Status.COMPLETE],
)
# A User does a Wall Completion in Session=1
w1 = s1.wall_events[0]
- tx = thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started)
+ tx = thl_ledger_manager.create_tx_task_complete(
+ wall=w1, user=user, created=w1.started
+ )
assert isinstance(tx, LedgerTransaction)
# A User does another Wall Completion in Session=1
w2 = s1.wall_events[1]
- tx = thl_lm.create_tx_task_complete(wall=w2, user=user, created=w2.started)
+ tx = thl_ledger_manager.create_tx_task_complete(
+ wall=w2, user=user, created=w2.started
+ )
assert isinstance(tx, LedgerTransaction)
# That first Wall Complete was "adjusted" to instead be marked
@@ -77,7 +90,7 @@ class TestLedgerLocks:
adjusted_cpi=0,
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
- tx = thl_lm.create_tx_task_adjustment(wall=w1, user=user)
+ tx = thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user)
assert isinstance(tx, LedgerTransaction)
# A User does another! Wall Completion in Session=1; however, we
@@ -86,60 +99,63 @@ class TestLedgerLocks:
# Make sure we clear any flags/locks first
lock_key = f"{currency.value}:thl_wall:{w3.uuid}"
- lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}"
- flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}"
- lm.redis_client.delete(lock_name)
- lm.redis_client.delete(flag_name)
+ lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}"
+ flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}"
+ ledger_manager.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(flag_name)
# Despite the
f1 = generate_condition_mp_payment(wall=w1)
f2 = generate_condition_mp_payment(wall=w2)
f3 = generate_condition_mp_payment(wall=w3)
- assert f1(lm=lm) is False
- assert f2(lm=lm) is False
- assert f3(lm=lm) is True
+ assert f1(ledger_manager) is False
+ assert f2(lm=ledger_manager) is False
+ assert f3(lm=ledger_manager) is True
condition = f3
- create_tx_func = lambda: thl_lm.create_tx_task_complete_(wall=w3, user=user)
+ create_tx_func = lambda: thl_ledger_manager.create_tx_task_complete_(
+ wall=w3, user=user
+ )
assert isinstance(create_tx_func, Callable)
- assert f3(lm) is True
+ assert f3(ledger_manager) is True
- lm.redis_client.delete(flag_name)
- lm.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(flag_name)
+ ledger_manager.redis_client.delete(lock_name)
- tx = thl_lm.create_tx_protected(
+ tx = thl_ledger_manager.create_tx_protected(
lock_key=lock_key, condition=condition, create_tx_func=create_tx_func
)
- assert f3(lm) is False
+ assert f3(ledger_manager) is False
# purposely hold the lock open
tx = None
- lm.redis_client.set(lock_name, "1")
- with caplog.at_level(logging.ERROR):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_lm.create_tx_protected(
- lock_key=lock_key,
- condition=condition,
- create_tx_func=create_tx_func,
- )
- assert tx is None
+ ledger_manager.redis_client.set(lock_name, "1")
+ with caplog.at_level(logging.ERROR), pytest.raises(
+ expected_exception=LedgerTransactionCreateLockError
+ ):
+ tx = thl_ledger_manager.create_tx_protected(
+ lock_key=lock_key,
+ condition=condition,
+ create_tx_func=create_tx_func,
+ )
+ assert tx is None
assert "Unable to acquire lock within the time specified" in caplog.text
- lm.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(lock_name)
def test_locking(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
- delete_ledger_db,
- caplog,
- thl_lm,
- lm,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ caplog: Generator[LogCaptureFixture],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
delete_ledger_db()
create_main_accounts()
- now = datetime.now(timezone.utc) - timedelta(hours=1)
+ now = datetime.now(UTC) - timedelta(hours=1)
user: User = user_factory(product=product_user_wallet_no)
# A User does a Wall complete on Session.id=1 and the transaction is
@@ -155,7 +171,9 @@ class TestLedgerLocks:
started=now,
finished=now + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
+ )
# A User does a Wall complete on Session.id=1 and the transaction is
# logged to the ledger
@@ -170,7 +188,9 @@ class TestLedgerLocks:
started=now,
finished=now + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall=wall2, user=user, created=wall2.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall2, user=user, created=wall2.started
+ )
# An hour later, the first wall complete is adjusted to a Failure and
# it's tracked in the ledger
@@ -179,7 +199,7 @@ class TestLedgerLocks:
adjusted_cpi=0,
adjusted_timestamp=now + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall=wall1, user=user)
+ thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user)
# A User does a Wall complete on Session.id=1 and the transaction
# IS NOT logged to the ledger
@@ -187,7 +207,7 @@ class TestLedgerLocks:
user_id=user.user_id,
source=Source.DYNATA,
req_survey_id="xxx",
- req_cpi=Decimal("4"),
+ req_cpi=Decimal(4),
session_id=1,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
@@ -196,52 +216,53 @@ class TestLedgerLocks:
uuid="867a282d8b4d40d2a2093d75b802b629",
)
- revenue_account = thl_lm.get_account_task_complete_revenue()
- assert 0 == thl_lm.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ assert 0 == thl_ledger_manager.get_account_filtered_balance(
account=revenue_account,
metadata_key="thl_wall",
metadata_value=wall3.uuid,
)
# Make sure we clear any flags/locks first
lock_key = f"test:thl_wall:{wall3.uuid}"
- lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}"
- flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}"
- lm.redis_client.delete(lock_name)
- lm.redis_client.delete(flag_name)
+ lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}"
+ flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}"
+ ledger_manager.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(flag_name)
# Purposely hold the lock open
- lm.redis_client.set(name=lock_name, value="1")
- with caplog.at_level(logging.DEBUG):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_lm.create_tx_task_complete(
- wall=wall3, user=user, created=wall3.started
- )
- assert isinstance(tx, LedgerTransaction)
+ ledger_manager.redis_client.set(name=lock_name, value="1")
+ with caplog.at_level(logging.DEBUG), pytest.raises(
+ expected_exception=LedgerTransactionCreateLockError
+ ):
+ tx = thl_ledger_manager.create_tx_task_complete(
+ wall=wall3, user=user, created=wall3.started
+ )
+ assert isinstance(tx, LedgerTransaction)
assert "Unable to acquire lock within the time specified" in caplog.text
# Release the lock
- lm.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(lock_name)
# Set the redis flag to indicate it has been run
- lm.redis_client.set(flag_name, "1")
+ ledger_manager.redis_client.set(flag_name, "1")
# with self.assertLogs(logger=logger, level=logging.DEBUG) as cm2:
with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError):
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall3, user=user, created=wall3.started
)
# self.assertIn("entered_lock: True, flag_set: True", cm2.output[0])
# Unset the flag
- lm.redis_client.delete(flag_name)
+ ledger_manager.redis_client.delete(flag_name)
- assert 0 == lm.get_account_filtered_balance(
+ assert 0 == ledger_manager.get_account_filtered_balance(
account=revenue_account,
metadata_key="thl_wall",
metadata_value=wall3.uuid,
)
# Now actually run it
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall3, user=user, created=wall3.started
)
assert tx is not None
@@ -250,29 +271,34 @@ class TestLedgerLocks:
# Confirm the Exception inheritance works
tx = None
with pytest.raises(expected_exception=LedgerTransactionCreateError):
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall3, user=user, created=wall3.started
)
assert tx is None
# clear the redis flag, it should query the db
- assert lm.redis_client.get(flag_name) is not None
- lm.redis_client.delete(flag_name)
- assert lm.redis_client.get(flag_name) is None
+ assert ledger_manager.redis_client.get(flag_name) is not None
+ ledger_manager.redis_client.delete(flag_name)
+ assert ledger_manager.redis_client.get(flag_name) is None
with pytest.raises(expected_exception=LedgerTransactionCreateError):
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall3, user=user, created=wall3.started
)
- assert 400 == thl_lm.get_account_filtered_balance(
+ assert 400 == thl_ledger_manager.get_account_filtered_balance(
account=revenue_account,
metadata_key="thl_wall",
metadata_value=wall3.uuid,
)
def test_bp_payment_without_locks(
- self, user_factory, product_user_wallet_no, create_main_accounts, thl_lm, lm
+ self,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
user: User = user_factory(product=product_user_wallet_no)
wall1 = Wall(
@@ -283,39 +309,46 @@ class TestLedgerLocks:
session_id=3,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
+ )
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": session.started + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=session.started + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
print(thl_net, commission_amount, bp_pay, user_pay)
# Run it 3 times without any checks, and it gets made three times!
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
- thl_lm.create_tx_bp_payment_(session=session, created=wall1.started)
- thl_lm.create_tx_bp_payment_(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment_(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment_(session=session, created=wall1.started)
- bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- assert 48 * 3 == lm.get_account_balance(account=bp_wallet)
- assert 48 * 3 == thl_lm.get_account_filtered_balance(
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
+ assert 48 * 3 == ledger_manager.get_account_balance(account=bp_wallet)
+ assert 48 * 3 == thl_ledger_manager.get_account_filtered_balance(
account=bp_wallet, metadata_key="thl_session", metadata_value=session.uuid
)
- assert lm.check_ledger_balanced()
+ assert ledger_manager.check_ledger_balanced()
def test_bp_payment_with_locks(
- self, user_factory, product_user_wallet_no, create_main_accounts, thl_lm, lm
+ self,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
user: User = user_factory(product=product_user_wallet_no)
@@ -327,45 +360,49 @@ class TestLedgerLocks:
session_id=3,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall1, user, created=wall1.started)
+ thl_ledger_manager.create_tx_task_complete(wall1, user, created=wall1.started)
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": session.started + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=session.started + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
print(thl_net, commission_amount, bp_pay, user_pay)
# Make sure we clear any flags/locks first
lock_key = f"test:thl_wall:{wall1.uuid}"
- lock_name = f"{lm.cache_prefix}:transaction_lock:{lock_key}"
- flag_name = f"{lm.cache_prefix}:transaction_flag:{lock_key}"
- lm.redis_client.delete(lock_name)
- lm.redis_client.delete(flag_name)
+ lock_name = f"{ledger_manager.cache_prefix}:transaction_lock:{lock_key}"
+ flag_name = f"{ledger_manager.cache_prefix}:transaction_flag:{lock_key}"
+ ledger_manager.redis_client.delete(lock_name)
+ ledger_manager.redis_client.delete(flag_name)
# Run it 3 times with check, and it gets made once!
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started)
with pytest.raises(expected_exception=LedgerTransactionCreateError):
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(
+ session=session, created=wall1.started
+ )
with pytest.raises(expected_exception=LedgerTransactionCreateError):
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(
+ session=session, created=wall1.started
+ )
- bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- assert 48 == thl_lm.get_account_balance(bp_wallet)
- assert 48 == thl_lm.get_account_filtered_balance(
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
+ assert 48 == thl_ledger_manager.get_account_balance(bp_wallet)
+ assert 48 == thl_ledger_manager.get_account_filtered_balance(
account=bp_wallet,
metadata_key="thl_session",
metadata_value=session.uuid,
)
- assert lm.check_ledger_balanced()
+ assert ledger_manager.check_ledger_balanced()
diff --git a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py
index 5d12633..3d8cf89 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx_metadata.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx_metadata.py
@@ -1,34 +1,55 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import (
+ LedgerManager,
+ LedgerTransaction,
+ )
+
+
class TestLedgerMetadataManager:
- def test_get_tx_metadata_by_txs(self, ledger_tx, lm):
+ def test_get_tx_metadata_by_txs(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
# First confirm the Ledger TX exists with 2 Entries
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
assert isinstance(res.metadata, dict)
- tx_metadatas = lm.get_tx_metadata_by_txs(transactions=[ledger_tx])
+ tx_metadatas = ledger_manager.get_tx_metadata_by_txs(transactions=[ledger_tx])
assert isinstance(tx_metadatas, dict)
assert isinstance(tx_metadatas[ledger_tx.id], dict)
assert res.metadata == tx_metadatas[ledger_tx.id]
- def test_get_tx_metadata_ids_by_tx(self, ledger_tx, lm):
+ def test_get_tx_metadata_ids_by_tx(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
# First confirm the Ledger TX exists with 2 Entries
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
tx_metadata_cnt = len(res.metadata.keys())
- tx_metadata_ids = lm.get_tx_metadata_ids_by_tx(transaction=ledger_tx)
+ tx_metadata_ids = ledger_manager.get_tx_metadata_ids_by_tx(
+ transaction=ledger_tx
+ )
assert isinstance(tx_metadata_ids, set)
- assert isinstance(list(tx_metadata_ids)[0], int)
+ assert isinstance(next(iter(tx_metadata_ids)), int)
assert tx_metadata_cnt == len(tx_metadata_ids)
- def test_get_tx_metadata_ids_by_txs(self, ledger_tx, lm):
+ def test_get_tx_metadata_ids_by_txs(
+ self, ledger_tx: LedgerTransaction, ledger_manager: LedgerManager
+ ):
# First confirm the Ledger TX exists with 2 Entries
- res = lm.get_tx_by_id(transaction_id=ledger_tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=ledger_tx.id)
tx_metadata_cnt = len(res.metadata.keys())
- tx_metadata_ids = lm.get_tx_metadata_ids_by_txs(transactions=[ledger_tx])
+ tx_metadata_ids = ledger_manager.get_tx_metadata_ids_by_txs(
+ transactions=[ledger_tx]
+ )
assert isinstance(tx_metadata_ids, set)
- assert isinstance(list(tx_metadata_ids)[0], int)
+ assert isinstance(next(iter(tx_metadata_ids)), int)
assert tx_metadata_cnt == len(tx_metadata_ids)
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
index 01d5fe1..107ff00 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
@@ -1,19 +1,41 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from generalresearch.currency import LedgerCurrency
+from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerAccountDoesntExistError,
+)
+from generalresearch.models.thl.ledger import (
+ AccountType,
+ Direction,
+ LedgerAccount,
+)
+from generalresearch.models.thl.product import Product
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import (
+ LedgerAccountManager,
+ LedgerManager,
+ )
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.user import User
+
class TestThlLedgerManagerAccounts:
- def test_get_account_or_create_user_wallet(self, user, thl_lm, lm):
- from generalresearch.currency import LedgerCurrency
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_get_account_or_create_user_wallet(
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
- account = thl_lm.get_account_or_create_user_wallet(user=user)
+ account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
assert isinstance(account, LedgerAccount)
assert user.uuid in account.qualified_name
@@ -25,18 +47,20 @@ class TestThlLedgerManagerAccounts:
assert account.currency == LedgerCurrency.TEST
# Actually query for it to confirm
- res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True)
+ res = ledger_manager.get_account(
+ qualified_name=account.qualified_name, raise_on_error=True
+ )
+ assert isinstance(res, LedgerAccount)
assert res.model_dump_json() == account.model_dump_json()
- def test_get_account_or_create_bp_wallet(self, product, thl_lm, lm):
- from generalresearch.currency import LedgerCurrency
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_get_account_or_create_bp_wallet(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
- account = thl_lm.get_account_or_create_bp_wallet(product=product)
+ account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
assert isinstance(account, LedgerAccount)
assert product.uuid in account.qualified_name
@@ -48,17 +72,22 @@ class TestThlLedgerManagerAccounts:
assert account.currency == LedgerCurrency.TEST
# Actually query for it to confirm
- res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True)
+ res = ledger_manager.get_account(
+ qualified_name=account.qualified_name, raise_on_error=True
+ )
+ assert isinstance(res, LedgerAccount)
assert res.model_dump_json() == account.model_dump_json()
- def test_get_account_or_create_bp_commission(self, product, thl_lm, lm):
- from generalresearch.currency import LedgerCurrency
- from generalresearch.models.thl.ledger import (
- Direction,
- AccountType,
- )
+ def test_get_account_or_create_bp_commission(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
- account = thl_lm.get_account_or_create_bp_commission(product=product)
+ account = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=product
+ )
assert product.uuid in account.qualified_name
assert account.display_name == f"Revenue from commission {product.uuid}"
@@ -69,18 +98,21 @@ class TestThlLedgerManagerAccounts:
assert account.currency == LedgerCurrency.TEST
# Actually query for it to confirm
- res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True)
+ res = ledger_manager.get_account(
+ qualified_name=account.qualified_name, raise_on_error=True
+ )
+ assert isinstance(res, LedgerAccount)
assert res.model_dump_json() == account.model_dump_json()
@pytest.mark.parametrize("expense", ["tango", "paypal", "gift", "tremendous"])
- def test_get_account_or_create_bp_expense(self, product, expense, thl_lm, lm):
- from generalresearch.currency import LedgerCurrency
- from generalresearch.models.thl.ledger import (
- Direction,
- AccountType,
- )
-
- account = thl_lm.get_account_or_create_bp_expense(
+ def test_get_account_or_create_bp_expense(
+ self,
+ product: Product,
+ expense,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
+ account = thl_ledger_manager.get_account_or_create_bp_expense(
product=product, expense_name=expense
)
assert product.uuid in account.qualified_name
@@ -92,17 +124,22 @@ class TestThlLedgerManagerAccounts:
assert account.currency == LedgerCurrency.TEST
# Actually query for it to confirm
- res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True)
+ res = ledger_manager.get_account(
+ qualified_name=account.qualified_name, raise_on_error=True
+ )
+ assert isinstance(res, LedgerAccount)
assert res.model_dump_json() == account.model_dump_json()
- def test_get_or_create_bp_pending_payout_account(self, product, thl_lm, lm):
- from generalresearch.currency import LedgerCurrency
- from generalresearch.models.thl.ledger import (
- Direction,
- AccountType,
- )
+ def test_get_or_create_bp_pending_payout_account(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
- account = thl_lm.get_or_create_bp_pending_payout_account(product=product)
+ account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
+ product=product
+ )
assert product.uuid in account.qualified_name
assert account.display_name == f"BP Wallet Pending {product.uuid}"
@@ -113,11 +150,17 @@ class TestThlLedgerManagerAccounts:
assert account.currency == LedgerCurrency.TEST
# Actually query for it to confirm
- res = lm.get_account(qualified_name=account.qualified_name, raise_on_error=True)
+ res = ledger_manager.get_account(
+ qualified_name=account.qualified_name, raise_on_error=True
+ )
+ assert isinstance(res, LedgerAccount)
assert res.model_dump_json() == account.model_dump_json()
def test_get_account_task_complete_revenue_raises(
- self, delete_ledger_db, thl_lm, lm
+ self,
+ delete_ledger_db: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerAccountDoesntExistError,
@@ -126,63 +169,79 @@ class TestThlLedgerManagerAccounts:
delete_ledger_db()
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- thl_lm.get_account_task_complete_revenue()
+ thl_ledger_manager.get_account_task_complete_revenue()
def test_get_account_task_complete_revenue(
- self, account_cash, account_revenue_task_complete, thl_lm, lm
+ self, thl_ledger_manager: ThlLedgerManager, create_main_accounts
):
from generalresearch.models.thl.ledger import (
- LedgerAccount,
AccountType,
+ LedgerAccount,
)
- res = thl_lm.get_account_task_complete_revenue()
+ create_main_accounts()
+
+ res = thl_ledger_manager.get_account_task_complete_revenue()
assert isinstance(res, LedgerAccount)
assert res.reference_type is None
assert res.reference_uuid is None
assert res.account_type == AccountType.REVENUE
assert res.display_name == "Cash flow task complete"
- def test_get_account_cash_raises(self, delete_ledger_db, thl_lm, lm):
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
+ def test_get_account_cash_raises(
+ self,
+ delete_ledger_db: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ):
delete_ledger_db()
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- thl_lm.get_account_cash()
+ thl_ledger_manager.get_account_cash()
- def test_get_account_cash(self, account_cash, thl_lm, lm):
+ def test_get_account_cash(
+ self,
+ thl_ledger_manager: ThlLedgerManager,
+ create_main_accounts
+ ):
+ create_main_accounts()
from generalresearch.models.thl.ledger import (
- LedgerAccount,
AccountType,
+ LedgerAccount,
)
- res = thl_lm.get_account_cash()
+ res = thl_ledger_manager.get_account_cash()
assert isinstance(res, LedgerAccount)
assert res.reference_type is None
assert res.reference_uuid is None
assert res.account_type == AccountType.CASH
assert res.display_name == "Operating Cash Account"
- def test_get_accounts(self, setup_accounts, product, user_factory, thl_lm, lm, lam):
- from generalresearch.models.thl.user import User
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
+ def test_get_accounts(
+ self,
+ setup_accounts: Callable[..., None],
+ product: Product,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ledger_account_manager: LedgerAccountManager,
+ ):
+ setup_accounts()
- user1: User = user_factory(product=product)
- user2: User = user_factory(product=product)
+ _: User = user_factory(product=product)
+ _: User = user_factory(product=product)
- account1 = thl_lm.get_account_or_create_bp_wallet(product=product)
+ account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
# (1) known account and confirm it comes back
- res = lm.get_account(qualified_name=account1.qualified_name)
+ res = ledger_manager.get_account(qualified_name=account1.qualified_name)
+ assert isinstance(res, LedgerAccount)
assert account1.model_dump_json() == res.model_dump_json()
# (2) known accounts and confirm they both come back
- res = lam.get_accounts(qualified_names=[account1.qualified_name])
+ res = ledger_account_manager.get_accounts(
+ qualified_names=[account1.qualified_name]
+ )
assert isinstance(res, list)
assert len(res) == 1
assert account1 in res
@@ -190,28 +249,34 @@ class TestThlLedgerManagerAccounts:
# Get 2 known and 1 made up qualified names, and confirm it raises
# an error
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_accounts(
+ ledger_account_manager.get_accounts(
qualified_names=[
account1.qualified_name,
f"test:bp_wall:{uuid4().hex}",
]
)
- def test_get_accounts_if_exists(self, product_factory, currency, thl_lm, lm):
- from generalresearch.models.thl.product import Product
+ def test_get_accounts_if_exists(
+ self,
+ product_factory: Callable[..., Product],
+ currency: LedgerCurrency,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
p1: Product = product_factory()
p2: Product = product_factory()
- account1 = thl_lm.get_account_or_create_bp_wallet(product=p1)
- account2 = thl_lm.get_account_or_create_bp_wallet(product=p2)
+ account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ account2 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
# (1) known account and confirm it comes back
- res = lm.get_account(qualified_name=account1.qualified_name)
+ res = ledger_manager.get_account(qualified_name=account1.qualified_name)
+ assert isinstance(res, LedgerAccount)
assert account1.model_dump_json() == res.model_dump_json()
# (2) known accounts and confirm they both come back
- res = lm.get_accounts(
+ res = ledger_manager.get_accounts(
qualified_names=[account1.qualified_name, account2.qualified_name]
)
assert isinstance(res, list)
@@ -221,7 +286,7 @@ class TestThlLedgerManagerAccounts:
# Get 2 known and 1 made up qualified names, and confirm only 2
# come back
- lm.get_accounts_if_exists(
+ ledger_manager.get_accounts_if_exists(
qualified_names=[
account1.qualified_name,
account2.qualified_name,
@@ -233,53 +298,50 @@ class TestThlLedgerManagerAccounts:
assert len(res) == 2
# Confirm an empty array comes back for all unknown qualified names
- res = lm.get_accounts_if_exists(
+ assert isinstance(ledger_manager.currency, LedgerCurrency)
+ res = ledger_manager.get_accounts_if_exists(
qualified_names=[
- f"{lm.currency.value}:bp_wall:{uuid4().hex}" for i in range(5)
+ f"{ledger_manager.currency.value}:bp_wall:{uuid4().hex}"
+ for _ in range(5)
]
)
assert isinstance(res, list)
assert len(res) == 0
- def test_get_accounts_for_products(self, product_factory, thl_lm, lm):
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- )
-
+ def test_get_accounts_for_products(
+ self,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ ):
# Create 5 Products
product_uuids = []
- for i in range(5):
+ for _ in range(5):
_p = product_factory()
product_uuids.append(_p.uuid)
# Confirm that this fails.. because none of those accounts have been
# created yet
with pytest.raises(expected_exception=LedgerAccountDoesntExistError):
- thl_lm.get_accounts_bp_wallet_for_products(product_uuids=product_uuids)
+ thl_ledger_manager.get_accounts_bp_wallet_for_products(
+ product_uuids=product_uuids
+ )
# Create the bp_wallet accounts and then try again
for p_uuid in product_uuids:
- thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=p_uuid)
+ thl_ledger_manager.get_account_or_create_bp_wallet_by_uuid(
+ product_uuid=p_uuid
+ )
- res = thl_lm.get_accounts_bp_wallet_for_products(product_uuids=product_uuids)
+ res = thl_ledger_manager.get_accounts_bp_wallet_for_products(
+ product_uuids=product_uuids
+ )
assert len(res) == len(product_uuids)
- assert all([isinstance(i, LedgerAccount) for i in res])
+ assert all(isinstance(i, LedgerAccount) for i in res)
class TestLedgerAccountManager:
- def test_get_or_create(self, thl_lm, lm, lam):
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_get_or_create(self, ledger_account_manager: LedgerAccountManager):
u = uuid4().hex
name = f"test-{u[:8]}"
@@ -297,48 +359,52 @@ class TestLedgerAccountManager:
# First we want to validate that using the get_account method raises
# an error for a random LedgerAccount which we know does not exist.
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_account(qualified_name=account.qualified_name)
+ ledger_account_manager.get_account(qualified_name=account.qualified_name)
# Now that we know it doesn't exist, get_or_create for it
- instance = lam.get_account_or_create(account=account)
+ instance = ledger_account_manager.get_account_or_create(account=account)
# It should always return
assert isinstance(instance, LedgerAccount)
assert instance.reference_uuid == u
- def test_get(self, user, thl_lm, lm, lam):
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- AccountType,
- )
+ def test_get(
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_account_manager: LedgerAccountManager,
+ ):
+
+ assert isinstance(user.product, Product)
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_account(qualified_name=f"test:bp_wallet:{user.product.id}")
+ ledger_account_manager.get_account(
+ qualified_name=f"test:bp_wallet:{user.product.id}"
+ )
- thl_lm.get_account_or_create_bp_wallet(product=user.product)
- account = lam.get_account(qualified_name=f"test:bp_wallet:{user.product.id}")
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=user.product)
+ account = ledger_account_manager.get_account(
+ qualified_name=f"test:bp_wallet:{user.product.id}"
+ )
assert isinstance(account, LedgerAccount)
assert AccountType.BP_WALLET == account.account_type
assert user.product.uuid == account.reference_uuid
- def test_get_many(self, product_factory, thl_lm, lm, lam, currency):
- from generalresearch.models.thl.product import Product
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
-
+ def test_get_many(
+ self,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_account_manager: LedgerAccountManager,
+ ):
p1: Product = product_factory()
p2: Product = product_factory()
- account1 = thl_lm.get_account_or_create_bp_wallet(product=p1)
- account2 = thl_lm.get_account_or_create_bp_wallet(product=p2)
+ account1 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ account2 = thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
# Get 1 known account and confirm it comes back
- res = lam.get_account_many(
+ res = ledger_account_manager.get_account_many(
qualified_names=[account1.qualified_name, account2.qualified_name]
)
assert isinstance(res, list)
@@ -346,7 +412,7 @@ class TestLedgerAccountManager:
assert account1 in res
# Get 2 known accounts and confirm they both come back
- res = lam.get_account_many(
+ res = ledger_account_manager.get_account_many(
qualified_names=[account1.qualified_name, account2.qualified_name]
)
assert isinstance(res, list)
@@ -356,7 +422,7 @@ class TestLedgerAccountManager:
# Get 2 known and 1 made up qualified names, and confirm only 2 come
# back. Don't raise on error, so we can confirm the array is "short"
- res = lam.get_account_many(
+ res = ledger_account_manager.get_account_many(
qualified_names=[
account1.qualified_name,
account2.qualified_name,
@@ -369,7 +435,7 @@ class TestLedgerAccountManager:
# Same as above, but confirm the raise works on checking res length
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_account_many(
+ ledger_account_manager.get_account_many(
qualified_names=[
account1.qualified_name,
account2.qualified_name,
@@ -379,19 +445,14 @@ class TestLedgerAccountManager:
)
# Confirm an empty array comes back for all unknown qualified names
- res = lam.get_account_many(
- qualified_names=[f"test:bp_wall:{uuid4().hex}" for i in range(5)],
+ res = ledger_account_manager.get_account_many(
+ qualified_names=[f"test:bp_wall:{uuid4().hex}" for _ in range(5)],
raise_on_error=False,
)
assert isinstance(res, list)
assert len(res) == 0
- def test_create_account(self, thl_lm, lm, lam):
- from generalresearch.models.thl.ledger import (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_create_account(self, ledger_account_manager: LedgerAccountManager):
u = uuid4().hex
name = f"test-{u[:8]}"
@@ -406,6 +467,6 @@ class TestLedgerAccountManager:
reference_uuid=u,
)
- lam.create_account(account=account)
- assert lam.get_account(f"test:bp_wallet:{u}") == account
- assert lam.get_account_or_create(account) == account
+ ledger_account_manager.create_account(account=account)
+ assert ledger_account_manager.get_account(f"test:bp_wallet:{u}") == account
+ assert ledger_account_manager.get_account_or_create(account) == account
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
index 294d092..27ddc29 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
@@ -1,7 +1,11 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -11,27 +15,35 @@ from redis.lock import Lock
from generalresearch.currency import USDCent
from generalresearch.managers.base import Permission
-from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerTransactionFlagAlreadyExistsError,
LedgerTransactionConditionFailedError,
- LedgerTransactionReleaseLockError,
LedgerTransactionCreateError,
+ LedgerTransactionFlagAlreadyExistsError,
+ LedgerTransactionReleaseLockError,
)
from generalresearch.managers.thl.ledger_manager.ledger import LedgerTransaction
-from generalresearch.models import Source
+from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import PayoutStatus
from generalresearch.models.thl.ledger import Direction, TransactionType
from generalresearch.models.thl.session import (
- Wall,
+ Session,
Status,
StatusCode1,
- Session,
+ Wall,
)
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.redis_helper import RedisConfig
+if TYPE_CHECKING:
+ from generalresearch.currency import LedgerCurrency
+ from generalresearch.managers.thl.payout import (
+ BrokerageProductPayoutEventManager,
+ BusinessPayoutEventManager,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+
def broken_acquire(self, *args, **kwargs):
raise redis.exceptions.TimeoutError("Simulated timeout during acquire")
@@ -42,20 +54,18 @@ def broken_release(self, *args, **kwargs):
class TestThlLedgerManagerBPPayout:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
def test_create_tx_with_bp_payment(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
caplog,
- thl_lm,
- delete_ledger_db,
+ thl_ledger_manager: ThlLedgerManager,
):
- delete_ledger_db()
- create_main_accounts()
-
- now = datetime.now(timezone.utc) - timedelta(hours=1)
+ now = datetime.now(UTC) - timedelta(hours=1)
user: User = user_factory(product=product_user_wallet_no)
wall1 = Wall(
@@ -69,31 +79,29 @@ class TestThlLedgerManagerBPPayout:
started=now,
finished=now + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
+ _, _, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": now + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=now + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started)
lock_key = f"test:bp_payout:{user.product.id}"
- flag_name = f"{thl_lm.cache_prefix}:transaction_flag:{lock_key}"
- thl_lm.redis_client.delete(flag_name)
+ flag_name = f"{thl_ledger_manager.cache_prefix}:transaction_flag:{lock_key}"
+ thl_ledger_manager.redis_client.delete(flag_name)
payoutevent_uuid = uuid4().hex
- thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=user.product,
amount=USDCent(200),
created=now,
@@ -101,7 +109,7 @@ class TestThlLedgerManagerBPPayout:
)
payoutevent_uuid = uuid4().hex
- thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=user.product,
amount=USDCent(200),
created=now + timedelta(minutes=2),
@@ -109,13 +117,15 @@ class TestThlLedgerManagerBPPayout:
payoutevent_uuid=payoutevent_uuid,
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- assert 170 == thl_lm.get_account_balance(bp_wallet_account)
- assert 200 == thl_lm.get_account_balance(cash)
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ assert 170 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 200 == thl_ledger_manager.get_account_balance(cash)
with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError):
- thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
user.product,
amount=USDCent(200),
created=now + timedelta(minutes=2),
@@ -125,19 +135,21 @@ class TestThlLedgerManagerBPPayout:
)
payoutevent_uuid = uuid4().hex
- with caplog.at_level(logging.INFO):
- with pytest.raises(LedgerTransactionConditionFailedError):
- thl_lm.create_tx_bp_payout(
- user.product,
- amount=USDCent(10_000),
- created=now + timedelta(minutes=2),
- skip_one_per_day_check=True,
- skip_wallet_balance_check=False,
- payoutevent_uuid=payoutevent_uuid,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ user.product,
+ amount=USDCent(10_000),
+ created=now + timedelta(minutes=2),
+ skip_one_per_day_check=True,
+ skip_wallet_balance_check=False,
+ payoutevent_uuid=payoutevent_uuid,
+ )
assert "failed condition check balance:" in caplog.text
- thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=user.product,
amount=USDCent(10_00),
created=now + timedelta(minutes=2),
@@ -145,20 +157,26 @@ class TestThlLedgerManagerBPPayout:
skip_wallet_balance_check=True,
payoutevent_uuid=payoutevent_uuid,
)
- assert 170 - 1000 == thl_lm.get_account_balance(bp_wallet_account)
+ assert 170 - 1000 == thl_ledger_manager.get_account_balance(bp_wallet_account)
- def test_create_tx(self, product, caplog, thl_lm, currency):
+ def test_create_tx(
+ self,
+ product: Product,
+ caplog,
+ thl_ledger_manager: ThlLedgerManager,
+ currency: LedgerCurrency,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
# Create a BP Payout for a Product without any activity. By issuing,
# the skip_* checks, we should be able to force it to work, and will
# then ultimately result in a negative balance
- tx = thl_lm.create_tx_bp_payout(
+ tx = thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
skip_wallet_balance_check=True,
skip_one_per_day_check=True,
skip_flag_check=True,
@@ -177,31 +195,38 @@ class TestThlLedgerManagerBPPayout:
# Check the Product's balance, it should be negative the amount that was
# paid out. That's because the Product earned nothing.. and then was
# sent something.
- balance = thl_lm.get_account_balance(
- account=thl_lm.get_account_or_create_bp_wallet(product=product)
+ balance = thl_ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
)
assert balance == int(rand_amount) * -1
# Test some basic assertions
- with caplog.at_level(logging.INFO):
- with pytest.raises(expected_exception=Exception):
- thl_lm.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(expected_exception=LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
- def test_create_tx_redis_failure(self, product, thl_web_rw, thl_lm):
+ def test_create_tx_redis_failure(
+ self,
+ product: Product,
+ thl_web_rw: PostgresConfig,
+ thl_ledger_manager: ThlLedgerManager,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
- thl_lm.create_tx_plug_bp_wallet(
+ thl_ledger_manager.create_tx_plug_bp_wallet(
product, rand_amount, now, direction=Direction.CREDIT
)
@@ -222,43 +247,49 @@ class TestThlLedgerManagerBPPayout:
)
with pytest.raises(expected_exception=Exception) as e:
- tx = thl_lm_redis_0.create_tx_bp_payout(
+ thl_lm_redis_0.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
assert e.type is redis.exceptions.TimeoutError
# No txs were created
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 0
- def test_create_tx_multiple_per_day(self, product, thl_lm):
+ def test_create_tx_multiple_per_day(
+ self, product: Product, thl_ledger_manager: ThlLedgerManager
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
- thl_lm.create_tx_plug_bp_wallet(
+ thl_ledger_manager.create_tx_plug_bp_wallet(
product, rand_amount * USDCent(2), now, direction=Direction.CREDIT
)
- tx = thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
# Try to create another
# Will fail b/c it has the same payout event uuid
with pytest.raises(expected_exception=Exception) as e:
- tx = thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
assert e.type is LedgerTransactionFlagAlreadyExistsError
@@ -266,251 +297,261 @@ class TestThlLedgerManagerBPPayout:
# Will fail due to multiple per day
payoutevent_uuid2 = uuid4().hex
with pytest.raises(expected_exception=Exception) as e:
- tx = thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid2,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
assert e.type is LedgerTransactionConditionFailedError
assert str(e.value) == ">1 tx per day"
# Make it run by skipping one per day check
- tx = thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid2,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
skip_one_per_day_check=True,
)
- def test_create_tx_redis_lock_release_error(self, product, thl_lm):
+ def test_create_tx_redis_lock_release_error(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ monkeypatch: pytest.MonkeyPatch,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
- now = datetime.now(tz=timezone.utc)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
+ now = datetime.now(tz=UTC)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
- thl_lm.create_tx_plug_bp_wallet(
+ thl_ledger_manager.create_tx_plug_bp_wallet(
product, rand_amount * USDCent(2), now, direction=Direction.CREDIT
)
- original_acquire = Lock.acquire
- original_release = Lock.release
- Lock.acquire = broken_acquire
-
# Create TX will fail on lock enter, no tx will actually get created
- with pytest.raises(expected_exception=Exception) as e:
- tx = thl_lm.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
- )
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "acquire", broken_acquire)
+ with pytest.raises(expected_exception=Exception) as e:
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=payoutevent_uuid,
+ created=datetime.now(tz=UTC),
+ )
assert e.type is LedgerTransactionCreateError
assert str(e.value) == "Redis error: Simulated timeout during acquire"
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 0
- Lock.acquire = original_acquire
- Lock.release = broken_release
-
# Create TX will fail on lock exit, after the tx was created!
- with pytest.raises(expected_exception=Exception) as e:
- tx = thl_lm.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
- )
- assert e.type is LedgerTransactionReleaseLockError
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "release", broken_release)
+ with pytest.raises(LedgerTransactionReleaseLockError) as e:
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=payoutevent_uuid,
+ created=datetime.now(tz=UTC),
+ )
assert str(e.value) == "Redis error: Simulated timeout during release"
# Transaction was still created!
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 1
- Lock.release = original_release
class TestPayoutEventManagerBPPayout:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
- def test_create(self, product, thl_lm, brokerage_product_payout_event_manager):
+ def test_create(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.utc)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
- assert thl_lm.get_account_balance(bp_wallet_account) == 0
- thl_lm.create_tx_plug_bp_wallet(
+ now = datetime.now(tz=UTC)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0
+ thl_ledger_manager.create_tx_plug_bp_wallet(
product, rand_amount, now, direction=Direction.CREDIT
)
- assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ bpe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
+ ext_ref_id=uuid4().hex,
)
+ bp_pe = bpe.bp_payouts[0]
assert brokerage_product_payout_event_manager.check_for_ledger_tx(
- thl_ledger_manager=thl_lm,
- product_id=product.id,
- amount=rand_amount,
- payout_event=pe,
+ thl_ledger_manager=thl_ledger_manager,
+ payout_event=bp_pe,
)
- assert thl_lm.get_account_balance(bp_wallet_account) == 0
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0
def test_create_with_redis_error(
- self, product, caplog, thl_lm, brokerage_product_payout_event_manager
+ self,
+ product: Product,
+ caplog,
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ monkeypatch: pytest.MonkeyPatch,
):
caplog.set_level("WARNING")
- original_acquire = Lock.acquire
- original_release = Lock.release
+ ext_ref_id = uuid4().hex
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.utc)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
- assert thl_lm.get_account_balance(bp_wallet_account) == 0
- thl_lm.create_tx_plug_bp_wallet(
- product, rand_amount, now, direction=Direction.CREDIT
+ now = datetime.now(tz=UTC)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0
+ thl_ledger_manager.create_tx_plug_bp_wallet(
+ product=product, amount=rand_amount, created=now, direction=Direction.CREDIT
)
- assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
# Will fail on lock enter, no tx will actually get created
- Lock.acquire = broken_acquire
- with pytest.raises(expected_exception=Exception) as e:
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- assert e.type is LedgerTransactionCreateError
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "acquire", broken_acquire)
+ with pytest.raises(LedgerTransactionCreateError) as e:
+ business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ product=product,
+ created=now,
+ amount=rand_amount,
+ ext_ref_id=ext_ref_id,
+ )
assert str(e.value) == "Redis error: Simulated timeout during acquire"
- assert any(
- "Simulated timeout during acquire. No ledger tx was created" in m
- for m in caplog.messages
- )
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
# One payout event is created, status is failed, and no ledger txs exist
assert len(txs) == 0
pes = (
brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm, product_uuids=[product.id]
+ product_uuids=[product.id]
)
)
assert len(pes) == 1
assert pes[0].status == PayoutStatus.FAILED
pe = pes[0]
- # Fix the redis method
- Lock.acquire = original_acquire
-
# Try to fix the failed payout, by trying ledger tx again
brokerage_product_payout_event_manager.retry_create_bp_payout_event_tx(
- product=product, thl_ledger_manager=thl_lm, payout_event_uuid=pe.uuid
+ product=product,
+ thl_ledger_manager=thl_ledger_manager,
+ bp_pe=pe,
+ )
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
)
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 1
- assert thl_lm.get_account_balance(bp_wallet_account) == 0
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0
# And then try to run it again, it'll fail because a payout event with the same info exists
- with pytest.raises(expected_exception=Exception) as e:
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ with pytest.raises(expected_exception=ValueError) as e:
+ pe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
+ ext_ref_id=ext_ref_id,
)
- assert e.type is ValueError
- assert "Payout event already exists!" in str(e.value)
+ assert (
+ "Cannot create a BusinessPayoutEvent with an existing transaction_id"
+ in str(e.value)
+ )
# We wouldn't do this in practice, because this is paying out the BP again, but
# we can if want to.
- # Change the timestamp so it'll create a new payout event
- now = datetime.now(tz=timezone.utc)
- with pytest.raises(LedgerTransactionConditionFailedError) as e:
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- # But it will fail due to 1 per day check
- assert str(e.value) == ">1 tx per day"
- pe = brokerage_product_payout_event_manager.get_by_uuid(e.value.pe_uuid)
- assert pe.status == PayoutStatus.FAILED
-
- # And if we really want to, we can make it again
- now = datetime.now(tz=timezone.utc)
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ # Change the ext_ref_id so it'll create a new payout event
+ pe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
- skip_one_per_day_check=True,
- skip_wallet_balance_check=True,
+ ext_ref_id=uuid4().hex,
)
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 2
# since they were paid twice
- assert thl_lm.get_account_balance(bp_wallet_account) == 0 - rand_amount
-
- Lock.release = original_release
- Lock.acquire = original_acquire
+ assert (
+ thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 - rand_amount
+ )
def test_create_with_redis_error_release(
- self, product, caplog, thl_lm, brokerage_product_payout_event_manager
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ monkeypatch: pytest.MonkeyPatch,
+ caplog: pytest.LogCaptureFixture,
):
caplog.set_level("WARNING")
- original_release = Lock.release
-
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.utc)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ now = datetime.now(tz=UTC)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
- assert thl_lm.get_account_balance(bp_wallet_account) == 0
- thl_lm.create_tx_plug_bp_wallet(
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == 0
+ thl_ledger_manager.create_tx_plug_bp_wallet(
product, rand_amount, now, direction=Direction.CREDIT
)
- assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
# Will fail on lock exit, after the tx was created!
# But it'll see that the tx was created and so everything will be fine
- Lock.release = broken_release
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- assert any(
- "Simulated timeout during release but ledger tx exists" in m
- for m in caplog.messages
- )
+ caplog.clear()
+ with monkeypatch.context() as m, caplog.at_level("WARNING"):
+ m.setattr(Lock, "release", broken_release)
+ business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ product=product,
+ created=now,
+ amount=rand_amount,
+ ext_ref_id=uuid4().hex,
+ )
+ assert "Redis error: Simulated timeout during release" in caplog.messages
- txs = thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet_account.uuid)
+ txs = thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet_account.uuid
+ )
txs = [tx for tx in txs if tx.metadata["tx_type"] != "plug"]
assert len(txs) == 1
pes = (
brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm, product_uuids=[product.uuid]
+ product_uuids=[product.uuid]
)
)
assert len(pes) == 1
assert pes[0].status == PayoutStatus.COMPLETE
- Lock.release = original_release
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py
index 31c7107..aa3b378 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py
@@ -1,112 +1,137 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.currency import USDCent
+from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerTransactionConditionFailedError,
+)
from generalresearch.managers.thl.ledger_manager.ledger import (
LedgerTransaction,
)
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
WALL_ALLOWED_STATUS_STATUS_CODE,
)
-from generalresearch.models.thl.ledger import Direction
-from generalresearch.models.thl.ledger import TransactionType
+from generalresearch.models.thl.ledger import (
+ Direction,
+ TransactionType,
+)
+from generalresearch.models.thl.payout import UserPayoutEvent
from generalresearch.models.thl.product import (
PayoutConfig,
PayoutTransformation,
+ Product,
UserWalletConfig,
)
from generalresearch.models.thl.session import (
- Wall,
+ Session,
Status,
StatusCode1,
- Session,
+ Wall,
WallAdjustedStatus,
)
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.payout import UserPayoutEvent
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.currency import LedgerCurrency
+ from generalresearch.managers.thl.ledger_manager.ledger import (
+ LedgerManager,
+ )
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.thl.ledger import (
+ LedgerAccount,
+ )
+ from generalresearch.models.thl.user import User
logger = logging.getLogger("LedgerManager")
class TestThlLedgerTxManager:
+ @pytest.fixture(autouse=True)
+ def setup(self, delete_ledger_db, create_main_accounts):
+ delete_ledger_db()
+ create_main_accounts()
def test_create_tx_task_complete(
self,
- wall,
- user,
- account_revenue_task_complete,
- create_main_accounts,
- thl_lm,
- lm,
+ wall: Wall,
+ user: User,
+ account_revenue_task_complete: LedgerAccount,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
- create_main_accounts()
- tx = thl_lm.create_tx_task_complete(wall=wall, user=user)
+ tx = thl_ledger_manager.create_tx_task_complete(wall=wall, user=user)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.created == tx.created
def test_create_tx_task_complete_(
- self, wall, user, account_revenue_task_complete, thl_lm, lm
+ self,
+ wall: Wall,
+ user: User,
+ account_revenue_task_complete: LedgerAccount,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
- tx = thl_lm.create_tx_task_complete_(wall=wall, user=user)
+ tx = thl_ledger_manager.create_tx_task_complete_(wall=wall, user=user)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.created == tx.created
def test_create_tx_bp_payment(
self,
- session_factory,
- user,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- lm,
- session_manager,
+ session_factory: Callable[..., Session],
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ session_manager: SessionManager,
):
- delete_ledger_db()
- create_main_accounts()
+
s1 = session_factory(user=user)
- status, status_code_1 = s1.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
+ _, status_code_1 = s1.determine_session_status()
+ _, _, bp_pay, user_pay = s1.determine_payments()
session_manager.finish_with_status(
session=s1,
status=Status.COMPLETE,
status_code_1=status_code_1,
- finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10),
+ finished=datetime.now(tz=UTC) + timedelta(minutes=10),
payout=bp_pay,
user_payout=user_pay,
)
- tx = thl_lm.create_tx_bp_payment(session=s1)
+ tx = thl_ledger_manager.create_tx_bp_payment(session=s1)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.created == tx.created
def test_create_tx_bp_payment_amt(
self,
- session_factory,
- user_factory,
- product_manager,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- lm,
- session_manager,
+ session_factory: Callable[..., Session],
+ user_factory: Callable[..., User],
+ product_manager: ProductManager,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ session_manager: SessionManager,
+ product_factory: Callable[..., Product],
):
- delete_ledger_db()
- create_main_accounts()
- product = product_manager.create_dummy(
+
+ product = product_factory(
payout_config=PayoutConfig(
payout_transformation=PayoutTransformation(
f="payout_transformation_amt"
@@ -115,42 +140,41 @@ class TestThlLedgerTxManager:
user_wallet_config=UserWalletConfig(amt=True, enabled=True),
)
user = user_factory(product=product)
- s1 = session_factory(user=user, wall_req_cpi=Decimal("1"))
+ s1 = session_factory(user=user, wall_req_cpi=Decimal(1))
status, status_code_1 = s1.determine_session_status()
assert status == Status.COMPLETE
thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments(
- thl_ledger_manager=thl_lm
+ thl_ledger_manager=thl_ledger_manager
)
print(thl_net, commission_amount, bp_pay, user_pay)
session_manager.finish_with_status(
session=s1,
status=Status.COMPLETE,
status_code_1=status_code_1,
- finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10),
+ finished=datetime.now(tz=UTC) + timedelta(minutes=10),
payout=bp_pay,
user_payout=user_pay,
)
- tx = thl_lm.create_tx_bp_payment(session=s1)
+ tx = thl_ledger_manager.create_tx_bp_payment(session=s1)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.created == tx.created
def test_create_tx_bp_payment_(
self,
- session_factory,
- user,
- create_main_accounts,
- thl_lm,
- lm,
- session_manager,
- utc_hour_ago,
+ session_factory: Callable[..., Session],
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
):
s1 = session_factory(user=user)
status, status_code_1 = s1.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
+ _, _, bp_pay, user_pay = s1.determine_payments()
session_manager.finish_with_status(
session=s1,
status=status,
@@ -161,14 +185,19 @@ class TestThlLedgerTxManager:
)
s1.determine_payments()
- tx = thl_lm.create_tx_bp_payment_(session=s1)
+ tx = thl_ledger_manager.create_tx_bp_payment_(session=s1)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.created == tx.created
def test_create_tx_task_adjustment(
- self, wall_factory, session, user, create_main_accounts, thl_lm, lm
+ self,
+ wall_factory: Callable[..., Wall],
+ bare_session: Session,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
"""Create Wall event Complete, and Create a Tx Task Adjustment
@@ -176,29 +205,34 @@ class TestThlLedgerTxManager:
the transaction comes back with balanced amounts, and that
the name of the Source is in the Tx description
"""
-
wall_status = Status.COMPLETE
- wall: Wall = wall_factory(session=session, wall_status=wall_status)
+ wall: Wall = wall_factory(session=bare_session, wall_status=wall_status)
- tx = thl_lm.create_tx_task_adjustment(wall=wall, user=user)
+ tx = thl_ledger_manager.create_tx_task_adjustment(wall=wall, user=user)
assert isinstance(tx, LedgerTransaction)
- res = lm.get_tx_by_id(transaction_id=tx.id)
+ res = ledger_manager.get_tx_by_id(transaction_id=tx.id)
assert res.entries[0].amount == int(wall.cpi * 100)
assert res.entries[1].amount == int(wall.cpi * 100)
assert wall.source.name in res.ext_description
assert res.created == tx.created
- def test_create_tx_bp_adjustment(self, session, user, caplog, thl_lm, lm):
+ def test_create_tx_bp_adjustment(
+ self,
+ session: Session,
+ user: User,
+ caplog,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ ):
status, status_code_1 = session.determine_session_status()
thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
- # The default session fixture is just an unfinished wall event
assert len(session.wall_events) == 1
assert session.finished is None
- assert status == Status.TIMEOUT
+ assert status == Status.FAIL
assert status_code_1 in list(
- WALL_ALLOWED_STATUS_STATUS_CODE.get(Status.TIMEOUT, {})
+ WALL_ALLOWED_STATUS_STATUS_CODE.get(Status.FAIL, {})
)
assert thl_net == Decimal(0)
assert commission_amount == Decimal(0)
@@ -208,28 +242,32 @@ class TestThlLedgerTxManager:
# Update the finished timestamp, but nothing else. This means that
# there is no financial changes needed
session.update(
- **{
- "finished": datetime.now(tz=timezone.utc) + timedelta(minutes=10),
- }
+ finished=datetime.now(tz=UTC) + timedelta(minutes=10), status=Status.FAIL
)
assert session.finished
with caplog.at_level(logging.INFO):
- tx = thl_lm.create_tx_bp_adjustment(session=session)
+ tx = thl_ledger_manager.create_tx_bp_adjustment(session=session)
assert tx is None
assert "No transactions needed." in caplog.text
- def test_create_tx_bp_payout(self, product, caplog, thl_lm, currency):
+ def test_create_tx_bp_payout(
+ self,
+ product: Product,
+ caplog,
+ thl_ledger_manager: ThlLedgerManager,
+ currency: LedgerCurrency,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
# Create a BP Payout for a Product without any activity. By issuing,
# the skip_* checks, we should be able to force it to work, and will
# then ultimately result in a negative balance
- tx = thl_lm.create_tx_bp_payout(
+ tx = thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
skip_wallet_balance_check=True,
skip_one_per_day_check=True,
skip_flag_check=True,
@@ -240,7 +278,7 @@ class TestThlLedgerTxManager:
assert tx.ext_description == "BP Payout"
assert (
tx.tag
- == f"{thl_lm.currency.value}:{TransactionType.BP_PAYOUT.value}:{payoutevent_uuid}"
+ == f"{thl_ledger_manager.currency.value}:{TransactionType.BP_PAYOUT.value}:{payoutevent_uuid}"
)
assert tx.entries[0].amount == rand_amount
assert tx.entries[1].amount == rand_amount
@@ -248,35 +286,42 @@ class TestThlLedgerTxManager:
# Check the Product's balance, it should be negative the amount that was
# paid out. That's because the Product earned nothing.. and then was
# sent something.
- balance = thl_lm.get_account_balance(
- account=thl_lm.get_account_or_create_bp_wallet(product=product)
+ balance = thl_ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
)
assert balance == int(rand_amount) * -1
# Test some basic assertions
- with caplog.at_level(logging.INFO):
- with pytest.raises(expected_exception=Exception):
- thl_lm.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(expected_exception=LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
- def test_create_tx_bp_payout_(self, product, thl_lm, lm, currency):
+ def test_create_tx_bp_payout_(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ currency: LedgerCurrency,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
# Create a BP Payout for a Product without any activity.
- tx = thl_lm.create_tx_bp_payout_(
+ tx = thl_ledger_manager.create_tx_bp_payout_(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
# Check the basic attributes
@@ -290,17 +335,21 @@ class TestThlLedgerTxManager:
assert tx.entries[1].amount == rand_amount
def test_create_tx_plug_bp_wallet(
- self, product, create_main_accounts, thl_lm, lm, currency
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
"""A BP Wallet "plug" is a way to makeup discrepancies and simply
add or remove money
"""
rand_amount: USDCent = USDCent(randint(100, 1_000))
- tx = thl_lm.create_tx_plug_bp_wallet(
+ tx = thl_ledger_manager.create_tx_plug_bp_wallet(
product=product,
amount=rand_amount,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
direction=Direction.DEBIT,
skip_flag_check=False,
)
@@ -309,13 +358,17 @@ class TestThlLedgerTxManager:
# We issued the BP money they didn't earn, so now they have a
# negative balance
- balance = thl_lm.get_account_balance(
- account=thl_lm.get_account_or_create_bp_wallet(product=product)
+ balance = thl_ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
)
assert balance == int(rand_amount) * -1
def test_create_tx_plug_bp_wallet_(
- self, product, create_main_accounts, thl_lm, lm, currency
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
"""A BP Wallet "plug" is a way to fix discrepancies and simply
add or remove money.
@@ -325,10 +378,10 @@ class TestThlLedgerTxManager:
"""
rand_amount: USDCent = USDCent(randint(100, 1_000))
- tx = thl_lm.create_tx_plug_bp_wallet_(
+ tx = thl_ledger_manager.create_tx_plug_bp_wallet_(
product=product,
amount=rand_amount,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
direction=Direction.DEBIT,
)
@@ -336,32 +389,32 @@ class TestThlLedgerTxManager:
# We issued the BP money they didn't earn, so now they have a
# negative balance
- balance = thl_lm.get_account_balance(
- account=thl_lm.get_account_or_create_bp_wallet(product=product)
+ balance = thl_ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
)
assert balance == int(rand_amount) * -1
# Issue a positive one now, and confirm the balance goes positive
- thl_lm.create_tx_plug_bp_wallet_(
+ thl_ledger_manager.create_tx_plug_bp_wallet_(
product=product,
amount=rand_amount + rand_amount,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
direction=Direction.CREDIT,
)
- balance = thl_lm.get_account_balance(
- account=thl_lm.get_account_or_create_bp_wallet(product=product)
+ balance = thl_ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
)
assert balance == int(rand_amount)
def test_create_tx_user_payout_request(
self,
- user,
- product_user_wallet_yes,
- user_factory,
- delete_df_collection,
- thl_lm,
- lm,
- currency,
+ user: User,
+ product_user_wallet_yes: Product,
+ user_factory: Callable[..., User],
+ delete_df_collection: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
pe = UserPayoutEvent(
uuid=uuid4().hex,
@@ -374,7 +427,7 @@ class TestThlLedgerTxManager:
# The default user fixture uses a product that doesn't have wallet
# mode enabled
with pytest.raises(expected_exception=AssertionError):
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
skip_flag_check=True,
@@ -385,12 +438,12 @@ class TestThlLedgerTxManager:
u2 = user_factory(product=product_user_wallet_yes)
# User's pre-balance is 0 because no activity has occurred yet
- pre_balance = lm.get_account_balance(
- account=thl_lm.get_account_or_create_user_wallet(user=u2)
+ pre_balance = ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_user_wallet(user=u2)
)
assert pre_balance == 0
- tx = thl_lm.create_tx_user_payout_request(
+ tx = thl_ledger_manager.create_tx_user_payout_request(
user=u2,
payout_event=pe,
skip_flag_check=True,
@@ -411,21 +464,19 @@ class TestThlLedgerTxManager:
# Post balance is -$5.00 because it comes out of the wallet before
# it's Approved or Completed
- post_balance = lm.get_account_balance(
- account=thl_lm.get_account_or_create_user_wallet(user=u2)
+ post_balance = ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_user_wallet(user=u2)
)
assert post_balance == -500
def test_create_tx_user_payout_request_(
self,
- user,
- product_user_wallet_yes,
- user_factory,
- delete_ledger_db,
- thl_lm,
- lm,
+ user: User,
+ product_user_wallet_yes: Product,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
- delete_ledger_db()
pe = UserPayoutEvent(
uuid=uuid4().hex,
@@ -436,36 +487,32 @@ class TestThlLedgerTxManager:
)
rand_description = uuid4().hex
- tx = thl_lm.create_tx_user_payout_request_(
+ tx = thl_ledger_manager.create_tx_user_payout_request_(
user=user, payout_event=pe, description=rand_description
)
assert tx.ext_description == rand_description
- post_balance = lm.get_account_balance(
- account=thl_lm.get_account_or_create_user_wallet(user=user)
+ post_balance = ledger_manager.get_account_balance(
+ account=thl_ledger_manager.get_account_or_create_user_wallet(user=user)
)
assert post_balance == -500
def test_create_tx_user_payout_complete(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
# Ensure the user starts out with nothing...
- assert lm.get_account_balance(account=user_account) == 0
+ assert ledger_manager.get_account_balance(account=user_account) == 0
pe = UserPayoutEvent(
uuid=uuid4().hex,
@@ -477,7 +524,7 @@ class TestThlLedgerTxManager:
# Confirm it's not possible unless a request occurred happen
with pytest.raises(expected_exception=ValueError):
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user=user,
payout_event=pe,
fee_amount=None,
@@ -485,17 +532,19 @@ class TestThlLedgerTxManager:
)
# (1) Make a request first
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
skip_flag_check=True,
skip_wallet_balance_check=True,
)
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == rand_amount * -1
+ assert (
+ ledger_manager.get_account_balance(account=user_account) == rand_amount * -1
+ )
# (2) Complete the request
- tx = thl_lm.create_tx_user_payout_complete(
+ tx = thl_ledger_manager.create_tx_user_payout_complete(
user=user,
payout_event=pe,
fee_amount=Decimal(0),
@@ -508,18 +557,19 @@ class TestThlLedgerTxManager:
# The amount that comes out of the user wallet doesn't change after
# it's approved becuase it's already been withdrawn
- assert lm.get_account_balance(account=user_account) == rand_amount * -1
+ assert (
+ ledger_manager.get_account_balance(account=user_account) == rand_amount * -1
+ )
def test_create_tx_user_payout_complete_(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
pe = UserPayoutEvent(
@@ -531,7 +581,7 @@ class TestThlLedgerTxManager:
)
# (1) Make a request first
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
skip_flag_check=True,
@@ -541,12 +591,14 @@ class TestThlLedgerTxManager:
# (2) Complete the request
rand_desc = uuid4().hex
- bp_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
product=user.product, expense_name="paypal"
)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
- tx = thl_lm.create_tx_user_payout_complete_(
+ tx = thl_ledger_manager.create_tx_user_payout_complete_(
user=user,
payout_event=pe,
fee_amount=Decimal("0.00"),
@@ -555,19 +607,20 @@ class TestThlLedgerTxManager:
description=rand_desc,
)
assert tx.ext_description == rand_desc
- assert lm.get_account_balance(account=user_account) == rand_amount * -1
+ assert (
+ ledger_manager.get_account_balance(account=user_account) == rand_amount * -1
+ )
def test_create_tx_user_payout_cancelled(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
pe = UserPayoutEvent(
@@ -579,17 +632,19 @@ class TestThlLedgerTxManager:
)
# (1) Make a request first
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
skip_flag_check=True,
skip_wallet_balance_check=True,
)
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == rand_amount * -1
+ assert (
+ ledger_manager.get_account_balance(account=user_account) == rand_amount * -1
+ )
# (2) Cancel the request
- tx = thl_lm.create_tx_user_payout_cancelled(
+ tx = thl_ledger_manager.create_tx_user_payout_cancelled(
user=user,
payout_event=pe,
skip_flag_check=False,
@@ -600,19 +655,18 @@ class TestThlLedgerTxManager:
assert isinstance(tx, LedgerTransaction)
# Assert the balance comes back to 0 after it was cancelled
- assert lm.get_account_balance(account=user_account) == 0
+ assert ledger_manager.get_account_balance(account=user_account) == 0
def test_create_tx_user_payout_cancelled_(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
pe = UserPayoutEvent(
@@ -624,43 +678,44 @@ class TestThlLedgerTxManager:
)
# (1) Make a request first
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
skip_flag_check=True,
skip_wallet_balance_check=True,
)
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == rand_amount * -1
+ assert (
+ ledger_manager.get_account_balance(account=user_account) == rand_amount * -1
+ )
# (2) Cancel the request
rand_desc = uuid4().hex
- tx = thl_lm.create_tx_user_payout_cancelled_(
+ tx = thl_ledger_manager.create_tx_user_payout_cancelled_(
user=user, payout_event=pe, description=rand_desc
)
assert isinstance(tx, LedgerTransaction)
assert tx.ext_description == rand_desc
- assert lm.get_account_balance(account=user_account) == 0
+ assert ledger_manager.get_account_balance(account=user_account) == 0
def test_create_tx_user_bonus(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
rand_ref_uuid = uuid4().hex
rand_desc = uuid4().hex
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == 0
+ assert ledger_manager.get_account_balance(account=user_account) == 0
- tx = thl_lm.create_tx_user_bonus(
+ tx = thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(rand_amount / 100),
ref_uuid=rand_ref_uuid,
@@ -668,44 +723,47 @@ class TestThlLedgerTxManager:
skip_flag_check=True,
)
assert tx.ext_description == rand_desc
- assert tx.tag == f"{thl_lm.currency.value}:user_bonus:{rand_ref_uuid}"
+ assert (
+ tx.tag == f"{thl_ledger_manager.currency.value}:user_bonus:{rand_ref_uuid}"
+ )
assert tx.entries[0].amount == rand_amount
assert tx.entries[1].amount == rand_amount
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == rand_amount
+ assert ledger_manager.get_account_balance(account=user_account) == rand_amount
def test_create_tx_user_bonus_(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
user: User = user_factory(product=product_user_wallet_yes)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
rand_amount = randint(100, 1_000)
rand_ref_uuid = uuid4().hex
rand_desc = uuid4().hex
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == 0
+ assert ledger_manager.get_account_balance(account=user_account) == 0
- tx = thl_lm.create_tx_user_bonus_(
+ tx = thl_ledger_manager.create_tx_user_bonus_(
user=user,
amount=Decimal(rand_amount / 100),
ref_uuid=rand_ref_uuid,
description=rand_desc,
)
assert tx.ext_description == rand_desc
- assert tx.tag == f"{thl_lm.currency.value}:user_bonus:{rand_ref_uuid}"
+ assert (
+ tx.tag == f"{thl_ledger_manager.currency.value}:user_bonus:{rand_ref_uuid}"
+ )
assert tx.entries[0].amount == rand_amount
assert tx.entries[1].amount == rand_amount
# Assert the balance came out of their user wallet
- assert lm.get_account_balance(account=user_account) == rand_amount
+ assert ledger_manager.get_account_balance(account=user_account) == rand_amount
class TestThlLedgerTxManagerFlows:
@@ -713,12 +771,19 @@ class TestThlLedgerTxManagerFlows:
examples
"""
- def test_create_tx_task_complete(
- self, user, create_main_accounts, thl_lm, lm, currency, delete_ledger_db
- ):
+ @pytest.fixture(autouse=True)
+ def setup(self, delete_ledger_db, create_main_accounts):
delete_ledger_db()
create_main_accounts()
+ def test_create_tx_task_complete(
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ ):
+
wall1 = Wall(
user_id=1,
source=Source.DYNATA,
@@ -727,10 +792,12 @@ class TestThlLedgerTxManagerFlows:
session_id=1,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
+ )
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
wall2 = Wall(
user_id=1,
@@ -740,41 +807,43 @@ class TestThlLedgerTxManagerFlows:
session_id=1,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
+ )
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall2, user=user, created=wall2.started
)
- thl_lm.create_tx_task_complete(wall=wall2, user=user, created=wall2.started)
- cash = thl_lm.get_account_cash()
- revenue = thl_lm.get_account_task_complete_revenue()
+ cash = thl_ledger_manager.get_account_cash()
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
- assert lm.get_account_balance(cash) == 123 + 321
- assert lm.get_account_balance(revenue) == 123 + 321
- assert lm.check_ledger_balanced()
+ assert ledger_manager.get_account_balance(cash) == 123 + 321
+ assert ledger_manager.get_account_balance(revenue) == 123 + 321
+ assert ledger_manager.check_ledger_balanced()
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="d"
)
== 123
)
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="f"
)
== 321
)
assert (
- lm.get_account_filtered_balance(
+ ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="x"
)
== 0
)
assert (
- thl_lm.get_account_filtered_balance(
+ thl_ledger_manager.get_account_filtered_balance(
account=revenue,
metadata_key="thl_wall",
metadata_value=wall1.uuid,
@@ -783,7 +852,11 @@ class TestThlLedgerTxManagerFlows:
)
def test_create_transaction_task_complete_1_cent(
- self, user, create_main_accounts, thl_lm, lm, currency
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
wall1 = Wall(
user_id=1,
@@ -793,10 +866,10 @@ class TestThlLedgerTxManagerFlows:
session_id=1,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
@@ -804,17 +877,13 @@ class TestThlLedgerTxManagerFlows:
def test_create_transaction_bp_payment(
self,
- user,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
- delete_ledger_db,
- session_factory,
- utc_hour_ago,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ session_factory: Callable[..., Session],
+ utc_hour_ago: datetime,
):
- delete_ledger_db()
- create_main_accounts()
s1: Session = session_factory(
user=user,
@@ -824,50 +893,53 @@ class TestThlLedgerTxManagerFlows:
)
w1: Wall = s1.wall_events[0]
- tx = thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started)
+ tx = thl_ledger_manager.create_tx_task_complete(
+ wall=w1, user=user, created=w1.started
+ )
assert isinstance(tx, LedgerTransaction)
status, status_code_1 = s1.determine_session_status()
thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
s1.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": s1.started + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=s1.started + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
print(thl_net, commission_amount, bp_pay, user_pay)
- thl_lm.create_tx_bp_payment(session=s1, created=w1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started)
- revenue = thl_lm.get_account_task_complete_revenue()
- bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- bp_commission = thl_lm.get_account_or_create_bp_commission(product=user.product)
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
+ bp_commission = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=user.product
+ )
- assert 0 == lm.get_account_balance(account=revenue)
- assert 50 == lm.get_account_filtered_balance(
+ assert 0 == ledger_manager.get_account_balance(account=revenue)
+ assert 50 == ledger_manager.get_account_filtered_balance(
account=revenue,
metadata_key="source",
metadata_value=Source.TESTING,
)
- assert 48 == lm.get_account_balance(account=bp_wallet)
- assert 48 == lm.get_account_filtered_balance(
+ assert 48 == ledger_manager.get_account_balance(account=bp_wallet)
+ assert 48 == ledger_manager.get_account_filtered_balance(
account=bp_wallet,
metadata_key="thl_session",
metadata_value=s1.uuid,
)
- assert 2 == thl_lm.get_account_balance(account=bp_commission)
- assert thl_lm.check_ledger_balanced()
+ assert 2 == thl_ledger_manager.get_account_balance(account=bp_commission)
+ assert thl_ledger_manager.check_ledger_balanced()
def test_create_transaction_bp_payment_round(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
product_user_wallet_no.commission_pct = Decimal("0.085")
user: User = user_factory(product=product_user_wallet_no)
@@ -880,11 +952,11 @@ class TestThlLedgerTxManagerFlows:
session_id=3,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
@@ -893,24 +965,27 @@ class TestThlLedgerTxManagerFlows:
status, status_code_1 = session.determine_session_status()
thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": session.started + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=session.started + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
print(thl_net, commission_amount, bp_pay, user_pay)
- tx = thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ tx = thl_ledger_manager.create_tx_bp_payment(
+ session=session, created=wall1.started
+ )
assert isinstance(tx, LedgerTransaction)
def test_create_transaction_bp_payment_round2(
- self, delete_ledger_db, user, create_main_accounts, thl_lm, lm, currency
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
+
# user must be no user wallet
# e.g. session 869b5bfa47f44b4f81cd095ed01df2ff this fails if you dont round properly
@@ -922,34 +997,33 @@ class TestThlLedgerTxManagerFlows:
session_id=3,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
+ )
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
# thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": session.started + timedelta(minutes=10),
- "payout": Decimal("1.53"),
- "user_payout": Decimal("1.53"),
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=session.started + timedelta(minutes=10),
+ payout=Decimal("1.53"),
+ user_payout=Decimal("1.53"),
)
- thl_lm.create_tx_bp_payment(session=session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=session, created=wall1.started)
def test_create_transaction_bp_payment_round3(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
# e.g. session ___ fails b/c we rounded incorrectly
# before, and now we are off by a penny...
@@ -963,22 +1037,22 @@ class TestThlLedgerTxManagerFlows:
session_id=3,
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
+ )
+ thl_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
# thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": session.started + timedelta(minutes=10),
- "payout": Decimal("0.39"),
- "user_payout": Decimal("0.26"),
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=session.started + timedelta(minutes=10),
+ payout=Decimal("0.39"),
+ user_payout=Decimal("0.26"),
)
# with pytest.logs(logger, level=logging.WARNING) as cm:
# tx = thl_lm.create_transaction_bp_payment(session, created=wall1.started)
@@ -986,22 +1060,19 @@ class TestThlLedgerTxManagerFlows:
def test_create_transaction_bp_payment_user_wallet(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- session_manager,
- wall_manager,
- lm,
- session_factory,
- currency,
- utc_hour_ago,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
+ ledger_manager: LedgerManager,
+ session_factory: Callable[..., Session],
+ currency: LedgerCurrency,
+ utc_hour_ago: datetime,
):
- delete_ledger_db()
- create_main_accounts()
user: User = user_factory(product=product_user_wallet_yes)
+ assert isinstance(user.product, Product)
assert user.product.user_wallet_enabled
s1: Session = session_factory(
@@ -1013,10 +1084,12 @@ class TestThlLedgerTxManagerFlows:
)
w1: Wall = s1.wall_events[0]
- thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=w1, user=user, created=w1.started
+ )
status, status_code_1 = s1.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
+ _, _, bp_pay, user_pay = s1.determine_payments()
session_manager.finish_with_status(
session=s1,
status=status,
@@ -1025,55 +1098,59 @@ class TestThlLedgerTxManagerFlows:
payout=bp_pay,
user_payout=user_pay,
)
- thl_lm.create_tx_bp_payment(session=s1, created=w1.started)
+ thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started)
- revenue = thl_lm.get_account_task_complete_revenue()
- bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- bp_commission = thl_lm.get_account_or_create_bp_commission(product=user.product)
- user_wallet = thl_lm.get_account_or_create_user_wallet(user=user)
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
+ bp_commission = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=user.product
+ )
+ user_wallet = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
- assert 0 == thl_lm.get_account_balance(account=revenue)
- assert 50 == thl_lm.get_account_filtered_balance(
+ assert 0 == thl_ledger_manager.get_account_balance(account=revenue)
+ assert 50 == thl_ledger_manager.get_account_filtered_balance(
account=revenue,
metadata_key="source",
metadata_value=Source.TESTING,
)
- assert 48 - 19 == thl_lm.get_account_balance(account=bp_wallet)
- assert 48 - 19 == thl_lm.get_account_filtered_balance(
+ assert 48 - 19 == thl_ledger_manager.get_account_balance(account=bp_wallet)
+ assert 48 - 19 == thl_ledger_manager.get_account_filtered_balance(
account=bp_wallet,
metadata_key="thl_session",
metadata_value=s1.uuid,
)
- assert 2 == thl_lm.get_account_balance(bp_commission)
- assert 19 == thl_lm.get_account_balance(user_wallet)
- assert 19 == thl_lm.get_account_filtered_balance(
+ assert 2 == thl_ledger_manager.get_account_balance(bp_commission)
+ assert 19 == thl_ledger_manager.get_account_balance(user_wallet)
+ assert 19 == thl_ledger_manager.get_account_filtered_balance(
account=user_wallet,
metadata_key="thl_session",
metadata_value=s1.uuid,
)
- assert 0 == thl_lm.get_account_filtered_balance(
+ assert 0 == thl_ledger_manager.get_account_filtered_balance(
account=user_wallet, metadata_key="thl_session", metadata_value="x"
)
- assert thl_lm.check_ledger_balanced()
+ assert thl_ledger_manager.check_ledger_balanced()
class TestThlLedgerManagerAdj:
+ @pytest.fixture(autouse=True)
+ def setup(self, delete_ledger_db, create_main_accounts):
+ delete_ledger_db()
+ create_main_accounts()
def test_create_tx_task_adjustment(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- lm,
- utc_hour_ago,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ utc_hour_ago: datetime,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
user: User = user_factory(product=product_user_wallet_no)
@@ -1089,7 +1166,7 @@ class TestThlLedgerManagerAdj:
finished=utc_hour_ago + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall1, user, created=wall1.started)
+ thl_ledger_manager.create_tx_task_complete(wall1, user, created=wall1.started)
wall2 = Wall(
user_id=1,
@@ -1102,7 +1179,7 @@ class TestThlLedgerManagerAdj:
started=utc_hour_ago,
finished=utc_hour_ago + timedelta(seconds=1),
)
- thl_lm.create_tx_task_complete(wall2, user, created=wall2.started)
+ thl_ledger_manager.create_tx_task_complete(wall2, user, created=wall2.started)
wall1.update(
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
@@ -1110,24 +1187,26 @@ class TestThlLedgerManagerAdj:
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
print(wall1.get_cpi_after_adjustment())
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
- cash = thl_lm.get_account_cash()
- revenue = thl_lm.get_account_task_complete_revenue()
+ cash = thl_ledger_manager.get_account_cash()
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
- assert 123 + 321 - 123 == thl_lm.get_account_balance(account=cash)
- assert 123 + 321 - 123 == thl_lm.get_account_balance(account=revenue)
- assert thl_lm.check_ledger_balanced()
- assert 0 == thl_lm.get_account_filtered_balance(
+ assert 123 + 321 - 123 == thl_ledger_manager.get_account_balance(account=cash)
+ assert 123 + 321 - 123 == thl_ledger_manager.get_account_balance(
+ account=revenue
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
+ assert 0 == thl_ledger_manager.get_account_filtered_balance(
revenue, metadata_key="source", metadata_value="d"
)
- assert 321 == thl_lm.get_account_filtered_balance(
+ assert 321 == thl_ledger_manager.get_account_filtered_balance(
revenue, metadata_key="source", metadata_value="f"
)
- assert 0 == thl_lm.get_account_filtered_balance(
+ assert 0 == thl_ledger_manager.get_account_filtered_balance(
revenue, metadata_key="source", metadata_value="x"
)
- assert 123 - 123 == thl_lm.get_account_filtered_balance(
+ assert 123 - 123 == thl_ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="thl_wall", metadata_value=wall1.uuid
)
@@ -1138,46 +1217,42 @@ class TestThlLedgerManagerAdj:
adjusted_timestamp=utc_hour_ago + timedelta(minutes=45),
)
print(wall1.get_cpi_after_adjustment())
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
# and then run it again to make sure it does nothing
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
- cash = thl_lm.get_account_cash()
- revenue = thl_lm.get_account_task_complete_revenue()
+ cash = thl_ledger_manager.get_account_cash()
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
- assert 123 + 321 - 123 + 123 == thl_lm.get_account_balance(cash)
- assert 123 + 321 - 123 + 123 == thl_lm.get_account_balance(revenue)
- assert thl_lm.check_ledger_balanced()
- assert 123 == thl_lm.get_account_filtered_balance(
+ assert 123 + 321 - 123 + 123 == thl_ledger_manager.get_account_balance(cash)
+ assert 123 + 321 - 123 + 123 == thl_ledger_manager.get_account_balance(revenue)
+ assert thl_ledger_manager.check_ledger_balanced()
+ assert 123 == thl_ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="d"
)
- assert 321 == thl_lm.get_account_filtered_balance(
+ assert 321 == thl_ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="f"
)
- assert 0 == thl_lm.get_account_filtered_balance(
+ assert 0 == thl_ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="source", metadata_value="x"
)
- assert 123 - 123 + 123 == thl_lm.get_account_filtered_balance(
+ assert 123 - 123 + 123 == thl_ledger_manager.get_account_filtered_balance(
account=revenue, metadata_key="thl_wall", metadata_value=wall1.uuid
)
def test_create_tx_bp_adjustment(
self,
- user,
- product_user_wallet_no,
- create_main_accounts,
+ user: User,
+ product_user_wallet_no: Product,
caplog,
- thl_lm,
- lm,
- currency,
- session_manager,
- wall_manager,
- session_factory,
- utc_hour_ago,
- delete_ledger_db,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ utc_hour_ago: datetime,
):
- delete_ledger_db()
- create_main_accounts()
s1 = session_factory(
user=user,
@@ -1190,11 +1265,15 @@ class TestThlLedgerManagerAdj:
w1: Wall = s1.wall_events[0]
w2: Wall = s1.wall_events[1]
- thl_lm.create_tx_task_complete(wall=w1, user=user, created=w1.started)
- thl_lm.create_tx_task_complete(wall=w2, user=user, created=w2.started)
+ thl_ledger_manager.create_tx_task_complete(
+ wall=w1, user=user, created=w1.started
+ )
+ thl_ledger_manager.create_tx_task_complete(
+ wall=w2, user=user, created=w2.started
+ )
status, status_code_1 = s1.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
+ _, _, bp_pay, user_pay = s1.determine_payments()
session_manager.finish_with_status(
session=s1,
status=status,
@@ -1203,21 +1282,25 @@ class TestThlLedgerManagerAdj:
payout=bp_pay,
user_payout=user_pay,
)
- thl_lm.create_tx_bp_payment(session=s1, created=w1.started)
- revenue = thl_lm.get_account_task_complete_revenue()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- bp_commission_account = thl_lm.get_account_or_create_bp_commission(
+ thl_ledger_manager.create_tx_bp_payment(session=s1, created=w1.started)
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
product=user.product
)
- assert 380 == thl_lm.get_account_balance(account=bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(account=revenue)
- assert 20 == thl_lm.get_account_balance(account=bp_commission_account)
- thl_lm.check_ledger_balanced()
+ bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=user.product
+ )
+ assert 380 == thl_ledger_manager.get_account_balance(account=bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(account=revenue)
+ assert 20 == thl_ledger_manager.get_account_balance(
+ account=bp_commission_account
+ )
+ thl_ledger_manager.check_ledger_balanced()
# This should do nothing (since we haven't adjusted any wall events)
s1.adjust_status()
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
assert (
"create_transaction_bp_adjustment. No transactions needed." in caplog.text
@@ -1235,22 +1318,22 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=Decimal(0),
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall=w1, user=user)
+ thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user)
# -$1.00 b/c the MP took the $1 back, but we haven't yet taken the BP payment back
- assert -100 == thl_lm.get_account_balance(revenue)
+ assert -100 == thl_ledger_manager.get_account_balance(revenue)
s1.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
assert (
"create_transaction_bp_adjustment. No transactions needed." in caplog.text
)
- assert 380 - 95 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 - 5 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ assert 380 - 95 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 - 5 == thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
# unrecon the $1 survey
wall_manager.adjust_status(
@@ -1259,32 +1342,28 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=None,
adjusted_timestamp=utc_hour_ago + timedelta(minutes=45),
)
- thl_lm.create_tx_task_adjustment(
+ thl_ledger_manager.create_tx_task_adjustment(
wall=w1,
user=user,
created=utc_hour_ago + timedelta(minutes=45),
)
- new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts()
+ _, _, _ = s1.determine_new_status_and_payouts()
s1.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=s1)
- assert 380 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20, thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
+ assert 380 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20, thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
def test_create_tx_bp_adjustment_small(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
- delete_ledger_db,
- thl_lm,
- lm,
- utc_hour_ago,
- currency,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ utc_hour_ago: datetime,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
# This failed when I didn't check that `change_commission` > 0 in
# create_transaction_bp_adjustment
@@ -1302,51 +1381,46 @@ class TestThlLedgerManagerAdj:
finished=utc_hour_ago + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
session = Session(started=wall1.started, user=user, wall_events=[wall1])
status, status_code_1 = session.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
+ _, _, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": utc_hour_ago + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
- thl_lm.create_tx_bp_payment(session, created=wall1.started)
+ thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started)
wall1.update(
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
adjusted_cpi=0,
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
session.adjust_status()
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
def test_create_tx_bp_adjustment_abandon(
self,
- user_factory,
- product_user_wallet_no,
- delete_ledger_db,
- session_factory,
- create_main_accounts,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ session_factory: Callable[..., Session],
caplog,
- thl_lm,
- lm,
- currency,
- utc_hour_ago,
- session_manager,
- wall_manager,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ utc_hour_ago: datetime,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
):
- delete_ledger_db()
- create_main_accounts()
+
user: User = user_factory(product=product_user_wallet_no)
s1: Session = session_factory(
user=user, final_status=Status.ABANDON, wall_req_cpi=Decimal(1)
@@ -1360,9 +1434,9 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=w1.cpi,
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall=w1, user=user)
+ thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user)
s1.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
# And then adjust it back (it was abandon before, but now it should be
# fail (?) or back to abandon?)
wall_manager.adjust_status(
@@ -1371,24 +1445,26 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=None,
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall=w1, user=user)
+ thl_ledger_manager.create_tx_task_adjustment(wall=w1, user=user)
s1.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
- revenue = thl_lm.get_account_task_complete_revenue()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- bp_commission_account = thl_lm.get_account_or_create_bp_commission(
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
product=user.product
)
- assert 0 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 0 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=user.product
+ )
+ assert 0 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 0 == thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
# This should do nothing
s1.adjust_status()
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session=s1)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
assert "No transactions needed" in caplog.text
# Now back to complete again
@@ -1399,24 +1475,20 @@ class TestThlLedgerManagerAdj:
adjusted_timestamp=utc_hour_ago + timedelta(hours=1),
)
s1.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=s1)
- assert 95 == thl_lm.get_account_balance(bp_wallet_account)
+ thl_ledger_manager.create_tx_bp_adjustment(session=s1)
+ assert 95 == thl_ledger_manager.get_account_balance(bp_wallet_account)
def test_create_tx_bp_adjustment_user_wallet(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
caplog,
- thl_lm,
- lm,
- currency,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
- now = datetime.now(timezone.utc) - timedelta(days=1)
+ now = datetime.now(UTC) - timedelta(days=1)
user: User = user_factory(product=product_user_wallet_yes)
# Create 2 Wall completes and create the respective transaction for
@@ -1447,7 +1519,7 @@ class TestThlLedgerManagerAdj:
started=now_w1,
finished=now_w1 + timedelta(minutes=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
@@ -1464,7 +1536,7 @@ class TestThlLedgerManagerAdj:
started=now_w2,
finished=now_w2 + timedelta(minutes=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall2, user=user, created=wall2.started
)
assert isinstance(tx, LedgerTransaction)
@@ -1485,34 +1557,38 @@ class TestThlLedgerManagerAdj:
assert user_pay == Decimal("1.52")
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": now + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=now + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
- tx = thl_lm.create_tx_bp_adjustment(session=session, created=wall1.started)
+ tx = thl_ledger_manager.create_tx_bp_adjustment(
+ session=session, created=wall1.started
+ )
assert isinstance(tx, LedgerTransaction)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=user.product)
- assert 228 == thl_lm.get_account_balance(account=bp_wallet_account)
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=user.product
+ )
+ assert 228 == thl_ledger_manager.get_account_balance(account=bp_wallet_account)
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
- assert 152 == thl_lm.get_account_balance(account=user_account)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
+ assert 152 == thl_ledger_manager.get_account_balance(account=user_account)
- revenue = thl_lm.get_account_task_complete_revenue()
- assert 0 == thl_lm.get_account_balance(account=revenue)
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ assert 0 == thl_ledger_manager.get_account_balance(account=revenue)
- bp_commission_account = thl_lm.get_account_or_create_bp_commission(
+ bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
product=user.product
)
- assert 20 == thl_lm.get_account_balance(account=bp_commission_account)
+ assert 20 == thl_ledger_manager.get_account_balance(
+ account=bp_commission_account
+ )
# the total (4.00) = 2.28 + 1.52 + .20
- assert thl_lm.check_ledger_balanced()
+ assert thl_ledger_manager.check_ledger_balanced()
# This should do nothing (since we haven't adjusted any wall events)
session.adjust_status()
@@ -1522,7 +1598,7 @@ class TestThlLedgerManagerAdj:
session.get_user_payout_after_adjustment(),
)
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
assert (
"create_transaction_bp_adjustment. No transactions needed." in caplog.text
)
@@ -1533,16 +1609,16 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=0,
adjusted_timestamp=now + timedelta(hours=1),
)
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
# -$1.00 b/c the MP took the $1 back, but we haven't yet taken the BP payment back
- assert -100 == thl_lm.get_account_balance(revenue)
+ assert -100 == thl_ledger_manager.get_account_balance(revenue)
session.adjust_status()
print(
session.get_status_after_adjustment(),
session.get_payout_after_adjustment(),
session.get_user_payout_after_adjustment(),
)
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
# running this twice b/c it should do nothing the 2nd time
print(
@@ -1551,16 +1627,16 @@ class TestThlLedgerManagerAdj:
session.get_user_payout_after_adjustment(),
)
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
assert (
"create_transaction_bp_adjustment. No transactions needed." in caplog.text
)
- assert 228 - 57 == thl_lm.get_account_balance(bp_wallet_account)
- assert 152 - 38 == thl_lm.get_account_balance(user_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 - 5 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ assert 228 - 57 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 152 - 38 == thl_ledger_manager.get_account_balance(user_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 - 5 == thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
# unrecon the $1 survey
wall1.update(
@@ -1568,7 +1644,7 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=None,
adjusted_timestamp=now + timedelta(hours=2),
)
- tx = thl_lm.create_tx_task_adjustment(wall=wall1, user=user)
+ tx = thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user)
assert isinstance(tx, LedgerTransaction)
new_status, new_payout, new_user_payout = (
@@ -1581,13 +1657,17 @@ class TestThlLedgerManagerAdj:
session.get_payout_after_adjustment(),
session.get_user_payout_after_adjustment(),
)
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
- assert 228 - 57 + 57 == thl_lm.get_account_balance(bp_wallet_account)
- assert 152 - 38 + 38 == thl_lm.get_account_balance(user_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 - 5 + 5 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ assert 228 - 57 + 57 == thl_ledger_manager.get_account_balance(
+ bp_wallet_account
+ )
+ assert 152 - 38 + 38 == thl_ledger_manager.get_account_balance(user_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 - 5 + 5 == thl_ledger_manager.get_account_balance(
+ bp_commission_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
# make the $2 failure into a complete also
wall3.update(
@@ -1595,7 +1675,7 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=wall3.cpi,
adjusted_timestamp=now + timedelta(hours=2),
)
- thl_lm.create_tx_task_adjustment(wall3, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall3, user)
new_status, new_payout, new_user_payout = (
session.determine_new_status_and_payouts()
)
@@ -1606,27 +1686,30 @@ class TestThlLedgerManagerAdj:
session.get_payout_after_adjustment(),
session.get_user_payout_after_adjustment(),
)
- thl_lm.create_tx_bp_adjustment(session)
- assert 228 - 57 + 57 + 114 == thl_lm.get_account_balance(bp_wallet_account)
- assert 152 - 38 + 38 + 76 == thl_lm.get_account_balance(user_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 - 5 + 5 + 10 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ thl_ledger_manager.create_tx_bp_adjustment(session)
+ assert 228 - 57 + 57 + 114 == thl_ledger_manager.get_account_balance(
+ bp_wallet_account
+ )
+ assert 152 - 38 + 38 + 76 == thl_ledger_manager.get_account_balance(
+ user_account
+ )
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 - 5 + 5 + 10 == thl_ledger_manager.get_account_balance(
+ bp_commission_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
def test_create_transaction_bp_adjustment_cpi_adjustment(
self,
- user_factory,
- product_user_wallet_no,
- create_main_accounts,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
caplog,
- thl_lm,
- lm,
- utc_hour_ago,
- currency,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ utc_hour_ago: datetime,
+ currency: LedgerCurrency,
):
- delete_ledger_db()
- create_main_accounts()
+
user: User = user_factory(product=product_user_wallet_no)
wall1 = Wall(
@@ -1640,7 +1723,7 @@ class TestThlLedgerManagerAdj:
started=utc_hour_ago,
finished=utc_hour_ago + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
@@ -1656,32 +1739,34 @@ class TestThlLedgerManagerAdj:
started=utc_hour_ago,
finished=utc_hour_ago + timedelta(seconds=1),
)
- tx = thl_lm.create_tx_task_complete(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall2, user=user, created=wall2.started
)
assert isinstance(tx, LedgerTransaction)
session = Session(started=wall1.started, user=user, wall_events=[wall1, wall2])
status, status_code_1 = session.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = session.determine_payments()
+ _, _, bp_pay, user_pay = session.determine_payments()
session.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": utc_hour_ago + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
- )
- thl_lm.create_tx_bp_payment(session, created=wall1.started)
-
- revenue = thl_lm.get_account_task_complete_revenue()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_commission_account = thl_lm.get_account_or_create_bp_commission(user.product)
- assert 380 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
+ )
+ thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started)
+
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ user.product
+ )
+ assert 380 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 == thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
# cpi adjustment $1 -> $.60.
wall1.update(
@@ -1689,17 +1774,17 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=Decimal("0.60"),
adjusted_timestamp=utc_hour_ago + timedelta(minutes=30),
)
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
# -$0.40 b/c the MP took $0.40 back, but we haven't yet taken the BP payment back
- assert -40 == thl_lm.get_account_balance(revenue)
+ assert -40 == thl_ledger_manager.get_account_balance(revenue)
session.adjust_status()
print(
session.get_status_after_adjustment(),
session.get_payout_after_adjustment(),
session.get_user_payout_after_adjustment(),
)
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
# running this twice b/c it should do nothing the 2nd time
print(
@@ -1708,14 +1793,14 @@ class TestThlLedgerManagerAdj:
session.get_user_payout_after_adjustment(),
)
with caplog.at_level(logging.INFO):
- thl_lm.create_tx_bp_adjustment(session)
+ thl_ledger_manager.create_tx_bp_adjustment(session)
assert "create_transaction_bp_adjustment." in caplog.text
assert "No transactions needed." in caplog.text
- assert 380 - 38 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 20 - 2 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ assert 380 - 38 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 20 - 2 == thl_ledger_manager.get_account_balance(bp_commission_account)
+ assert thl_ledger_manager.check_ledger_balanced()
# adjust it to failure
wall1.update(
@@ -1723,13 +1808,17 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=0,
adjusted_timestamp=utc_hour_ago + timedelta(minutes=45),
)
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
session.adjust_status()
- thl_lm.create_tx_bp_adjustment(session)
- assert 300 - (300 * 0.05) == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 300 * 0.05 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ thl_ledger_manager.create_tx_bp_adjustment(session)
+ assert 300 - (300 * 0.05) == thl_ledger_manager.get_account_balance(
+ bp_wallet_account
+ )
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 300 * 0.05 == thl_ledger_manager.get_account_balance(
+ bp_commission_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
# and then back to cpi adj again, but this time for more than the orig amount
wall1.update(
@@ -1737,13 +1826,17 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=Decimal("2.00"),
adjusted_timestamp=utc_hour_ago + timedelta(minutes=45),
)
- thl_lm.create_tx_task_adjustment(wall1, user)
+ thl_ledger_manager.create_tx_task_adjustment(wall1, user)
session.adjust_status()
- thl_lm.create_tx_bp_adjustment(session)
- assert 500 - (500 * 0.05) == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(revenue)
- assert 500 * 0.05 == thl_lm.get_account_balance(bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ thl_ledger_manager.create_tx_bp_adjustment(session)
+ assert 500 - (500 * 0.05) == thl_ledger_manager.get_account_balance(
+ bp_wallet_account
+ )
+ assert 0 == thl_ledger_manager.get_account_balance(revenue)
+ assert 500 * 0.05 == thl_ledger_manager.get_account_balance(
+ bp_commission_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
# And adjust again
wall1.update(
@@ -1751,12 +1844,14 @@ class TestThlLedgerManagerAdj:
adjusted_cpi=Decimal("3.00"),
adjusted_timestamp=utc_hour_ago + timedelta(minutes=45),
)
- thl_lm.create_tx_task_adjustment(wall=wall1, user=user)
+ thl_ledger_manager.create_tx_task_adjustment(wall=wall1, user=user)
session.adjust_status()
- thl_lm.create_tx_bp_adjustment(session=session)
- assert 600 - (600 * 0.05) == thl_lm.get_account_balance(
+ thl_ledger_manager.create_tx_bp_adjustment(session=session)
+ assert 600 - (600 * 0.05) == thl_ledger_manager.get_account_balance(
account=bp_wallet_account
)
- assert 0 == thl_lm.get_account_balance(account=revenue)
- assert 600 * 0.05 == thl_lm.get_account_balance(account=bp_commission_account)
- assert thl_lm.check_ledger_balanced()
+ assert 0 == thl_ledger_manager.get_account_balance(account=revenue)
+ assert 600 * 0.05 == thl_ledger_manager.get_account_balance(
+ account=bp_commission_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py
index 1e7146a..3fd21dc 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py
@@ -1,30 +1,37 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerTransactionFlagAlreadyExistsError,
LedgerTransactionConditionFailedError,
+ LedgerTransactionFlagAlreadyExistsError,
)
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.payout import UserPayoutEvent
-from test_utils.managers.ledger.conftest import create_main_accounts
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
class TestLedgerManagerAMT:
def test_create_transaction_amt_ass_request(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
@@ -40,16 +47,16 @@ class TestLedgerManagerAMT:
)
flag_key = f"test:user_payout:{pe.uuid}:request"
flag_name = f"ledger-manager:transaction_flag:{flag_key}"
- lm.redis_client.delete(flag_name)
+ ledger_manager.redis_client.delete(flag_name)
# User has $0 in their wallet. They are allowed amt_assignment payouts until -$1.00
- thl_lm.create_tx_user_payout_request(user=user, payout_event=pe)
+ thl_ledger_manager.create_tx_user_payout_request(user=user, payout_event=pe)
with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError):
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user, payout_event=pe, skip_flag_check=False
)
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user, payout_event=pe, skip_flag_check=True
)
pe2 = UserPayoutEvent(
@@ -62,36 +69,40 @@ class TestLedgerManagerAMT:
flag_key = f"test:user_payout:{pe2.uuid}:request"
flag_name = f"ledger-manager:transaction_flag:{flag_key}"
- lm.redis_client.delete(flag_name)
+ ledger_manager.redis_client.delete(flag_name)
# 96 cents would put them over the -$1.00 limit
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- thl_lm.create_tx_user_payout_request(user, payout_event=pe2)
+ thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe2)
# But they could do 0.95 cents
pe2.amount = 95
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe2, skip_flag_check=True
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
product=user.product
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(
+ user=user
+ )
- assert 0 == lm.get_account_balance(account=bp_wallet_account)
- assert 0 == lm.get_account_balance(account=cash)
- assert 100 == lm.get_account_balance(account=bp_pending_account)
- assert -100 == lm.get_account_balance(account=user_wallet_account)
- assert thl_lm.check_ledger_balanced()
- assert -5 == thl_lm.get_account_filtered_balance(
+ assert 0 == ledger_manager.get_account_balance(account=bp_wallet_account)
+ assert 0 == ledger_manager.get_account_balance(account=cash)
+ assert 100 == ledger_manager.get_account_balance(account=bp_pending_account)
+ assert -100 == ledger_manager.get_account_balance(account=user_wallet_account)
+ assert thl_ledger_manager.check_ledger_balanced()
+ assert -5 == thl_ledger_manager.get_account_filtered_balance(
account=user_wallet_account,
metadata_key="payoutevent",
metadata_value=pe.uuid,
)
- assert -95 == thl_lm.get_account_filtered_balance(
+ assert -95 == thl_ledger_manager.get_account_filtered_balance(
account=user_wallet_account,
metadata_key="payoutevent",
metadata_value=pe2.uuid,
@@ -99,12 +110,12 @@ class TestLedgerManagerAMT:
def test_create_transaction_amt_ass_complete(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
@@ -118,40 +129,42 @@ class TestLedgerManagerAMT:
debit_account_uuid=uuid4().hex,
)
flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:request"
- lm.redis_client.delete(flag)
+ ledger_manager.redis_client.delete(flag)
flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:complete"
- lm.redis_client.delete(flag)
+ ledger_manager.redis_client.delete(flag)
# User has $0 in their wallet. They are allowed amt_assignment payouts until -$1.00
- thl_lm.create_tx_user_payout_request(user, payout_event=pe)
- thl_lm.create_tx_user_payout_complete(user, payout_event=pe)
+ thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe)
+ thl_ledger_manager.create_tx_user_payout_complete(user, payout_event=pe)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
user.product
)
- bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
+ user.product
+ )
+ bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
user.product, expense_name="amt"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user)
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user)
# BP wallet pays the 1cent fee
- assert -1 == thl_lm.get_account_balance(bp_wallet_account)
- assert -5 == thl_lm.get_account_balance(cash)
- assert -1 == thl_lm.get_account_balance(bp_amt_expense_account)
- assert 0 == thl_lm.get_account_balance(bp_pending_account)
- assert -5 == lm.get_account_balance(user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ assert -1 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert -5 == thl_ledger_manager.get_account_balance(cash)
+ assert -1 == thl_ledger_manager.get_account_balance(bp_amt_expense_account)
+ assert 0 == thl_ledger_manager.get_account_balance(bp_pending_account)
+ assert -5 == ledger_manager.get_account_balance(user_wallet_account)
+ assert thl_ledger_manager.check_ledger_balanced()
def test_create_transaction_amt_bonus(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
@@ -166,15 +179,15 @@ class TestLedgerManagerAMT:
debit_account_uuid=uuid4().hex,
)
flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:request"
- lm.redis_client.delete(flag)
+ ledger_manager.redis_client.delete(flag)
flag = f"ledger-manager:transaction_flag:test:user_payout:{pe.uuid}:complete"
- lm.redis_client.delete(flag)
+ ledger_manager.redis_client.delete(flag)
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
# User has $0 in their wallet. No amt bonus allowed
- thl_lm.create_tx_user_payout_request(user, payout_event=pe)
+ thl_ledger_manager.create_tx_user_payout_request(user, payout_event=pe)
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user,
amount=Decimal(5),
ref_uuid="e703830dec124f17abed2d697d8d7701",
@@ -182,68 +195,68 @@ class TestLedgerManagerAMT:
skip_flag_check=True,
)
pe.amount = 101
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe, skip_flag_check=False
)
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe, skip_flag_check=False
)
with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError):
# duplicate, even if amount changed
pe.amount = 200
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe, skip_flag_check=False
)
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
# duplicate
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe, skip_flag_check=True
)
pe.uuid = "533364150de4451198e5774e221a2acb"
pe.amount = 9900
with pytest.raises(expected_exception=ValueError):
# Trying to complete payout with no pending tx
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe, skip_flag_check=True
)
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
# trying to payout $99 with only a $5 balance
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe, skip_flag_check=True
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
user.product
)
- bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
user.product, expense_name="amt"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user)
- assert -500 + round(-101 * 0.20) == thl_lm.get_account_balance(
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user)
+ assert -500 + round(-101 * 0.20) == thl_ledger_manager.get_account_balance(
bp_wallet_account
)
- assert -101 == lm.get_account_balance(cash)
- assert -20 == lm.get_account_balance(bp_amt_expense_account)
- assert 0 == lm.get_account_balance(bp_pending_account)
- assert 500 - 101 == lm.get_account_balance(user_wallet_account)
- assert lm.check_ledger_balanced() is True
+ assert -101 == ledger_manager.get_account_balance(cash)
+ assert -20 == ledger_manager.get_account_balance(bp_amt_expense_account)
+ assert 0 == ledger_manager.get_account_balance(bp_pending_account)
+ assert 500 - 101 == ledger_manager.get_account_balance(user_wallet_account)
+ assert ledger_manager.check_ledger_balanced() is True
def test_create_transaction_amt_bonus_cancel(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
caplog,
- thl_lm,
- lm,
- delete_ledger_db,
+ thl_ledger_manager: ThlLedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
- now = datetime.now(timezone.utc) - timedelta(hours=1)
user: User = user_factory(product=product_amt_true)
pe = UserPayoutEvent(
@@ -254,41 +267,48 @@ class TestLedgerManagerAMT:
debit_account_uuid=uuid4().hex,
)
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user,
amount=Decimal(5),
ref_uuid="c44f4da2db1d421ebc6a5e5241ca4ce6",
description="Bribe",
skip_flag_check=True,
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe, skip_flag_check=True
)
- thl_lm.create_tx_user_payout_cancelled(
+ thl_ledger_manager.create_tx_user_payout_cancelled(
user, payout_event=pe, skip_flag_check=True
)
- with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- with caplog.at_level(logging.WARNING):
- thl_lm.create_tx_user_payout_complete(
- user, payout_event=pe, skip_flag_check=True
- )
+ with pytest.raises(
+ expected_exception=LedgerTransactionConditionFailedError
+ ), caplog.at_level(logging.WARNING):
+ thl_ledger_manager.create_tx_user_payout_complete(
+ user, payout_event=pe, skip_flag_check=True
+ )
assert "trying to complete payout that was already cancelled" in caplog.text
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
user.product
)
- bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
user.product, expense_name="amt"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user)
- assert -500 == thl_lm.get_account_balance(account=bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(account=cash)
- assert 0 == thl_lm.get_account_balance(account=bp_amt_expense_account)
- assert 0 == thl_lm.get_account_balance(account=bp_pending_account)
- assert 500 == thl_lm.get_account_balance(account=user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user)
+ assert -500 == thl_ledger_manager.get_account_balance(account=bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(account=cash)
+ assert 0 == thl_ledger_manager.get_account_balance(
+ account=bp_amt_expense_account
+ )
+ assert 0 == thl_ledger_manager.get_account_balance(account=bp_pending_account)
+ assert 500 == thl_ledger_manager.get_account_balance(
+ account=user_wallet_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
pe2 = UserPayoutEvent(
uuid=uuid4().hex,
@@ -297,17 +317,18 @@ class TestLedgerManagerAMT:
cashout_method_uuid=uuid4().hex,
debit_account_uuid=uuid4().hex,
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe2, skip_flag_check=True
)
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe2, skip_flag_check=True
)
- with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- with caplog.at_level(logging.WARNING):
- thl_lm.create_tx_user_payout_cancelled(
- user, payout_event=pe2, skip_flag_check=True
- )
+ with pytest.raises(
+ expected_exception=LedgerTransactionConditionFailedError
+ ), caplog.at_level(logging.WARNING):
+ thl_ledger_manager.create_tx_user_payout_cancelled(
+ user, payout_event=pe2, skip_flag_check=True
+ )
assert "trying to cancel payout that was already completed" in caplog.text
@@ -315,12 +336,12 @@ class TestLedgerManagerTango:
def test_create_transaction_tango_request(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
@@ -337,64 +358,65 @@ class TestLedgerManagerTango:
)
flag_key = f"test:user_payout:{pe.uuid}:request"
flag_name = f"ledger-manager:transaction_flag:{flag_key}"
- lm.redis_client.delete(flag_name)
- thl_lm.create_tx_user_bonus(
+ ledger_manager.redis_client.delete(flag_name)
+ thl_ledger_manager.create_tx_user_bonus(
user,
amount=Decimal(6),
ref_uuid="e703830dec124f17abed2d697d8d7701",
description="Bribe",
skip_flag_check=True,
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe, skip_flag_check=True
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
user.product
)
- bp_tango_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_tango_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
user.product, expense_name="tango"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user)
- assert -600 == thl_lm.get_account_balance(bp_wallet_account)
- assert 0 == thl_lm.get_account_balance(cash)
- assert 0 == thl_lm.get_account_balance(bp_tango_expense_account)
- assert 500 == thl_lm.get_account_balance(bp_pending_account)
- assert 600 - 500 == thl_lm.get_account_balance(user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(user)
+ assert -600 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert 0 == thl_ledger_manager.get_account_balance(cash)
+ assert 0 == thl_ledger_manager.get_account_balance(bp_tango_expense_account)
+ assert 500 == thl_ledger_manager.get_account_balance(bp_pending_account)
+ assert 600 - 500 == thl_ledger_manager.get_account_balance(user_wallet_account)
+ assert thl_ledger_manager.check_ledger_balanced()
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user, payout_event=pe, skip_flag_check=True
)
- assert -600 - round(500 * 0.035) == thl_lm.get_account_balance(
+ assert -600 - round(500 * 0.035) == thl_ledger_manager.get_account_balance(
bp_wallet_account
)
- assert -500, thl_lm.get_account_balance(cash)
- assert round(-500 * 0.035) == thl_lm.get_account_balance(
+ assert -500, thl_ledger_manager.get_account_balance(cash)
+ assert round(-500 * 0.035) == thl_ledger_manager.get_account_balance(
bp_tango_expense_account
)
- assert 0 == lm.get_account_balance(bp_pending_account)
- assert 100 == lm.get_account_balance(user_wallet_account)
- assert lm.check_ledger_balanced()
+ assert 0 == ledger_manager.get_account_balance(bp_pending_account)
+ assert 100 == ledger_manager.get_account_balance(user_wallet_account)
+ assert ledger_manager.check_ledger_balanced()
class TestLedgerManagerPaypal:
def test_create_transaction_paypal_request(
self,
- user_factory,
- product_amt_true,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
- now = datetime.now(tz=timezone.utc) - timedelta(hours=1)
user: User = user_factory(product=product_amt_true)
# debit_account_uuid nothing checks they match the ledger ... todo?
@@ -407,8 +429,8 @@ class TestLedgerManagerPaypal:
)
flag_key = f"test:user_payout:{pe.uuid}:request"
flag_name = f"ledger-manager:transaction_flag:{flag_key}"
- lm.redis_client.delete(flag_name)
- thl_lm.create_tx_user_bonus(
+ ledger_manager.redis_client.delete(flag_name)
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(6),
ref_uuid="e703830dec124f17abed2d697d8d7701",
@@ -416,79 +438,91 @@ class TestLedgerManagerPaypal:
skip_flag_check=True,
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user, payout_event=pe, skip_flag_check=True
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
product=user.product
)
- bp_paypal_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_paypal_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
product=user.product, expense_name="paypal"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user)
- assert -600 == lm.get_account_balance(account=bp_wallet_account)
- assert 0 == lm.get_account_balance(account=cash)
- assert 0 == lm.get_account_balance(account=bp_paypal_expense_account)
- assert 500 == lm.get_account_balance(account=bp_pending_account)
- assert 600 - 500 == lm.get_account_balance(account=user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(
+ user=user
+ )
+ assert -600 == ledger_manager.get_account_balance(account=bp_wallet_account)
+ assert 0 == ledger_manager.get_account_balance(account=cash)
+ assert 0 == ledger_manager.get_account_balance(
+ account=bp_paypal_expense_account
+ )
+ assert 500 == ledger_manager.get_account_balance(account=bp_pending_account)
+ assert 600 - 500 == ledger_manager.get_account_balance(
+ account=user_wallet_account
+ )
+ assert thl_ledger_manager.check_ledger_balanced()
- thl_lm.create_tx_user_payout_complete(
+ thl_ledger_manager.create_tx_user_payout_complete(
user=user, payout_event=pe, skip_flag_check=True, fee_amount=Decimal("0.50")
)
- assert -600 - 50 == thl_lm.get_account_balance(bp_wallet_account)
- assert -500 == thl_lm.get_account_balance(cash)
- assert -50 == thl_lm.get_account_balance(bp_paypal_expense_account)
- assert 0 == thl_lm.get_account_balance(bp_pending_account)
- assert 100 == thl_lm.get_account_balance(user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ assert -600 - 50 == thl_ledger_manager.get_account_balance(bp_wallet_account)
+ assert -500 == thl_ledger_manager.get_account_balance(cash)
+ assert -50 == thl_ledger_manager.get_account_balance(bp_paypal_expense_account)
+ assert 0 == thl_ledger_manager.get_account_balance(bp_pending_account)
+ assert 100 == thl_ledger_manager.get_account_balance(user_wallet_account)
+ assert thl_ledger_manager.check_ledger_balanced()
class TestLedgerManagerBonus:
def test_create_transaction_bonus(
self,
- user_factory,
- product_user_wallet_yes,
- create_main_accounts,
- thl_lm,
- lm,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ product_user_wallet_yes: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ delete_ledger_db: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
user: User = user_factory(product=product_user_wallet_yes)
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(5),
ref_uuid="8d0aaf612462448a9ebdd57fab0fc660",
description="Bribe",
skip_flag_check=True,
)
- cash = thl_lm.get_account_cash()
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(user.product)
- bp_pending_account = thl_lm.get_or_create_bp_pending_payout_account(
+ cash = thl_ledger_manager.get_account_cash()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ user.product
+ )
+ bp_pending_account = thl_ledger_manager.get_or_create_bp_pending_payout_account(
product=user.product
)
- bp_amt_expense_account = thl_lm.get_account_or_create_bp_expense(
+ bp_amt_expense_account = thl_ledger_manager.get_account_or_create_bp_expense(
user.product, expense_name="amt"
)
- user_wallet_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(
+ user=user
+ )
- assert -500 == lm.get_account_balance(account=bp_wallet_account)
- assert 0 == lm.get_account_balance(account=cash)
- assert 0 == lm.get_account_balance(account=bp_amt_expense_account)
- assert 0 == lm.get_account_balance(account=bp_pending_account)
- assert 500 == lm.get_account_balance(account=user_wallet_account)
- assert thl_lm.check_ledger_balanced()
+ assert -500 == ledger_manager.get_account_balance(account=bp_wallet_account)
+ assert 0 == ledger_manager.get_account_balance(account=cash)
+ assert 0 == ledger_manager.get_account_balance(account=bp_amt_expense_account)
+ assert 0 == ledger_manager.get_account_balance(account=bp_pending_account)
+ assert 500 == ledger_manager.get_account_balance(account=user_wallet_account)
+ assert thl_ledger_manager.check_ledger_balanced()
with pytest.raises(expected_exception=LedgerTransactionFlagAlreadyExistsError):
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(5),
ref_uuid="8d0aaf612462448a9ebdd57fab0fc660",
@@ -496,7 +530,7 @@ class TestLedgerManagerBonus:
skip_flag_check=False,
)
with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- thl_lm.create_tx_user_bonus(
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(5),
ref_uuid="8d0aaf612462448a9ebdd57fab0fc660",
diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py
index 5fb9e7d..fb35aa4 100644
--- a/tests/managers/thl/test_ledger/test_thl_pem.py
+++ b/tests/managers/thl/test_ledger/test_thl_pem.py
@@ -1,21 +1,39 @@
-import uuid
+from __future__ import annotations
+
+from collections.abc import Callable
from random import randint
-from uuid import uuid4, UUID
+from typing import TYPE_CHECKING
+from uuid import UUID, uuid4
import pytest
from generalresearch.currency import USDCent
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.payout import BrokerageProductPayoutEvent
-from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+)
from generalresearch.models.thl.wallet.cashout_method import (
CashoutRequestInfo,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.payout import (
+ BrokerageProductPayoutEventManager,
+ UserPayoutEventManager,
+ )
+ from generalresearch.models.thl.payout import UserPayoutEvent
+ from generalresearch.models.thl.product import Product
+
class TestThlPayoutEventManager:
- def test_get_by_uuid(self, brokerage_product_payout_event_manager, thl_lm):
+ def test_get_by_uuid(
+ self, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager
+ ):
"""This validates that the method raises an exception if it
fails. There are plenty of other tests that use this method so
it seems silly to duplicate it here again
@@ -27,35 +45,31 @@ class TestThlPayoutEventManager:
def test_filter_by(
self,
- product_factory,
- usd_cent,
- bp_payout_event_factory,
- thl_lm,
- brokerage_product_payout_event_manager,
+ product_factory: Callable[..., Product],
+ usd_cent: USDCent,
+ bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent],
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
N_PRODUCTS = randint(3, 10)
N_PAYOUT_EVENTS = randint(3, 10)
amounts = []
products = []
- for x_idx in range(N_PRODUCTS):
+ for _ in range(N_PRODUCTS):
product: Product = product_factory()
- thl_lm.get_account_or_create_bp_wallet(product=product)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
products.append(product)
- brokerage_product_payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm
- )
- for y_idx in range(N_PAYOUT_EVENTS):
+ for _ in range(N_PAYOUT_EVENTS):
pe = bp_payout_event_factory(product=product, usd_cent=usd_cent)
amounts.append(int(usd_cent))
assert isinstance(pe, BrokerageProductPayoutEvent)
# We just added Payout Events for Products, now go ahead and
# query for them
- accounts = thl_lm.get_accounts_bp_wallet_for_products(
+ accounts = thl_ledger_manager.get_accounts_bp_wallet_for_products(
product_uuids=[i.uuid for i in products]
)
res = brokerage_product_payout_event_manager.filter_by(
@@ -67,36 +81,32 @@ class TestThlPayoutEventManager:
def test_get_bp_payout_events_for_product(
self,
- product_factory,
- usd_cent,
- bp_payout_event_factory,
- brokerage_product_payout_event_manager,
- thl_lm,
+ product_factory: Callable[..., Product],
+ usd_cent: USDCent,
+ bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent],
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
N_PRODUCTS = randint(3, 10)
N_PAYOUT_EVENTS = randint(3, 10)
amounts = []
products = []
- for x_idx in range(N_PRODUCTS):
+ for _ in range(N_PRODUCTS):
product: Product = product_factory()
products.append(product)
- thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm
- )
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
- for y_idx in range(N_PAYOUT_EVENTS):
+ for _ in range(N_PAYOUT_EVENTS):
pe = bp_payout_event_factory(product=product, usd_cent=usd_cent)
amounts.append(usd_cent)
assert isinstance(pe, BrokerageProductPayoutEvent)
- # We just added 5 Payouts for a specific Product, now go
+ # We just added 5 Payouts for a specific product: Product, now go
# ahead and query for them
res = brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm, product_uuids=[product.id]
+ product_uuids=[product.id]
)
assert len(res) == N_PAYOUT_EVENTS
@@ -105,7 +115,7 @@ class TestThlPayoutEventManager:
# ahead and query for them
res = (
brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm, product_uuids=[i.uuid for i in products]
+ product_uuids=[i.uuid for i in products],
)
)
@@ -113,13 +123,12 @@ class TestThlPayoutEventManager:
assert sum([i.amount for i in res]) == sum(amounts)
@pytest.mark.skip
- def test_get_payout_detail(self, user_payout_event_manager):
+ def test_get_payout_detail(self, user_payout_event_manager: UserPayoutEventManager):
"""This fails because the description coming back is None, but then
it tries to return a PayoutEvent which validates that the
description can't be None
"""
from generalresearch.models.thl.payout import (
- UserPayoutEvent,
PayoutType,
)
@@ -145,11 +154,15 @@ class TestThlPayoutEventManager:
# def test_filter_by(self):
# raise NotImplementedError
- def test_create(self, user_payout_event_manager):
+ def test_create(
+ self,
+ user_payout_event_factory: Callable[..., UserPayoutEvent],
+ user_payout_event_manager: UserPayoutEventManager,
+ ):
from generalresearch.models.thl.payout import UserPayoutEvent
# Confirm the creation method returns back an instance.
- pe = user_payout_event_manager.create_dummy()
+ pe = user_payout_event_factory()
assert isinstance(pe, UserPayoutEvent)
# Now query the DB for that PayoutEvent to confirm it was actually
@@ -167,27 +180,26 @@ class TestThlPayoutEventManager:
def test_create_bp_payout(
self,
- product,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- brokerage_product_payout_event_manager,
- lm,
+ product: Product,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ ledger_manager: LedgerManager,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
delete_ledger_db()
create_main_accounts()
- account_bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
+ account_bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
rand_amount = randint(a=99, b=999)
# Save a Brokerage Product Payout, so we have something in the
# Payout Event table and the respective ledger TX and Entry rows for it
pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ thl_ledger_manager=thl_ledger_manager,
product=product,
amount=USDCent(rand_amount),
skip_wallet_balance_check=True,
@@ -196,15 +208,17 @@ class TestThlPayoutEventManager:
assert isinstance(pe, BrokerageProductPayoutEvent)
# Now try to query for it!
- res = thl_lm.get_tx_bp_payouts(account_uuids=[account_bp_wallet.uuid])
+ res = thl_ledger_manager.get_tx_bp_payouts(
+ account_uuids=[account_bp_wallet.uuid]
+ )
assert len(res) == 1
- res = thl_lm.get_tx_bp_payouts(account_uuids=[uuid4().hex])
+ res = thl_ledger_manager.get_tx_bp_payouts(account_uuids=[uuid4().hex])
assert len(res) == 0
# Confirm it added to the users balance. The amount is negative because
- # money was sent to the Brokerage Product, but they didn't have
+ # money was sent to the Brokerage product: Product, but they didn't have
# any activity that earned them money
- bal = lm.get_account_balance(account=account_bp_wallet)
+ bal = ledger_manager.get_account_balance(account=account_bp_wallet)
assert rand_amount == bal * -1
@@ -212,13 +226,13 @@ class TestBPPayoutEvent:
def test_get_bp_bp_payout_events_for_products(
self,
- product_factory,
- bp_payout_event_factory,
- usd_cent,
- delete_ledger_db,
- create_main_accounts,
- brokerage_product_payout_event_manager,
- thl_lm,
+ product_factory: Callable[..., Product],
+ bp_payout_event_factory: Callable[..., BrokerageProductPayoutEvent],
+ usd_cent: USDCent,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
):
delete_ledger_db()
create_main_accounts()
@@ -227,10 +241,9 @@ class TestBPPayoutEvent:
amounts = []
product: Product = product_factory()
- thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
- for y_idx in range(N_PAYOUT_EVENTS):
+ for _ in range(N_PAYOUT_EVENTS):
bp_payout_event_factory(product=product, usd_cent=usd_cent)
amounts.append(usd_cent)
@@ -238,7 +251,7 @@ class TestBPPayoutEvent:
# array of BPPayoutEvents
bp_bp_res = (
brokerage_product_payout_event_manager.get_bp_bp_payout_events_for_products(
- thl_ledger_manager=thl_lm, product_uuids=[product.uuid]
+ product_uuids=[product.uuid]
)
)
assert isinstance(bp_bp_res, list)
diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py
index ecf146f..6b6ef5b 100644
--- a/tests/managers/thl/test_ledger/test_user_txs.py
+++ b/tests/managers/thl/test_ledger/test_user_txs.py
@@ -1,54 +1,58 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
-from typing import TYPE_CHECKING, Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
-from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.managers.thl.user_compensate import user_compensate
from generalresearch.models.thl.definitions import (
Status,
- WallAdjustedStatus,
)
from generalresearch.models.thl.ledger import (
TransactionType,
UserLedgerTransactionTypesSummary,
UserLedgerTransactionTypeSummary,
)
+from generalresearch.models.thl.wallet.definitions import PayoutType
if TYPE_CHECKING:
- from generalresearch.config import GRLSettings
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.payout import UserPayoutEventManager
from generalresearch.models.thl.product import Product
from generalresearch.models.thl.session import Session
from generalresearch.models.thl.user import User
- from generalresearch.models.thl.wallet import PayoutType
def test_user_txs(
- user_factory: Callable[..., "User"],
- product_amt_true: "Product",
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
create_main_accounts: Callable[..., None],
- thl_lm: ThlLedgerManager,
- lm,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
delete_ledger_db: Callable[..., None],
- session_with_tx_factory,
- adj_to_fail_with_tx_factory,
- adj_to_complete_with_tx_factory,
- session_factory,
- user_payout_event_manager,
+ session_with_tx_factory: Callable[..., Session],
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ adj_to_complete_with_tx_factory: Callable[..., None],
+ session_factory: Callable[..., Session],
+ user_payout_event_manager: UserPayoutEventManager,
utc_now: datetime,
- settings: "GRLSettings",
+ settings: GRLBaseSettings,
):
delete_ledger_db()
create_main_accounts()
user: User = user_factory(product=product_amt_true)
- account = thl_lm.get_account_or_create_user_wallet(user)
+ account = thl_ledger_manager.get_account_or_create_user_wallet(user)
print(f"{account.uuid=}")
s: Session = session_with_tx_factory(user=user, wall_req_cpi=Decimal("1.00"))
- bribe_uuid = user_compensate(
- ledger_manager=thl_lm,
+ user_compensate(
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
)
@@ -60,9 +64,9 @@ def test_user_txs(
amount=5,
created=utc_now,
payout_type=PayoutType.AMT_HIT,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
@@ -73,9 +77,9 @@ def test_user_txs(
amount=127,
created=utc_now,
payout_type=PayoutType.AMT_BONUS,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
@@ -92,16 +96,16 @@ def test_user_txs(
)
adj_to_complete_with_tx_factory(session=s_fail, created=utc_now)
- # txs = thl_lm.get_tx_filtered_by_account(account.uuid)
+ # txs = thl_ledger_manager.get_tx_filtered_by_account(account.uuid)
# print(len(txs), txs)
- txs = thl_lm.get_user_txs(user)
+ txs = thl_ledger_manager.get_user_txs(user)
assert len(txs.transactions) == 6
assert txs.total == 6
assert txs.page == 1
assert txs.size == 50
# print(len(txs.transactions), txs)
- d = txs.model_dump_json()
+ # d = txs.model_dump_json()
# print(d)
descriptions = {x.description for x in txs.transactions}
@@ -136,33 +140,29 @@ def test_user_txs(
def test_user_txs_pagination(
- user_factory: Callable[..., "User"],
- product_amt_true: "Product",
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
create_main_accounts: Callable[..., None],
- thl_lm: "ThlLedgerManager",
- lm: "LedgerManager",
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
delete_ledger_db: Callable[..., None],
- session_with_tx_factory: Callable[..., "Session"],
- adj_to_fail_with_tx_factory,
- user_payout_event_manager,
- utc_now: datetime,
):
delete_ledger_db()
create_main_accounts()
user: User = user_factory(product=product_amt_true)
- account = thl_lm.get_account_or_create_user_wallet(user)
+ account = thl_ledger_manager.get_account_or_create_user_wallet(user)
print(f"{account.uuid=}")
for _ in range(12):
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
)
- txs = thl_lm.get_user_txs(user, page=1, size=5)
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=5)
assert len(txs.transactions) == 5
assert txs.total == 12
assert txs.page == 1
@@ -171,7 +171,7 @@ def test_user_txs_pagination(
assert txs.summary.user_bonus.entry_count == 12
# Skip to the 3rd page. We made 12, so there are 2 left
- txs = thl_lm.get_user_txs(user, page=3, size=5)
+ txs = thl_ledger_manager.get_user_txs(user, page=3, size=5)
assert len(txs.transactions) == 2
assert txs.total == 12
assert txs.page == 3
@@ -179,7 +179,7 @@ def test_user_txs_pagination(
assert txs.summary.user_bonus.entry_count == 12
# Should be empty, not fail
- txs = thl_lm.get_user_txs(user, page=4, size=5)
+ txs = thl_ledger_manager.get_user_txs(user, page=4, size=5)
assert len(txs.transactions) == 0
assert txs.total == 12
assert txs.page == 4
@@ -187,14 +187,14 @@ def test_user_txs_pagination(
assert txs.summary.user_bonus.entry_count == 12
# Test filtering. We should pull back only this one
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
)
- txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now)
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=5, time_start=now)
assert len(txs.transactions) == 1
assert txs.total == 1
assert txs.page == 1
@@ -203,8 +203,8 @@ def test_user_txs_pagination(
assert txs.summary.user_bonus.entry_count == 1
# And filtering with 0 results
- now = datetime.now(tz=timezone.utc)
- txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now)
+ now = datetime.now(tz=UTC)
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=5, time_start=now)
assert len(txs.transactions) == 0
assert txs.total == 0
assert txs.page == 1
@@ -215,16 +215,13 @@ def test_user_txs_pagination(
def test_user_txs_rolling_balance(
- user_factory: Callable[..., "User"],
- product_amt_true: "Product",
- create_main_accounts,
- thl_lm,
- lm,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
delete_ledger_db: Callable[..., None],
- session_with_tx_factory,
- adj_to_fail_with_tx_factory,
- user_payout_event_manager,
- settings: "GRLSettings",
+ user_payout_event_manager: UserPayoutEventManager,
+ settings: GRLBaseSettings,
):
"""
Creates 3 $1.00 bonuses (postive),
@@ -237,11 +234,11 @@ def test_user_txs_rolling_balance(
create_main_accounts()
user: User = user_factory(product=product_amt_true)
- account = thl_lm.get_account_or_create_user_wallet(user)
+ account = thl_ledger_manager.get_account_or_create_user_wallet(user)
for _ in range(3):
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
@@ -253,21 +250,21 @@ def test_user_txs_rolling_balance(
cashout_method_uuid=settings.amt_bonus_cashout_method_id,
amount=150,
payout_type=PayoutType.AMT_BONUS,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
for _ in range(3):
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
)
- txs = thl_lm.get_user_txs(user, page=1, size=10)
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=10)
assert txs.transactions[0].balance_after == 100
assert txs.transactions[1].balance_after == 200
assert txs.transactions[2].balance_after == 300
@@ -278,7 +275,7 @@ def test_user_txs_rolling_balance(
# Ascending order, get 2nd page, make sure the balances include
# the previous txs. (will return last 3 txs)
- txs = thl_lm.get_user_txs(user, page=2, size=4)
+ txs = thl_ledger_manager.get_user_txs(user, page=2, size=4)
assert len(txs.transactions) == 3
assert txs.transactions[0].balance_after == 250
assert txs.transactions[1].balance_after == 350
@@ -286,7 +283,7 @@ def test_user_txs_rolling_balance(
# Descending order, get 1st page. Will
# return most recent 3 txs in desc order
- txs = thl_lm.get_user_txs(user, page=1, size=3, order_by="-created")
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=3, order_by="-created")
assert len(txs.transactions) == 3
assert txs.transactions[0].balance_after == 450
assert txs.transactions[1].balance_after == 350
diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py
index a0abd7c..0a1da73 100644
--- a/tests/managers/thl/test_ledger/test_wallet.py
+++ b/tests/managers/thl/test_ledger/test_wallet.py
@@ -1,20 +1,31 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.models.thl.product import (
- UserWalletConfig,
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ Product,
+ UserWalletConfig,
)
-from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.thl.user import User
@pytest.fixture()
-def schrute_product(product_manager):
- return product_manager.create_dummy(
+def schrute_product(
+ product_factory: Callable[..., Product], product_manager: ProductManager
+) -> Product:
+ return product_factory(
user_wallet_config=UserWalletConfig(enabled=True, amt=False),
payout_config=PayoutConfig(
payout_transformation=PayoutTransformation(
@@ -27,25 +38,31 @@ def schrute_product(product_manager):
class TestGetUserWalletBalance:
- def test_get_user_wallet_balance_non_managed(self, user, thl_lm):
+ def test_get_user_wallet_balance_non_managed(
+ self, user: User, thl_ledger_manager: ThlLedgerManager
+ ):
with pytest.raises(
AssertionError,
match="Can't get wallet balance on non-managed account.",
):
- thl_lm.get_user_wallet_balance(user=user)
+ thl_ledger_manager.get_user_wallet_balance(user=user)
def test_get_user_wallet_balance_managed_0(
- self, schrute_product, user_factory, thl_lm
+ self,
+ schrute_product: Product,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
):
assert (
schrute_product.payout_config.payout_format == "{payout:,.0f} Schrute Bucks"
)
- user: User = user_factory(schrute_product)
- balance = thl_lm.get_user_wallet_balance(user=user)
+ user: User = user_factory(product=schrute_product)
+ balance = thl_ledger_manager.get_user_wallet_balance(user=user)
assert balance == 0
+ assert isinstance(user.product, Product)
balance_string = user.product.format_payout_format(Decimal(balance) / 100)
assert balance_string == "0 Schrute Bucks"
- redeemable_balance = thl_lm.get_user_redeemable_wallet_balance(
+ redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance(
user=user, user_wallet_balance=balance
)
assert redeemable_balance == 0
@@ -55,10 +72,14 @@ class TestGetUserWalletBalance:
assert redeemable_balance_string == "0 Schrute Bucks"
def test_get_user_wallet_balance_managed(
- self, schrute_product, user_factory, thl_lm, session_with_tx_factory
+ self,
+ schrute_product: Product,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
+ session_with_tx_factory: Callable[..., None],
):
- user: User = user_factory(schrute_product)
- thl_lm.create_tx_user_bonus(
+ user: User = user_factory(product=schrute_product)
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(1),
ref_uuid=uuid4().hex,
@@ -69,10 +90,10 @@ class TestGetUserWalletBalance:
# This product has a payout xform of 40% and commission of 5%
# 1.23 * 0.05 = 0.06 of commission
# 1.17 of payout * 0.40 = 0.47 of user pay and (1.17-0.47) 0.70 bp pay
- balance = thl_lm.get_user_wallet_balance(user=user)
+ balance = thl_ledger_manager.get_user_wallet_balance(user=user)
assert balance == 47 + 100 # plus the $1 bribe
- redeemable_balance = thl_lm.get_user_redeemable_wallet_balance(
+ redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance(
user=user, user_wallet_balance=balance
)
assert redeemable_balance == 20 + 100
diff --git a/tests/managers/thl/test_maxmind.py b/tests/managers/thl/test_maxmind.py
index c588c58..e44fe49 100644
--- a/tests/managers/thl/test_maxmind.py
+++ b/tests/managers/thl/test_maxmind.py
@@ -1,23 +1,6 @@
-import json
-import logging
-from typing import Callable
-
-import geoip2.models
-import pytest
from faker import Faker
from faker.providers.address.en_US import Provider as USAddressProvider
-from generalresearch.managers.thl.ipinfo import GeoIpInfoManager
-from generalresearch.managers.thl.maxmind import MaxmindManager
-from generalresearch.managers.thl.maxmind.basic import (
- MaxmindBasicManager,
-)
-from generalresearch.models.thl.ipinfo import (
- GeoIPInformation,
- normalize_ip,
-)
-from generalresearch.models.thl.maxmind.definitions import UserType
-
fake = Faker()
US_STATES = {x.lower() for x in USAddressProvider.states}
@@ -29,245 +12,244 @@ IP_v6_US = "2600:1700:ece0:9410:55d:faf3:c15d:6e4"
IP_v6_US_SAME_64 = "2600:1700:ece0:9410:55d:faf3:c15d:aaaa"
-@pytest.fixture(scope="session")
-def delete_ipinfo(thl_web_rw) -> Callable:
- def _delete_ipinfo(ip):
- thl_web_rw.execute_write(
- query="DELETE FROM thl_geoname WHERE geoname_id IN (SELECT geoname_id FROM thl_ipinformation WHERE ip = %s);",
- params=[ip],
- )
- thl_web_rw.execute_write(
- query="DELETE FROM thl_ipinformation WHERE ip = %s;",
- params=[ip],
- )
-
- return _delete_ipinfo
-
-
-class TestMaxmindBasicManager:
-
- def test_init(self, maxmind_basic_manager):
-
- assert isinstance(maxmind_basic_manager, MaxmindBasicManager)
-
- def test_get_basic_ip_information(self, maxmind_basic_manager):
- ip = IP_v4_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
- assert isinstance(res1, geoip2.models.Country)
- assert res1.country.iso_code == "IN"
- assert res1.country.name == "India"
-
- res2 = maxmind_basic_manager.get_basic_ip_information(
- ip_address=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_country_iso_from_ip_geoip2db(self, maxmind_basic_manager):
- ip = IP_v4_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(ip=ip)
- assert res1 == "in"
-
- res2 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(
- ip=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_basic_ip_information_ipv6(self, maxmind_basic_manager):
- ip = IP_v6_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
- assert isinstance(res1, geoip2.models.Country)
- assert res1.country.iso_code == "IN"
- assert res1.country.name == "India"
-
-
-class TestMaxmindManager:
-
- def test_init(self, thl_web_rr, thl_redis_config, maxmind_manager: MaxmindManager):
- instance = MaxmindManager(pg_config=thl_web_rr, redis_config=thl_redis_config)
- assert isinstance(instance, MaxmindManager)
- assert isinstance(maxmind_manager, MaxmindManager)
-
- def test_create_basic(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in India, and so it should only do the basic lookup
- ip = IP_v4_INDIA
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
-
- maxmind_manager.run_ip_information(ip, force_insights=False)
- # Check that it is in the cache and in mysql
- res = geoipinfo_manager.get_cache(ip)
- assert res.ip == ip
- assert res.basic
- res = geoipinfo_manager.get_mysql(ip)
- assert res.ip == ip
- assert res.basic
-
- def test_create_basic_ipv6(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in India, and so it should only do the basic lookup
- ip = IP_v6_INDIA
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_cache(normalized_ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
-
- maxmind_manager.run_ip_information(ip, force_insights=False)
-
- # Check that it is in the cache
- res = geoipinfo_manager.get_cache(ip)
- # The looked up IP (/128) is returned,
- assert res.ip == ip
- assert res.lookup_prefix == "/64"
- assert res.basic
-
- # ... but the normalized version was stored (/64)
- assert geoipinfo_manager.get_cache_raw(ip) is None
- res = json.loads(geoipinfo_manager.get_cache_raw(normalized_ip))
- assert res["ip"] == normalized_ip
-
- # Check mysql
- res = geoipinfo_manager.get_mysql(ip)
- assert res.ip == ip
- assert res.lookup_prefix == "/64"
- assert res.basic
- with pytest.raises(AssertionError):
- geoipinfo_manager.get_mysql_raw(ip)
- res = geoipinfo_manager.get_mysql_raw(normalized_ip)
- assert res["ip"] == normalized_ip
-
- def test_create_insights(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in the US, so it should do insights
- ip = IP_v4_US
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
-
- res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
- assert isinstance(res1, GeoIPInformation)
-
- # Check that it is in the cache and in mysql
- res2 = geoipinfo_manager.get_cache(ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert not res2.basic
-
- res3 = geoipinfo_manager.get_mysql(ip)
- assert isinstance(res3, GeoIPInformation)
- assert res3.ip == ip
- assert not res3.basic
- assert res3.is_anonymous is False
- assert res3.subdivision_1_name.lower() in US_STATES
- # this might change ...
- assert res3.user_type == UserType.CELLULAR
-
- assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
-
- def test_create_insights_ipv6(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in the US, so it should do insights
- ip = IP_v6_US
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_cache(normalized_ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
-
- res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
- assert isinstance(res1, GeoIPInformation)
- assert res1.lookup_prefix == "/64"
-
- # Check that it is in the cache and in mysql
- res2 = geoipinfo_manager.get_cache(ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert not res2.basic
-
- res3 = geoipinfo_manager.get_mysql(ip)
- assert isinstance(res3, GeoIPInformation)
- assert res3.ip == ip
- assert not res3.basic
- assert res3.is_anonymous is False
- assert res3.subdivision_1_name.lower() in US_STATES
- # this might change ...
- assert res3.user_type == UserType.RESIDENTIAL
-
- assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
-
- def test_get_or_create_ip_information(self, maxmind_manager):
- ip = IP_v4_US
-
- res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res1, GeoIPInformation)
-
- res2 = maxmind_manager.get_or_create_ip_information(
- ip_address=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_or_create_ip_information_ipv6(
- self, maxmind_manager, delete_ipinfo, geoipinfo_manager, caplog
- ):
- ip = IP_v6_US
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
-
- with caplog.at_level(logging.INFO):
- res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res1, GeoIPInformation)
- assert res1.ip == ip
- # It looks up in insight using the normalize IP!
- assert f"get_insights_ip_information: {normalized_ip}" in caplog.text
-
- # And it should NOT do the lookup again with an ipv6 in the same /64 block!
- ip = IP_v6_US_SAME_64
- caplog.clear()
- with caplog.at_level(logging.INFO):
- res2 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert "get_insights_ip_information" not in caplog.text
-
- def test_run_ip_information(self, maxmind_manager):
- ip = IP_v4_US
-
- res = maxmind_manager.run_ip_information(ip_address=ip)
- assert isinstance(res, GeoIPInformation)
- assert res.country_name == "United States"
- assert res.country_iso == "us"
+# @pytest.fixture(scope="session")
+# def delete_ipinfo(thl_web_rw) -> Callable:
+# def _delete_ipinfo(ip):
+# thl_web_rw.execute_write(
+# query="DELETE FROM thl_geoname WHERE geoname_id IN (SELECT geoname_id FROM thl_ipinformation WHERE ip = %s);",
+# params=[ip],
+# )
+# thl_web_rw.execute_write(
+# query="DELETE FROM thl_ipinformation WHERE ip = %s;",
+# params=[ip],
+# )
+
+# return _delete_ipinfo
+
+
+# @pytest.skip("TODO: Replace with GRIP Client")
+# class TestMaxmindBasicManager:
+
+# def test_init(self,):
+
+# def test_get_basic_ip_information(self, maxmind_basic_manager):
+# ip = IP_v4_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
+# # assert isinstance(res1, geoip2.models.Country)
+# assert res1.country.iso_code == "IN"
+# assert res1.country.name == "India"
+
+# res2 = maxmind_basic_manager.get_basic_ip_information(
+# ip_address=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_country_iso_from_ip_geoip2db(self, maxmind_basic_manager):
+# ip = IP_v4_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(ip=ip)
+# assert res1 == "in"
+
+# res2 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(
+# ip=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_basic_ip_information_ipv6(self, maxmind_basic_manager):
+# ip = IP_v6_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
+# assert isinstance(res1, geoip2.models.Country)
+# assert res1.country.iso_code == "IN"
+# assert res1.country.name == "India"
+
+
+# class TestMaxmindManager:
+
+# def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, maxmind_manager: MaxmindManager):
+# instance = MaxmindManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config)
+# assert isinstance(instance, MaxmindManager)
+# assert isinstance(maxmind_manager, MaxmindManager)
+
+# def test_create_basic(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in India, and so it should only do the basic lookup
+# ip = IP_v4_INDIA
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+
+# maxmind_manager.run_ip_information(ip, force_insights=False)
+# # Check that it is in the cache and in mysql
+# res = geoipinfo_manager.get_cache(ip)
+# assert res.ip == ip
+# assert res.basic
+# res = geoipinfo_manager.get_mysql(ip)
+# assert res.ip == ip
+# assert res.basic
+
+# def test_create_basic_ipv6(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in India, and so it should only do the basic lookup
+# ip = IP_v6_INDIA
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_cache(normalized_ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
+
+# maxmind_manager.run_ip_information(ip, force_insights=False)
+
+# # Check that it is in the cache
+# res = geoipinfo_manager.get_cache(ip)
+# # The looked up IP (/128) is returned,
+# assert res.ip == ip
+# assert res.lookup_prefix == "/64"
+# assert res.basic
+
+# # ... but the normalized version was stored (/64)
+# assert geoipinfo_manager.get_cache_raw(ip) is None
+# res = json.loads(geoipinfo_manager.get_cache_raw(normalized_ip))
+# assert res["ip"] == normalized_ip
+
+# # Check mysql
+# res = geoipinfo_manager.get_mysql(ip)
+# assert res.ip == ip
+# assert res.lookup_prefix == "/64"
+# assert res.basic
+# with pytest.raises(AssertionError):
+# geoipinfo_manager.get_mysql_raw(ip)
+# res = geoipinfo_manager.get_mysql_raw(normalized_ip)
+# assert res["ip"] == normalized_ip
+
+# def test_create_insights(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in the US, so it should do insights
+# ip = IP_v4_US
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+
+# res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
+# assert isinstance(res1, GeoIPInformation)
+
+# # Check that it is in the cache and in mysql
+# res2 = geoipinfo_manager.get_cache(ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert not res2.basic
+
+# res3 = geoipinfo_manager.get_mysql(ip)
+# assert isinstance(res3, GeoIPInformation)
+# assert res3.ip == ip
+# assert not res3.basic
+# assert res3.is_anonymous is False
+# assert res3.subdivision_1_name.lower() in US_STATES
+# # this might change ...
+# assert res3.user_type == UserType.CELLULAR
+
+# assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
+
+# def test_create_insights_ipv6(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in the US, so it should do insights
+# ip = IP_v6_US
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_cache(normalized_ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
+
+# res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
+# assert isinstance(res1, GeoIPInformation)
+# assert res1.lookup_prefix == "/64"
+
+# # Check that it is in the cache and in mysql
+# res2 = geoipinfo_manager.get_cache(ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert not res2.basic
+
+# res3 = geoipinfo_manager.get_mysql(ip)
+# assert isinstance(res3, GeoIPInformation)
+# assert res3.ip == ip
+# assert not res3.basic
+# assert res3.is_anonymous is False
+# assert res3.subdivision_1_name.lower() in US_STATES
+# # this might change ...
+# assert res3.user_type == UserType.RESIDENTIAL
+
+# assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
+
+# def test_get_or_create_ip_information(self, maxmind_manager):
+# ip = IP_v4_US
+
+# res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res1, GeoIPInformation)
+
+# res2 = maxmind_manager.get_or_create_ip_information(
+# ip_address=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_or_create_ip_information_ipv6(
+# self, maxmind_manager, delete_ipinfo, geoipinfo_manager, caplog
+# ):
+# ip = IP_v6_US
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+
+# with caplog.at_level(logging.INFO):
+# res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res1, GeoIPInformation)
+# assert res1.ip == ip
+# # It looks up in insight using the normalize IP!
+# assert f"get_insights_ip_information: {normalized_ip}" in caplog.text
+
+# # And it should NOT do the lookup again with an ipv6 in the same /64 block!
+# ip = IP_v6_US_SAME_64
+# caplog.clear()
+# with caplog.at_level(logging.INFO):
+# res2 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert "get_insights_ip_information" not in caplog.text
+
+# def test_run_ip_information(self, maxmind_manager):
+# ip = IP_v4_US
+
+# res = maxmind_manager.run_ip_information(ip_address=ip)
+# assert isinstance(res, GeoIPInformation)
+# assert res.country_name == "United States"
+# assert res.country_iso == "us"
diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py
index 31087b8..52bbbec 100644
--- a/tests/managers/thl/test_payout.py
+++ b/tests/managers/thl/test_payout.py
@@ -1,25 +1,52 @@
+import io
import logging
import os
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from random import choice as rand_choice, randint
-from typing import Optional
+from random import choice as rand_choice
+from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
from generalresearch.currency import USDCent
-from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerTransactionConditionFailedError,
-)
-from generalresearch.managers.thl.payout import UserPayoutEventManager
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.ledger import LedgerEntry, Direction
-from generalresearch.models.thl.payout import BusinessPayoutEvent
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.ledger import LedgerAccount
+from generalresearch.models.thl.finance import BusinessBalances
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+ BusinessPayoutEvent,
+)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.thl_web import (
+ LedgerDFCollection,
+ )
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.payout import (
+ BrokerageProductPayoutEventManager,
+ BusinessPayoutEventManager,
+ PayoutEventManager,
+ UserPayoutEventManager,
+ )
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.thl.ledger import LedgerAccount
+ from generalresearch.models.thl.payout import (
+ UserPayoutEvent,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
logger = logging.getLogger()
@@ -27,17 +54,15 @@ cashout_method_uuid = uuid4().hex
class TestPayout:
-
def test_get_by_uuid_and_create(
self,
- user,
+ user: User,
user_payout_event_manager: UserPayoutEventManager,
- thl_lm,
- utc_now,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
):
-
- user_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet(
- user=user
+ user_account: LedgerAccount = (
+ thl_ledger_manager.get_account_or_create_user_wallet(user=user)
)
pe1: UserPayoutEvent = user_payout_event_manager.create(
@@ -57,11 +82,14 @@ class TestPayout:
assert pe1 == pe2
- def test_update(self, user, user_payout_event_manager, lm, thl_lm, utc_now):
- from generalresearch.models.thl.definitions import PayoutStatus
- from generalresearch.models.thl.wallet import PayoutType
-
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ def test_update(
+ self,
+ user: User,
+ user_payout_event_manager: UserPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
+ ):
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
pe1 = user_payout_event_manager.create(
status=PayoutStatus.PENDING,
@@ -89,142 +117,84 @@ class TestPayout:
def test_create_bp_payout(
self,
- user,
- thl_web_rr,
- user_payout_event_manager,
- lm,
- thl_lm,
- product,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
):
- delete_ledger_db()
- create_main_accounts()
- from generalresearch.models.thl.ledger import LedgerAccount
-
- thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- # wallet balance failure
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- )
-
- # (we don't have a special method for this) Put money in the BP's account
- amount_cents = 100
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- bp_wallet: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(
- product=product
- )
-
- entries = [
- LedgerEntry(
- direction=Direction.DEBIT,
- account_uuid=cash_account.uuid,
- amount=amount_cents,
- ),
- LedgerEntry(
- direction=Direction.CREDIT,
- account_uuid=bp_wallet.uuid,
- amount=amount_cents,
- ),
- ]
-
- lm.create_tx(entries=entries)
- assert 100 == lm.get_account_balance(account=bp_wallet)
-
- # Then run it again for $1.00
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- )
- assert 0 == lm.get_account_balance(account=bp_wallet)
-
- # Run again should without balance check, should still fail due to day check
- with pytest.raises(LedgerTransactionConditionFailedError):
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=False,
- )
+ # create_bp_payout_event does not get called directly. We have tests
+ # for the ledger methods already
+ pass
- # And then we can run again skip both checks
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
+ @pytest.fixture
+ def pending_bp_pe(
+ self,
+ thl_web_rw: PostgresConfig,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ utc_now: datetime,
+ ) -> BrokerageProductPayoutEvent:
+ account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
+ bp_pe = BrokerageProductPayoutEvent(
+ product_id=product.uuid,
amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
- )
- assert -100 == lm.get_account_balance(account=bp_wallet)
-
- pe = brokerage_product_payout_event_manager.get_by_uuid(pe.uuid)
- txs = lm.get_tx_filtered_by_metadata(
- metadata_key="event_payout", metadata_value=pe.uuid
+ payout_type=PayoutType.ACH,
+ debit_account_uuid=account.uuid,
+ cashout_method_uuid=brokerage_product_payout_event_manager.CASHOUT_METHOD_UUID,
+ created=utc_now,
)
-
- assert 1 == len(txs)
+ params = bp_pe.model_dump_postgres()
+ # This shouldn't exist. For testing only, so no supplier_payout
+ params["supplier_payout_id"] = None
+ thl_web_rw.execute_write(
+ """
+ INSERT INTO event_payout (uuid, debit_account_uuid, created, cashout_method_uuid,
+ amount, status, ext_ref_id, payout_type, order_data,
+ request_data, supplier_payout_id)
+ VALUES (%(uuid)s, %(debit_account_uuid)s, %(created)s, %(cashout_method_uuid)s,
+ %(amount)s, %(status)s, %(ext_ref_id)s, %(payout_type)s, %(order_data)s,
+ %(request_data)s, %(supplier_payout_id)s);
+ """,
+ params,
+ )
+ return bp_pe
def test_create_bp_payout_quick_dupe(
self,
- user,
- product,
- thl_web_rw,
- brokerage_product_payout_event_manager,
- thl_lm,
- lm,
- utc_now,
- create_main_accounts,
+ product: Product,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
+ pending_bp_pe: BrokerageProductPayoutEvent,
):
- thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ bp_pe=pending_bp_pe,
product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
created=utc_now,
)
with pytest.raises(ValueError) as cm:
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ bp_pe=pending_bp_pe,
created=utc_now,
)
assert "Payout event already exists!" in str(cm.value)
def test_filter(
self,
- thl_web_rw,
- thl_lm,
- lm,
- product,
- user,
- user_payout_event_manager,
- utc_now,
+ thl_ledger_manager: ThlLedgerManager,
+ product: Product,
+ user: User,
+ user_payout_event_manager: UserPayoutEventManager,
+ utc_now: datetime,
):
from generalresearch.models.thl.definitions import PayoutStatus
- from generalresearch.models.thl.wallet import PayoutType
+ from generalresearch.models.thl.wallet.definitions import PayoutType
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
- bp_account = thl_lm.get_account_or_create_bp_wallet(product=product)
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
+ bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
user_payout_event_manager.create(
status=PayoutStatus.PENDING,
@@ -289,166 +259,109 @@ class TestPayout:
class TestPayoutEventManager:
-
- def test_set_account_lookup_table(
- self, payout_event_manager, thl_redis_config, thl_lm, delete_ledger_db
- ):
- delete_ledger_db()
- rc = thl_redis_config.create_redis_client()
- rc.delete("pem:account_to_product")
- rc.delete("pem:product_to_account")
- N = 5
-
- for idx in range(N):
- thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=uuid4().hex)
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == 0
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == 0
-
- payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm,
- )
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == N
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == N
-
- thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=uuid4().hex)
- payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm,
- )
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == N + 1
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == N + 1
+ pass
class TestBusinessPayoutEventManager:
-
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
return "5d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return timedelta(days=10)
def test_base(
self,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- business,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ gr_business: Business,
):
delete_ledger_db()
create_main_accounts()
- from generalresearch.models.thl.product import Product
-
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
ach_id1 = uuid4().hex
ach_id2 = uuid4().hex
- bp_payout_factory(
- product=p1,
- amount=USDCent(1),
- ext_ref_id=None,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ # ext_ref_id is required now
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(1), ext_ref_id="none"
)
- bp_payout_factory(
- product=p1,
- amount=USDCent(1),
- ext_ref_id=ach_id1,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(1), ext_ref_id=ach_id1
)
+ with pytest.raises(
+ expected_exception=ValueError,
+ match="Cannot create a BusinessPayoutEvent with an existing transaction_id",
+ ):
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(25), ext_ref_id=ach_id1
+ )
- bp_payout_factory(
- product=p1,
- amount=USDCent(25),
- ext_ref_id=ach_id1,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(50), ext_ref_id=ach_id2
)
- bp_payout_factory(
- product=p1,
- amount=USDCent(50),
- ext_ref_id=ach_id2,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ gr_business.prebuild_payouts(
+ bpem=business_payout_event_manager,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
- bpem=business_payout_event_manager,
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 3
+ assert gr_business.payouts_total == sum(
+ [pe.amount for pe in gr_business.payouts]
)
+ assert gr_business.payouts[0].created > gr_business.payouts[1].created
+ assert len(gr_business.payouts[0].bp_payouts) == 1
- assert len(business.payouts) == 3
- assert business.payouts_total == sum([pe.amount for pe in business.payouts])
- assert business.payouts[0].created > business.payouts[1].created
- assert len(business.payouts[0].bp_payouts) == 1
- assert len(business.payouts[1].bp_payouts) == 2
+ # Cannot pay out the same product twice in the same business payout
+ # assert len(business.payouts[1].bp_payouts) == 2
+ assert len(gr_business.payouts[1].bp_payouts) == 1
- assert business.payouts[0].ext_ref_id == ach_id2
- assert business.payouts[1].ext_ref_id == ach_id1
- assert business.payouts[2].ext_ref_id is None
+ assert gr_business.payouts[0].ext_ref_id == ach_id2
+ assert gr_business.payouts[1].ext_ref_id == ach_id1
+ assert gr_business.payouts[2].ext_ref_id == "none"
def test_update_ext_reference_ids(
self,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- delete_df_collection,
- user_factory,
- ledger_collection,
- session_with_tx_factory,
- pop_ledger_merge,
- client_no_amm,
- mnt_filepath,
- lm,
- product_manager,
- start,
- business,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ user_factory: Callable[..., User],
+ ledger_collection: LedgerDFCollection,
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ product_manager: ProductManager,
+ start: datetime,
+ gr_business: Business,
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
# $250.00 to work with
for idx in range(1, 10):
@@ -461,32 +374,34 @@ class TestBusinessPayoutEventManager:
ach_id1 = uuid4().hex
ach_id2 = uuid4().hex
- with pytest.raises(expected_exception=Warning) as cm:
+ with pytest.raises(
+ expected_exception=AssertionError, match="No Business Payout found"
+ ):
business_payout_event_manager.update_ext_reference_ids(
new_value=ach_id2,
current_value=ach_id1,
)
- assert "No event_payouts found to UPDATE" in str(cm)
# We must build the balance to issue ACH/Wire
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
res = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(100_01),
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
transaction_id=ach_id1,
)
assert isinstance(res, BusinessPayoutEvent)
+ assert business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
# Okay, now that there is a payout_event, let's try to update the
# ext_reference_id
@@ -495,111 +410,17 @@ class TestBusinessPayoutEventManager:
current_value=ach_id1,
)
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- assert len(res) == 0
+ with pytest.raises(
+ expected_exception=AssertionError, match="No Business Payout found"
+ ):
+ business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id2)
- assert len(res) == 1
+ assert business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id2)
- def test_delete_failed_business_payout(
- self,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- currency,
- delete_df_collection,
- user_factory,
- ledger_collection,
- session_with_tx_factory,
- pop_ledger_merge,
- client_no_amm,
- mnt_filepath,
- lm,
- product_manager,
- start,
- business,
+ def test_recoup_empty(
+ self, business_payout_event_manager: BusinessPayoutEventManager
):
- delete_ledger_db()
- create_main_accounts()
- delete_df_collection(coll=ledger_collection)
-
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- u1: User = user_factory(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- # $250.00 to work with
- for idx in range(1, 10):
- session_with_tx_factory(
- user=u1,
- wall_req_cpi=Decimal("25.00"),
- started=start + timedelta(days=1, minutes=idx),
- )
-
- # We must build the balance to issue ACH/Wire
- ledger_collection.initial_load(client=None, sync=True)
- pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
- ds=mnt_filepath,
- client=client_no_amm,
- pop_ledger=pop_ledger_merge,
- )
-
- ach_id1 = uuid4().hex
-
- res = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
- amount=USDCent(100_01),
- pm=product_manager,
- thl_lm=thl_lm,
- transaction_id=ach_id1,
- )
- assert isinstance(res, BusinessPayoutEvent)
-
- # (1) Confirm the initial Event Payout, Tx, TxMeta, TxEntry all exist
- event_payouts = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- event_payout_uuids = [i.uuid for i in event_payouts]
- assert len(event_payout_uuids) == 1
- tags = [f"{currency.value}:bp_payout:{x}" for x in event_payout_uuids]
- transactions = thl_lm.get_txs_by_tags(tags=tags)
- assert len(transactions) == 1
- tx_metadata_ids = thl_lm.get_tx_metadata_ids_by_txs(transactions=transactions)
- assert len(tx_metadata_ids) == 2
- tx_entries = thl_lm.get_tx_entries_by_txs(transactions=transactions)
- assert len(tx_entries) == 2
-
- # (2) Delete!
- business_payout_event_manager.delete_failed_business_payout(
- ext_ref_id=ach_id1, thl_lm=thl_lm
- )
-
- # (3) Confirm the initial Event Payout, Tx, TxMeta, TxEntry have
- # all been deleted
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- assert len(res) == 0
-
- # Note: b/c the event_payout shouldn't exist anymore, we are taking
- # the tag strings and transactions from when they did..
- res = thl_lm.get_txs_by_tags(tags=tags)
- assert len(res) == 0
-
- tx_metadata_ids = thl_lm.get_tx_metadata_ids_by_txs(transactions=transactions)
- assert len(tx_metadata_ids) == 0
- tx_entries = thl_lm.get_tx_entries_by_txs(transactions=transactions)
- assert len(tx_entries) == 0
-
- def test_recoup_empty(self, business_payout_event_manager):
- res = {uuid4().hex: USDCent(0) for i in range(100)}
+ res = {uuid4().hex: USDCent(0) for _ in range(100)}
df = pd.DataFrame.from_dict(res, orient="index").reset_index()
df.columns = ["product_id", "available_balance"]
@@ -609,10 +430,12 @@ class TestBusinessPayoutEventManager:
)
assert "Total available amount is empty, cannot recoup" in str(cm)
- def test_recoup_exceeds(self, business_payout_event_manager):
+ def test_recoup_exceeds(
+ self, business_payout_event_manager: BusinessPayoutEventManager
+ ):
from random import randint
- res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for i in range(100)}
+ res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for _ in range(100)}
df = pd.DataFrame.from_dict(res, orient="index").reset_index()
df.columns = ["product_id", "available_balance"]
@@ -624,10 +447,10 @@ class TestBusinessPayoutEventManager:
)
assert " exceeds total available " in str(cm)
- def test_recoup(self, business_payout_event_manager):
+ def test_recoup(self, business_payout_event_manager: BusinessPayoutEventManager):
from random import randint, random
- res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for i in range(100)}
+ res = {uuid4().hex: USDCent(randint(a=0, b=1_000_00)) for _ in range(100)}
df = pd.DataFrame.from_dict(res, orient="index").reset_index()
df.columns = ["product_id", "available_balance"]
@@ -643,7 +466,9 @@ class TestBusinessPayoutEventManager:
assert res.deduction.sum() == random_recoup_amount
assert res.remaining_balance.sum() == avail_balance - random_recoup_amount
- def test_recoup_loop(self, business_payout_event_manager, request):
+ def test_recoup_loop(
+ self, business_payout_event_manager: BusinessPayoutEventManager, request
+ ):
# TODO: Generate this file at random
fp = os.path.join(
request.config.rootpath, "data/pytest_recoup_proportional.csv"
@@ -657,9 +482,11 @@ class TestBusinessPayoutEventManager:
assert int(res.deduction.sum()) == 1416089
- def test_recoup_loop_single_profitable_account(self, business_payout_event_manager):
- res = [{"product_id": uuid4().hex, "available_balance": 0} for i in range(1000)]
- for x in range(100):
+ def test_recoup_loop_single_profitable_account(
+ self, business_payout_event_manager: BusinessPayoutEventManager
+ ):
+ res = [{"product_id": uuid4().hex, "available_balance": 0} for _ in range(1000)]
+ for _ in range(100):
item = rand_choice(res)
item["available_balance"] = randint(8, 12)
@@ -670,14 +497,16 @@ class TestBusinessPayoutEventManager:
# res = res[res["remaining_balance"] > 0]
assert int(res.deduction.sum()) == 500
- def test_recoup_loop_assertions(self, business_payout_event_manager):
+ def test_recoup_loop_assertions(
+ self, business_payout_event_manager: BusinessPayoutEventManager
+ ):
df = pd.DataFrame(
[
{
"product_id": uuid4().hex,
"available_balance": randint(0, 999_999),
}
- for i in range(10_000)
+ for _ in range(10_000)
]
)
available_balance = int(df.available_balance.sum())
@@ -697,7 +526,7 @@ class TestBusinessPayoutEventManager:
assert int(res.deduction.sum()) == available_balance - 1
# Slightly less
- with pytest.raises(expected_exception=Exception) as cm:
+ with pytest.raises(expected_exception=ValueError):
res = business_payout_event_manager.recoup_proportional(
df=df, target_amount=available_balance + 1
)
@@ -707,8 +536,9 @@ class TestBusinessPayoutEventManager:
assert res.remaining_balance.sum() == available_balance
assert int(res.deduction.sum()) == 0
- def test_distribute_amount(self, business_payout_event_manager):
- import io
+ def test_distribute_amount(
+ self, business_payout_event_manager: BusinessPayoutEventManager
+ ):
df = pd.read_csv(
io.StringIO(
@@ -724,31 +554,31 @@ class TestBusinessPayoutEventManager:
def test_ach_payment_min_amount(
self,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- payout_event_manager,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
- product_manager,
+ product: Product,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ thl_redis_config: RedisConfig,
+ payout_event_manager: PayoutEventManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
):
- """Test having a Business with three products.. one that lost money
+ """Test having a Business with three products. One that lost money
and two that gained money. Ensure that the Business balance
reflects that to compensate for the Product in the negative and only
assigns Brokerage Product payments from the 2 accounts that have
@@ -759,12 +589,9 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
session_with_tx_factory(
user=u1,
@@ -776,8 +603,7 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("5.00"),
started=start + timedelta(days=6),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(475), # 95% of $5.00
created=start + timedelta(days=1, minutes=1),
@@ -785,9 +611,9 @@ class TestBusinessPayoutEventManager:
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
@@ -795,40 +621,151 @@ class TestBusinessPayoutEventManager:
with pytest.raises(expected_exception=AssertionError) as cm:
business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(500),
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
+ transaction_id=uuid4().hex,
)
assert "Must issue Supplier Payouts at least $100 minimum." in str(cm)
+ def test_create_from_ach_or_wire(
+ self,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., None],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ caplog,
+ ):
+ """Test having a Business with three products"""
+ # Now let's load it up and actually test some things
+ delete_ledger_db()
+ create_main_accounts()
+ delete_df_collection(coll=ledger_collection)
+
+ p1: Product = product_factory(business=gr_business)
+ p2: Product = product_factory(business=gr_business)
+ p3: Product = product_factory(business=gr_business)
+ _: User = user_factory(product=p1)
+ u2: User = user_factory(product=p2)
+ u3: User = user_factory(product=p3)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p3)
+
+ ach_id1 = uuid4().hex
+ ach_id2 = uuid4().hex
+
+ # Product 1: Complete $10 x 20
+ for idx in range(20):
+ session_with_tx_factory(
+ user=u2,
+ wall_req_cpi=Decimal("10.00"),
+ started=start + timedelta(days=1, hours=2, minutes=1 + idx),
+ )
+
+ # Product 2: Complete $10 x 30
+ for idx in range(30):
+ session_with_tx_factory(
+ user=u3,
+ wall_req_cpi=Decimal("10.00"),
+ started=start + timedelta(days=1, hours=3, minutes=1 + idx),
+ )
+
+ ledger_collection.initial_load(client=None, sync=True)
+ pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
+ gr_business.prebuild_balance(
+ thl_pg_config=thl_web_rr,
+ lm=ledger_manager,
+ ds=mnt_filepath,
+ client=client_no_amm,
+ pop_ledger=pop_ledger_merge,
+ )
+
+ bb = gr_business.balance
+ assert isinstance(bb, BusinessBalances)
+ assert bb.payout == 475_00 # $500 * .95% = $475
+ assert bb.net == 475_00
+
+ bp1 = business_payout_event_manager.create_from_ach_or_wire(
+ business=gr_business,
+ amount=USDCent(100_00),
+ pm=product_manager,
+ thl_lm=thl_ledger_manager,
+ created=start + timedelta(days=1, hours=5),
+ transaction_id=ach_id1,
+ )
+ print(f"{bp1=}")
+ assert isinstance(bp1, BusinessPayoutEvent)
+ assert len(bp1.bp_payouts) == 2
+
+ bp2 = business_payout_event_manager.create_from_ach_or_wire(
+ business=gr_business,
+ amount=USDCent(bb.available_balance),
+ pm=product_manager,
+ thl_lm=thl_ledger_manager,
+ created=start + timedelta(days=2, hours=5),
+ transaction_id=ach_id2,
+ )
+ print(f"{bp2=}")
+ assert isinstance(bp2, BusinessPayoutEvent)
+ assert len(bp2.bp_payouts) == 2
+
+ with caplog.at_level(logging.WARNING):
+ business_payout_event_manager.resume_failed_business_payout(
+ ext_ref_id=ach_id1, thl_lm=thl_ledger_manager, pm=product_manager
+ )
+ assert "Nothing to do!" in caplog.text
+
+ # bpe = business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
+ # bp_pe = bpe.bp_payouts[0]
+ # thl_web_rr.execute_write(
+ # """
+ # UPDATE event_payout
+ # SET status = %(status)s
+ # WHERE uuid = %(uuid)s""",
+ # {"uuid": bp_pe.uuid, "status": PayoutStatus.FAILED},
+ # )
+
def test_ach_payment(
self,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- payout_event_manager,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
- product_manager,
- rm_ledger_collection,
- rm_pop_ledger_merge,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., None],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ rm_ledger_collection: Callable[..., None],
+ rm_pop_ledger_merge: Callable[..., None],
):
"""Test having a Business with three products.. one that lost money
and two that gained money. Ensure that the Business balance
@@ -841,21 +778,17 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- p2: Product = product_factory(business=business)
- p3: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
+ p2: Product = product_factory(business=gr_business)
+ p3: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p2)
u3: User = user_factory(product=p3)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
- thl_lm.get_account_or_create_bp_wallet(product=p3)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p3)
ach_id1 = uuid4().hex
- ach_id2 = uuid4().hex
# Product 1: Complete, Payout, Recon..
s1 = session_with_tx_factory(
@@ -863,14 +796,11 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("5.00"),
started=start + timedelta(days=1),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(475), # 95% of $5.00
ext_ref_id=ach_id1,
created=start + timedelta(days=1, minutes=1),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
adj_to_fail_with_tx_factory(
session=s1,
@@ -893,18 +823,19 @@ class TestBusinessPayoutEventManager:
started=start + timedelta(days=1, hours=3, minutes=1 + idx),
)
- # Now that we paid out the business, let's confirm the updated balances
+ # Now that we paid out the gr_business: Business, let's confirm the updated balances
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- bb1 = business.balance
+ bb1 = gr_business.balance
+ assert isinstance(bb1, BusinessBalances)
pb1 = bb1.product_balances[0]
pb2 = bb1.product_balances[1]
pb3 = bb1.product_balances[2]
@@ -927,20 +858,21 @@ class TestBusinessPayoutEventManager:
assert pb2.recoup_usd_str == "$0.00"
assert pb3.recoup_usd_str == "$0.00"
- assert business.payouts is None
- business.prebuild_payouts(
+ assert gr_business.payouts is None
+ gr_business.prebuild_payouts(
thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert business.payouts[0].ext_ref_id == ach_id1
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 1
+ assert gr_business.payouts[0].ext_ref_id == ach_id1
bp1 = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(bb1.available_balance),
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=1, hours=5),
)
assert isinstance(bp1, BusinessPayoutEvent)
@@ -948,7 +880,7 @@ class TestBusinessPayoutEventManager:
assert bp1.bp_payouts[0].status == PayoutStatus.COMPLETE
assert bp1.bp_payouts[1].status == PayoutStatus.COMPLETE
bp1_tx = brokerage_product_payout_event_manager.check_for_ledger_tx(
- thl_ledger_manager=thl_lm,
+ thl_ledger_manager=thl_ledger_manager,
payout_event=bp1.bp_payouts[0],
product_id=bp1.bp_payouts[0].product_id,
amount=bp1.bp_payouts[0].amount,
@@ -956,14 +888,14 @@ class TestBusinessPayoutEventManager:
assert bp1_tx
bp2_tx = brokerage_product_payout_event_manager.check_for_ledger_tx(
- thl_ledger_manager=thl_lm,
+ thl_ledger_manager=thl_ledger_manager,
payout_event=bp1.bp_payouts[1],
product_id=bp1.bp_payouts[1].product_id,
amount=bp1.bp_payouts[1].amount,
)
assert bp2_tx
- # Now that we paid out the business, let's confirm the updated balances
+ # Now that we paid out the business: Business, let's confirm the updated balances
rm_ledger_collection()
rm_pop_ledger_merge()
ledger_collection.initial_load(client=None, sync=True)
@@ -971,16 +903,15 @@ class TestBusinessPayoutEventManager:
business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
bpem=business_payout_event_manager,
)
+ assert isinstance(business.payouts, list)
assert len(business.payouts) == 2
assert len(business.payouts[0].bp_payouts) == 2
assert len(business.payouts[1].bp_payouts) == 1
@@ -989,6 +920,8 @@ class TestBusinessPayoutEventManager:
# Okay os we have the balance before, and after the Business Payout
# of bb1.available_balance worth..
+ assert isinstance(bb1, BusinessBalances)
+ assert isinstance(bb2, BusinessBalances)
assert bb1.payout == bb2.payout
assert bb1.adjustment == bb2.adjustment
assert bb1.net == bb2.net
@@ -1005,34 +938,29 @@ class TestBusinessPayoutEventManager:
def test_ach_payment_partial_amount(
self,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- payout_event_manager,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
- product_manager,
- rm_ledger_collection,
- rm_pop_ledger_merge,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ payout_event_manager: PayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., None],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ rm_ledger_collection: Callable[..., None],
+ rm_pop_ledger_merge: Callable[..., None],
):
"""There are valid instances when we want issue a ACH or Wire to a
- Business, but not for the full Available Balance amount in their
+ gr_business: Business, but not for the full Available Balance amount in their
account.
To test this, we'll create a Business with multiple Products, and
@@ -1047,18 +975,15 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- p2: Product = product_factory(business=business)
- p3: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
+ p2: Product = product_factory(business=gr_business)
+ p3: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p2)
u3: User = user_factory(product=p3)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
- thl_lm.get_account_or_create_bp_wallet(product=p3)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p3)
# Product 1, 2, 3: Complete, and Payout multiple times.
for idx in range(5):
@@ -1068,27 +993,27 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("50.00"),
started=start + timedelta(days=1, hours=2, minutes=1 + idx),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
- # Now that we paid out the business, let's confirm the updated balances
+ # Now that we paid out the business: Business, let's confirm the updated balances
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
# Confirm the initial amounts.
- assert len(business.payouts) == 0
- bb1 = business.balance
+ assert len(gr_business.payouts) == 0
+ bb1 = gr_business.balance
+
+ assert isinstance(bb1, BusinessBalances)
assert bb1.payout == 3 * 5 * 4750
assert bb1.adjustment == 0
assert bb1.payout == bb1.net
@@ -1100,24 +1025,25 @@ class TestBusinessPayoutEventManager:
assert bb1.product_balances[x].balance == 5 * 4750
assert bb1.product_balances[x].available_balance_usd_str == "$178.13"
- assert business.payouts_total_str == "$0.00"
- assert business.balance.payment_usd_str == "$0.00"
- assert business.balance.available_balance_usd_str == "$534.39"
+ assert gr_business.payouts_total_str == "$0.00"
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payment_usd_str == "$0.00"
+ assert gr_business.balance.available_balance_usd_str == "$534.39"
# This is the important part, even those the Business has $534.39
# available to it, we are only trying to issue out a $250.00 ACH or
# Wire to the Business
bp1 = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(250_00),
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=1, hours=3),
)
assert isinstance(bp1, BusinessPayoutEvent)
assert len(bp1.bp_payouts) == 3
- # Now that we paid out the business, let's confirm the updated
+ # Now that we paid out the gr_business: Business, let's confirm the updated
# balances. Clear and rebuild the parquet files.
rm_ledger_collection()
rm_pop_ledger_merge()
@@ -1127,51 +1053,46 @@ class TestBusinessPayoutEventManager:
# Now rebuild and confirm the payouts, balance.payment, and the
# balance.available_balance are reflective of having a $250 ACH/Wire
# sent to the Business
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 3
- assert business.payouts_total_str == "$250.00"
- assert business.balance.payment_usd_str == "$250.00"
- assert business.balance.available_balance_usd_str == "$346.88"
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 1
+ assert len(gr_business.payouts[0].bp_payouts) == 3
+ assert gr_business.payouts_total_str == "$250.00"
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payment_usd_str == "$250.00"
+ assert gr_business.balance.available_balance_usd_str == "$346.88"
def test_ach_tx_id_reference(
self,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- payout_event_manager,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
- product_manager,
- rm_ledger_collection,
- rm_pop_ledger_merge,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ payout_event_manager: PayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ rm_ledger_collection: Callable[..., None],
+ rm_pop_ledger_merge: Callable[..., None],
):
# Now let's load it up and actually test some things
@@ -1179,18 +1100,15 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- p2: Product = product_factory(business=business)
- p3: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
+ p2: Product = product_factory(business=gr_business)
+ p3: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p2)
u3: User = user_factory(product=p3)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
- thl_lm.get_account_or_create_bp_wallet(product=p3)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p3)
ach_id1 = uuid4().hex
ach_id2 = uuid4().hex
@@ -1202,26 +1120,26 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("7.50"),
started=start + timedelta(days=1, hours=1 + iidx, minutes=1 + idx),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager)
rm_ledger_collection()
rm_pop_ledger_merge()
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
bp1 = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(100_01),
transaction_id=ach_id1,
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=2, hours=1),
)
@@ -1229,20 +1147,20 @@ class TestBusinessPayoutEventManager:
rm_pop_ledger_merge()
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
bp2 = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
+ business=gr_business,
amount=USDCent(100_02),
transaction_id=ach_id2,
pm=product_manager,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=4, hours=1),
)
@@ -1253,17 +1171,18 @@ class TestBusinessPayoutEventManager:
rm_pop_ledger_merge()
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_payouts(
+ gr_business.prebuild_payouts(
thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
- business.prebuild_balance(
+ gr_business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.payouts[0].ext_ref_id == ach_id2
- assert business.payouts[1].ext_ref_id == ach_id1
+ assert isinstance(gr_business.payouts, list)
+ assert gr_business.payouts[0].ext_ref_id == ach_id2
+ assert gr_business.payouts[1].ext_ref_id == ach_id1
diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py
index 78d5dde..8d72fa5 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -1,24 +1,35 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.product import (
Product,
+ ProfilingConfig,
SourceConfig,
- UserCreateConfig,
SourcesConfig,
- UserHealthConfig,
- ProfilingConfig,
- SupplyPolicy,
SupplyConfig,
+ SupplyPolicy,
+ UserCreateConfig,
+ UserHealthConfig,
)
-from test_utils.models.conftest import product_factory
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.team import Team
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager):
- product: Product = product_manager.create_dummy(
+ def test_get_by_uuid(
+ self,
+ product_manager: ProductManager,
+ product_factory: Callable[..., Product],
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -37,12 +48,14 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuid(product_uuid=uuid4().hex)
assert "product not found" in str(cm.value)
- def test_get_by_uuids(self, product_manager):
+ def test_get_by_uuids(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
cnt = 5
- product_uuids = [uuid4().hex for idx in range(cnt)]
+ product_uuids = [uuid4().hex for _ in range(cnt)]
for product_id in product_uuids:
- product_manager.create_dummy(
+ product_factory(
product_id=product_id,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -62,8 +75,10 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"])
assert "invalid uuid" in str(cm.value)
- def test_get_by_uuid_if_exists(self, product_manager):
- product: Product = product_manager.create_dummy(
+ def test_get_by_uuid_if_exists(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -74,10 +89,12 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
assert instance == None
- def test_get_by_uuids_if_exists(self, product_manager):
+ def test_get_by_uuids_if_exists(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
product_uuids = [uuid4().hex for _ in range(2)]
for product_id in product_uuids:
- product_manager.create_dummy(
+ product_factory(
product_id=product_id,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -106,13 +123,15 @@ class TestProductManagerGetMethods:
# for instance in res:
# assert isinstance(instance, Product)
- def test_get_by_business_ids(self, product_manager):
- business_ids = [uuid4().hex for i in range(5)]
+ def test_get_by_business_ids(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ business_ids = [uuid4().hex for _ in range(5)]
product_manager.fetch_uuids(business_uuids=business_ids)
for business_id in business_ids:
- product_manager.create(
+ product_factory(
product_id=uuid4().hex,
team_id=None,
business_id=business_id,
@@ -124,8 +143,10 @@ class TestProductManagerGetMethods:
class TestProductManagerCreation:
- def test_base(self, product_manager):
- instance = product_manager.create_dummy(
+ def test_base(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ instance = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"New Test Product {uuid4().hex[:6]}",
@@ -136,7 +157,7 @@ class TestProductManagerCreation:
class TestProductManagerCreate:
- def test_create_simple(self, product_manager):
+ def test_create_simple(self, product_manager: ProductManager):
# Always required: product_id, team_id, name, redirect_url
# Required internally - if not passed use default: harmonizer_domain,
# commission_pct, sources
@@ -179,20 +200,26 @@ class TestProductManager:
]
]
- def test_get_by_uuid1(self, product_manager, team, product, product_factory):
- p1 = product_factory(team=team)
+ def test_get_by_uuid1(
+ self,
+ product_manager: ProductManager,
+ gr_team: Team,
+ product: Product,
+ product_factory: Callable[..., Product],
+ ):
+ p1 = product_factory(team=gr_team)
instance = product_manager.get_by_uuid(product_uuid=p1.uuid)
assert instance.id == p1.id
# No Team and no user_create_config
- assert instance.team_id == team.uuid
+ assert instance.team_id == gr_team.uuid
# user_create_config can't be None, so ensure the default was set.
assert isinstance(instance.user_create_config, UserCreateConfig)
assert 0 == instance.user_create_config.min_hourly_create_limit
assert instance.user_create_config.max_hourly_create_limit is None
- def test_get_by_uuid2(self, product_manager, product_factory):
+ def test_get_by_uuid2(self, product_manager: ProductManager, product_factory):
p2 = product_factory()
instance = product_manager.get_by_uuid(p2.id)
assert instance.id, p2.id
@@ -204,7 +231,9 @@ class TestProductManager:
assert 0 == instance.user_create_config.min_hourly_create_limit
assert instance.user_create_config.max_hourly_create_limit is None
- def test_get_by_uuid3(self, product_manager, product_factory):
+ def test_get_by_uuid3(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
p3 = product_factory()
instance = product_manager.get_by_uuid(p3.id)
assert instance.id == p3.id
@@ -220,10 +249,12 @@ class TestProductManager:
assert instance.user_create_config.max_hourly_create_limit is None
assert not instance.user_wallet_config.enabled
- def test_sources(self, product_manager):
+ def test_sources(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
user_defined = [SourceConfig(name=Source.DYNATA, active=False)]
sources_config = SourcesConfig(user_defined=user_defined)
- p = product_manager.create_dummy(sources_config=sources_config)
+ p = product_factory(sources_config=sources_config)
p2 = product_manager.get_by_uuid(p.id)
@@ -235,7 +266,9 @@ class TestProductManager:
assert not dynata.active
assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA)
- def test_global_sources(self, product_manager):
+ def test_global_sources(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
sources_config = SupplyConfig(
policies=[
SupplyPolicy(
@@ -246,7 +279,7 @@ class TestProductManager:
)
]
)
- p1 = product_manager.create_dummy(sources_config=sources_config)
+ p1 = product_factory(sources_config=sources_config)
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
@@ -262,8 +295,10 @@ class TestProductManager:
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
- def test_user_health_config(self, product_manager):
- p = product_manager.create_dummy(
+ def test_user_health_config(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory(
user_health_config=UserHealthConfig(banned_countries=["ng", "in"])
)
@@ -273,10 +308,10 @@ class TestProductManager:
assert p2.user_health_config.banned_countries == ["in", "ng"]
assert p2.user_health_config.allow_ban_iphist
- def test_profiling_config(self, product_manager):
- p = product_manager.create_dummy(
- profiling_config=ProfilingConfig(max_questions=1)
- )
+ def test_profiling_config(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory(profiling_config=ProfilingConfig(max_questions=1))
p2 = product_manager.get_by_uuid(p.id)
assert p == p2
@@ -320,8 +355,10 @@ class TestProductManager:
class TestProductManagerUpdate:
- def test_update(self, product_manager):
- p = product_manager.create_dummy()
+ def test_update(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory()
p.name = "new name"
p.enabled = False
p.user_create_config = UserCreateConfig(min_hourly_create_limit=200)
@@ -341,8 +378,10 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
- def test_cache_clear(self, product_manager):
- p = product_manager.create_dummy()
+ def test_cache_clear(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory()
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.pg_config.execute_write(
diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py
index 7b4f677..d584527 100644
--- a/tests/managers/thl/test_product_prod.py
+++ b/tests/managers/thl/test_product_prod.py
@@ -1,17 +1,25 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.models.thl.product import Product
-from test_utils.models.conftest import product_factory
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.product import ProductManager
logger = logging.getLogger()
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager, product_factory):
+ def test_get_by_uuid(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
# Just test that we load properly
for p in [product_factory(), product_factory(), product_factory()]:
instance = product_manager.get_by_uuid(product_uuid=p.id)
@@ -23,7 +31,9 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuid(product_uuid=uuid4().hex)
assert "product not found" in str(cm.value)
- def test_get_by_uuids(self, product_manager, product_factory):
+ def test_get_by_uuids(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
products = [product_factory(), product_factory(), product_factory()]
cnt = len(products)
res = product_manager.get_by_uuids(product_uuids=[p.id for p in products])
@@ -43,7 +53,9 @@ class TestProductManagerGetMethods:
)
assert "invalid uuid passed" in str(cm.value)
- def test_get_by_uuid_if_exists(self, product_factory, product_manager):
+ def test_get_by_uuid_if_exists(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
products = [product_factory(), product_factory(), product_factory()]
instance = product_manager.get_by_uuid_if_exists(product_uuid=products[0].id)
@@ -52,7 +64,9 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
assert instance is None
- def test_get_by_uuids_if_exists(self, product_manager, product_factory):
+ def test_get_by_uuids_if_exists(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
products = [product_factory(), product_factory(), product_factory()]
res = product_manager.get_by_uuids_if_exists(
@@ -75,8 +89,7 @@ class TestProductManagerGetMethods:
class TestProductManagerGetAll:
@pytest.mark.skip(reason="TODO")
- def test_get_ALL_by_ids(self, product_manager):
+ def test_get_ALL_by_ids(self, product_manager: ProductManager):
products = product_manager.get_all(rand_limit=50)
logger.info(f"Fetching {len(products)} product uuids")
# todo: once timebucks stops spamming broken accounts, fetch more
- pass
diff --git a/tests/managers/thl/test_profiling/test_question.py b/tests/managers/thl/test_profiling/test_question.py
index 998466e..e4afb87 100644
--- a/tests/managers/thl/test_profiling/test_question.py
+++ b/tests/managers/thl/test_profiling/test_question.py
@@ -1,12 +1,20 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
-from generalresearch.managers.thl.profiling.question import QuestionManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.question import QuestionManager
class TestQuestionManager:
- def test_get_multi_upk(self, question_manager: QuestionManager, upk_data):
+ def test_get_multi_upk(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
qs = question_manager.get_multi_upk(
question_ids=[
"8a22de34f985476aac85e15547100db8",
@@ -17,13 +25,21 @@ class TestQuestionManager:
)
assert len(qs) == 3
- def test_get_questions_ranked(self, question_manager: QuestionManager, upk_data):
+ def test_get_questions_ranked(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
qs = question_manager.get_questions_ranked(country_iso="mx", language_iso="spa")
assert len(qs) >= 40
assert qs[0].importance.task_score > qs[40].importance.task_score
assert all(q.country_iso == "mx" and q.language_iso == "spa" for q in qs)
- def test_lookup_by_property(self, question_manager: QuestionManager, upk_data):
+ def test_lookup_by_property(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
q = question_manager.lookup_by_property(
property_code="i:industry", country_iso="us", language_iso="eng"
)
@@ -38,7 +54,11 @@ class TestQuestionManager:
)
assert q.explanation_template
- def test_filter_by_property(self, question_manager: QuestionManager, upk_data):
+ def test_filter_by_property(
+ self, question_manager: QuestionManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
lookup = [
("i:industry", "us", "eng"),
("i:industry", "mx", "eng"),
diff --git a/tests/managers/thl/test_profiling/test_schema.py b/tests/managers/thl/test_profiling/test_schema.py
index ae61527..feab902 100644
--- a/tests/managers/thl/test_profiling/test_schema.py
+++ b/tests/managers/thl/test_profiling/test_schema.py
@@ -1,9 +1,21 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
from generalresearch.models.thl.profiling.upk_property import PropertyType
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.schema import (
+ UpkSchemaManager,
+ )
+
class TestUpkSchemaManager:
- def test_get_props_info(self, upk_schema_manager, upk_data):
+ def test_get_props_info(
+ self, upk_schema_manager: UpkSchemaManager, upk_data: Callable[..., None]
+ ):
+ upk_data()
+
props = upk_schema_manager.get_props_info()
assert (
len(props) == 16955
@@ -35,10 +47,10 @@ class TestUpkSchemaManager:
assert age.prop_type == PropertyType.UPK_NUMERICAL
assert age.gold_standard
- cars = [
+ cars = next(
x
for x in props
if x.country_iso == "us" and x.property_label == "household_auto_type"
- ][0]
+ )
assert not cars.gold_standard
assert cars.categories[0].label == "Autos & Vehicles"
diff --git a/tests/managers/thl/test_profiling/test_uqa.py b/tests/managers/thl/test_profiling/test_uqa.py
deleted file mode 100644
index 8b13789..0000000
--- a/tests/managers/thl/test_profiling/test_uqa.py
+++ /dev/null
@@ -1 +0,0 @@
-
diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py
index 53bb8fe..0f3140c 100644
--- a/tests/managers/thl/test_profiling/test_user_upk.py
+++ b/tests/managers/thl/test_profiling/test_user_upk.py
@@ -1,8 +1,12 @@
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
-from generalresearch.managers.thl.profiling.user_upk import UserUpkManager
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.user_upk import UserUpkManager
+ from generalresearch.models.thl.user import User
-now = datetime.now(tz=timezone.utc)
+now = datetime.now(tz=UTC)
base = {
"country_iso": "us",
"language_iso": "eng",
@@ -21,11 +25,25 @@ for a in upk_ans_dict:
class TestUserUpkManager:
- def test_user_upk_empty(self, user_upk_manager: UserUpkManager, upk_data, user):
+ def test_user_upk_empty(
+ self,
+ user_upk_manager: UserUpkManager,
+ upk_data: Callable[..., None],
+ user: User,
+ ):
+ upk_data()
+
res = user_upk_manager.get_user_upk_mysql(user_id=user.user_id)
assert len(res) == 0
- def test_user_upk(self, user_upk_manager: UserUpkManager, upk_data, user):
+ def test_user_upk(
+ self,
+ user_upk_manager: UserUpkManager,
+ upk_data: Callable[..., None],
+ user: User,
+ ):
+ upk_data()
+
for x in upk_ans_dict:
x["user_id"] = user.user_id
user_upk = user_upk_manager.populate_user_upk_from_dict(upk_ans_dict)
diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py
index 6bedc2b..60edcb9 100644
--- a/tests/managers/thl/test_session_manager.py
+++ b/tests/managers/thl/test_session_manager.py
@@ -1,28 +1,42 @@
-from datetime import timedelta
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import datetime, timedelta
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
from faker import Faker
-from generalresearch.models import DeviceType
+from generalresearch.models.definitions import DeviceType
from generalresearch.models.legacy.bucket import Bucket
from generalresearch.models.thl.definitions import (
+ SessionStatusCode2,
Status,
StatusCode1,
- SessionStatusCode2,
)
-from test_utils.models.conftest import user
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.gr.team import Team
+ from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
fake = Faker()
class TestSessionManager:
- def test_create_session(self, session_manager, user, utc_hour_ago):
+ def test_create_session(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
bucket = Bucket(
loi_min=timedelta(seconds=60),
loi_max=timedelta(seconds=120),
- user_payout_min=Decimal("1"),
- user_payout_max=Decimal("2"),
+ user_payout_min=Decimal(1),
+ user_payout_max=Decimal(2),
)
s1 = session_manager.create(
@@ -40,7 +54,9 @@ class TestSessionManager:
s2 = session_manager.get_from_uuid(session_uuid=s1.uuid)
assert s1 == s2
- def test_finish_with_status(self, session_manager, user, utc_hour_ago):
+ def test_finish_with_status(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
uuid_1 = uuid4().hex
session = session_manager.create(
started=utc_hour_ago, user=user, uuid_id=uuid_1
@@ -60,7 +76,7 @@ class TestSessionManager:
class TestSessionManagerFilter:
- def test_base(self, session_manager, user, utc_now):
+ def test_base(self, session_manager: SessionManager, user: User, utc_now: datetime):
uuid_id = uuid4().hex
session_manager.create(started=utc_now, user=user, uuid_id=uuid_id)
res = session_manager.filter(limit=1)
@@ -68,7 +84,9 @@ class TestSessionManagerFilter:
assert isinstance(res, list)
assert res[0].uuid == uuid_id
- def test_user(self, session_manager, user, utc_hour_ago):
+ def test_user(
+ self, session_manager: SessionManager, user: User, utc_hour_ago: datetime
+ ):
session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex)
session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex)
@@ -76,14 +94,16 @@ class TestSessionManagerFilter:
assert len(res) == 2
def test_product(
- self, product_factory, user_factory, session_manager, user, utc_hour_ago
+ self,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
- from generalresearch.models.thl.user import User
p1 = product_factory()
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
@@ -96,42 +116,40 @@ class TestSessionManagerFilter:
def test_team(
self,
- product_factory,
- user_factory,
- team,
- session_manager,
- user,
- utc_hour_ago,
- thl_web_rr,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ gr_team: Team,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
+ thl_web_rr: PostgresConfig,
):
- p1 = product_factory(team=team)
+ p1 = product_factory(team=gr_team)
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
- team.prefetch_products(thl_pg_config=thl_web_rr)
- assert len(team.product_uuids) == 1
- res = session_manager.filter(product_uuids=team.product_uuids)
+ gr_team.prefetch_products(thl_pg_config=thl_web_rr)
+ assert len(gr_team.product_uuids) == 1
+ res = session_manager.filter(product_uuids=gr_team.product_uuids)
assert len(res) == 5
def test_business(
self,
- product_factory,
- business,
- user_factory,
- session_manager,
- user,
- utc_hour_ago,
- thl_web_rr,
+ product_factory: Callable[..., Product],
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
+ thl_web_rr: PostgresConfig,
):
- p1 = product_factory(business=business)
+ p1 = product_factory(business=gr_business)
- for n in range(5):
+ for _ in range(5):
u = user_factory(product=p1)
session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex)
- business.prefetch_products(thl_pg_config=thl_web_rr)
- assert len(business.product_uuids) == 1
- res = session_manager.filter(product_uuids=business.product_uuids)
+ gr_business.prefetch_products(thl_pg_config=thl_web_rr)
+ assert len(gr_business.product_uuids) == 1
+ res = session_manager.filter(product_uuids=gr_business.product_uuids)
assert len(res) == 5
diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py
index 58c4577..e114b70 100644
--- a/tests/managers/thl/test_survey.py
+++ b/tests/managers/thl/test_survey.py
@@ -1,29 +1,41 @@
+from __future__ import annotations
+
import uuid
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.legacy.bucket import (
- SurveyEligibilityCriterion,
- TopNPlusBucket,
DurationSummary,
PayoutSummary,
+ SurveyEligibilityCriterion,
+ TopNPlusBucket,
)
from generalresearch.models.thl.profiling.user_question_answer import (
UserQuestionAnswer,
)
from generalresearch.models.thl.survey.model import (
Survey,
- SurveyStat,
SurveyCategoryModel,
SurveyEligibilityDefinition,
+ SurveyStat,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.buyer import BuyerManager
+ from generalresearch.managers.thl.profiling.question import (
+ QuestionManager,
+ )
+ from generalresearch.managers.thl.profiling.uqa import UQAManager
+ from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager
+
@pytest.fixture(scope="session")
-def surveys_fixture():
+def surveys_fixture() -> list[Survey]:
return [
Survey(source=Source.TESTING, survey_id="a", buyer_code="buyer1"),
Survey(source=Source.TESTING, survey_id="b", buyer_code="buyer2"),
@@ -73,11 +85,13 @@ class TestSurvey:
def test(
self,
- delete_buyers_surveys,
- buyer_manager,
- survey_manager,
- surveys_fixture,
+ delete_buyers_surveys: Callable[..., None],
+ buyer_manager: BuyerManager,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
+ delete_buyers_surveys()
+
survey_manager.create_or_update(surveys_fixture)
survey_ids = {s.survey_id for s in surveys_fixture}
res = survey_manager.filter_by_natural_key(
@@ -98,7 +112,7 @@ class TestSurvey:
assert res2[0] == res[0]
assert len(res2) == len(surveys2)
- def test_category(self, survey_manager):
+ def test_category(self, survey_manager: SurveyManager):
survey1 = Survey(id=562289, survey_id="a", source=Source.TESTING)
survey2 = Survey(id=562290, survey_id="a", source=Source.TESTING)
categories = list(survey_manager.category_manager.categories.values())
@@ -110,8 +124,14 @@ class TestSurvey:
survey_manager.update_surveys_categories(surveys)
def test_survey_eligibility(
- self, survey_manager, upk_data, question_manager, uqa_manager
+ self,
+ survey_manager: SurveyManager,
+ upk_data: Callable[..., None],
+ question_manager: QuestionManager,
+ uqa_manager: UQAManager,
):
+ upk_data()
+
bucket = TopNPlusBucket(
id="c82cf98c578a43218334544ab376b00e",
contents=[],
@@ -161,9 +181,10 @@ class TestSurvey:
calc_answers={"i:adhoc_13126": ("3", "4")},
),
]
- uqad = dict()
+ uqad = {}
for uqa in uqas:
- for k, v in uqa.calc_answers.items():
+ assert uqa.calc_answers
+ for k in uqa.calc_answers:
if k in qualifying_questions:
uqad[k] = uqa
uqad[uqa.property_code] = uqa
@@ -174,9 +195,9 @@ class TestSurvey:
qs = sorted(qs, key=lambda x: x.importance.task_count if x.importance else 0)
# qd = {q.id: q for q in qs}
- q = [x for x in qs if x.ext_question_id == "i:adhoc_13126"][0]
+ q = next(x for x in qs if x.ext_question_id == "i:adhoc_13126")
q.explanation_template = "You have been diagnosed with: {answer}."
- q = [x for x in qs if x.ext_question_id == "gr:gender"][0]
+ q = next(x for x in qs if x.ext_question_id == "gr:gender")
q.explanation_template = "Your gender is {answer}."
ecs = []
@@ -205,10 +226,9 @@ class TestSurvey:
class TestSurveyStat:
def test(
self,
- delete_buyers_surveys,
surveystat_manager,
- survey_manager,
- surveys_fixture,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
survey_manager.create_or_update(surveys_fixture)
ss = [ssa, ssb]
@@ -234,7 +254,7 @@ class TestSurveyStat:
):
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(20_000):
+ for _ in range(20_000):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -251,14 +271,14 @@ class TestSurveyStat:
survey_stats.append(ss)
print(len(survey_stats))
print(survey_stats[12].natural_key, survey_stats[2000].natural_key)
- print(f"----a-----: {datetime.now().isoformat()}")
+ print(f"----a-----: {datetime.now(tz=UTC).isoformat()}")
res = surveystat_manager.update_or_create(survey_stats)
- print(f"----b-----: {datetime.now().isoformat()}")
+ print(f"----b-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res) == 20_000
return
# 1,000 of the 20,000 are "new"
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
for s in ss[:1000]:
s.survey__survey_id = "b"
s.updated_at = now
@@ -269,22 +289,24 @@ class TestSurveyStat:
s.conv_beta = 20
s.updated_at = now
# and 1,000 don't change
- print(f"----c-----: {datetime.now().isoformat()}")
+ print(f"----c-----: {datetime.now(tz=UTC).isoformat()}")
res2 = surveystat_manager.update_or_create(ss)
- print(f"----d-----: {datetime.now().isoformat()}")
+ print(f"----d-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res2) == 20_000
def test_ymsp(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
source = Source.TESTING
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(100):
+ for _ in range(100):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -298,14 +320,14 @@ class TestSurveyStat:
source=source, surveys=surveys, survey_stats=survey_stats
)
# UPDATE -------
- since = datetime.now(tz=timezone.utc)
+ since = datetime.now(tz=UTC)
print(f"{since=}")
# 10 survey disappear
surveys = surveys[10:]
# and 2 new ones are created
- for idx in range(2):
+ for _ in range(2):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -329,11 +351,13 @@ class TestSurveyStat:
def test_filter(
self,
- delete_buyers_surveys,
- surveys_fixture,
- survey_manager,
- surveystat_manager,
+ delete_buyers_surveys: Callable[..., None],
+ surveys_fixture: list[Survey],
+ survey_manager: SurveyManager,
+ surveystat_manager: SurveyStatManager,
):
+ delete_buyers_surveys()
+
surveys = []
survey = surveys_fixture[0].model_copy()
survey.source = Source.TESTING
diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py
index 4c7dc08..04f69d2 100644
--- a/tests/managers/thl/test_survey_penalty.py
+++ b/tests/managers/thl/test_survey_penalty.py
@@ -1,14 +1,19 @@
+from __future__ import annotations
+
import uuid
+from typing import TYPE_CHECKING
import pytest
-from cachetools.keys import _HashedTuple
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.survey.penalty import (
BPSurveyPenalty,
TeamSurveyPenalty,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager
+
@pytest.fixture
def product_uuid() -> str:
@@ -23,7 +28,9 @@ def team_uuid() -> str:
@pytest.fixture
-def penalties(product_uuid, team_uuid):
+def penalties(
+ product_uuid: str, team_uuid: str
+) -> list[BPSurveyPenalty | TeamSurveyPenalty]:
return [
BPSurveyPenalty(
source=Source.TESTING, survey_id="a", penalty=0.1, product_id=product_uuid
@@ -49,7 +56,13 @@ def penalties(product_uuid, team_uuid):
class TestSurveyPenalty:
- def test(self, surveypenalty_manager, penalties, product_uuid, team_uuid):
+ def test(
+ self,
+ surveypenalty_manager: SurveyPenaltyManager,
+ penalties: list[BPSurveyPenalty | TeamSurveyPenalty],
+ product_uuid: str,
+ team_uuid: str,
+ ):
surveypenalty_manager.set_penalties(penalties)
res = surveypenalty_manager.get_penalties_for(
@@ -89,10 +102,8 @@ class TestSurveyPenalty:
)
assert res == {"t:a": 0.1, "t:b": 0.2, "u:b": 0.1}
assert surveypenalty_manager.cache.currsize == 1
- cached_key = tuple(list(list(surveypenalty_manager.cache.keys())[0])[1:])
- assert cached_key == tuple(
- ["product_id", product_uuid, "team_id", team_id_random]
- )
+ cached_key = tuple(list(next(iter(surveypenalty_manager.cache.keys())))[1:])
+ assert cached_key == ("product_id", product_uuid, "team_id", team_id_random)
# Both don't exist, return nothing
res = surveypenalty_manager.get_penalties_for(
diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py
index 839bbe1..1e741fd 100644
--- a/tests/managers/thl/test_task_adjustment.py
+++ b/tests/managers/thl/test_task_adjustment.py
@@ -1,27 +1,43 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
+from decimal import Decimal
from random import randint
+from typing import TYPE_CHECKING
import pytest
-from datetime import datetime, timezone, timedelta
-from decimal import Decimal
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
Status,
StatusCode1,
WallAdjustedStatus,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.task_adjustment import (
+ TaskAdjustmentManager,
+ )
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
+
@pytest.fixture()
-def session_complete(session_with_tx_factory, user):
+def session_complete(session_with_tx_factory: Callable[..., Session], user: User):
return session_with_tx_factory(
user=user, final_status=Status.COMPLETE, wall_req_cpi=Decimal("1.23")
)
@pytest.fixture()
-def session_complete_with_wallet(session_with_tx_factory, user_with_wallet):
+def session_complete_with_wallet(
+ session_with_tx_factory: Callable[..., None], user_with_wallet: User
+):
return session_with_tx_factory(
user=user_with_wallet,
final_status=Status.COMPLETE,
@@ -30,17 +46,21 @@ def session_complete_with_wallet(session_with_tx_factory, user_with_wallet):
@pytest.fixture()
-def session_fail(user, session_manager, wall_manager):
- session = session_manager.create_dummy(
- started=datetime.now(timezone.utc), user=user
- )
- wall1 = wall_manager.create_dummy(
- session_id=session.id,
- user_id=user.user_id,
+def session_fail(
+ user: User,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
+) -> Session:
+ session = session_manager.create(started=datetime.now(UTC), user=user)
+ wall1 = wall_factory(
+ session=session,
+ user=user,
source=Source.DYNATA,
req_survey_id="72723",
req_cpi=Decimal("3.22"),
- started=datetime.now(timezone.utc),
+ started=datetime.now(UTC),
)
wall_manager.finish(
wall=wall1,
@@ -48,46 +68,47 @@ def session_fail(user, session_manager, wall_manager):
status_code_1=StatusCode1.PS_FAIL,
finished=wall1.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
)
- session.wall_events.append(wall1)
return session
class TestHandleRecons:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
def test_complete_to_recon(
self,
- session_complete,
- thl_lm,
- task_adjustment_manager,
- wall_manager,
- session_manager,
+ session_complete: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
caplog,
):
print(wall_manager.pg_config.dsn)
mid = session_complete.uuid
wall_uuid = session_complete.wall_events[-1].uuid
s = session_complete
- ledger_manager = thl_lm
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- current_amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
- assert (
- current_amount == 123
- ), "this is the amount of revenue from this task complete"
+ assert current_amount == 123, (
+ "this is the amount of revenue from this task complete"
+ )
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
s.user.product
)
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 117, "this is the amount paid to the BP"
# Do the work here !! ----v
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
@@ -95,21 +116,21 @@ class TestHandleRecons:
len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
- commission_account = ledger_manager.get_account_or_create_bp_commission(
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
s.user.product
)
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
# Now, say we get the exact same *adjust to incomplete* msg again. It should do nothing!
- adjusted_timestamp = datetime.now(tz=timezone.utc)
+ adjusted_timestamp = datetime.now(tz=UTC)
wall = wall_manager.get_from_uuid(wall_uuid=wall_uuid)
with pytest.raises(match=" is already "):
wall_manager.adjust_status(
@@ -122,7 +143,7 @@ class TestHandleRecons:
session = session_manager.get_from_id(wall.session_id)
user = session.user
with caplog.at_level(logging.INFO):
- ledger_manager.create_tx_task_adjustment(
+ thl_ledger_manager.create_tx_task_adjustment(
wall, user=user, created=adjusted_timestamp
)
assert "No transactions needed" in caplog.text
@@ -135,212 +156,219 @@ class TestHandleRecons:
assert "is already f" in caplog.text or "is already Status.FAIL" in caplog.text
with caplog.at_level(logging.INFO, logger="LedgerManager"):
- ledger_manager.create_tx_bp_adjustment(session, created=adjusted_timestamp)
+ thl_ledger_manager.create_tx_bp_adjustment(
+ session, created=adjusted_timestamp
+ )
assert "No transactions needed" in caplog.text
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
# And if we get an adj to fail, and handle it, it should do nothing at all
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
assert (
len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0, "after recon, it should be zeroed"
- def test_fail_to_complete(self, session_fail, thl_lm, task_adjustment_manager):
- s = session_fail
+ def test_fail_to_complete(
+ self,
+ session_fail: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
+ ):
mid = session_fail.uuid
wall_uuid = session_fail.wall_events[-1].uuid
- ledger_manager = thl_lm
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- current_amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", mid
)
- assert (
- current_amount == 0
- ), "this is the amount of revenue from this task complete"
+ assert current_amount == 0, (
+ "this is the amount of revenue from this task complete"
+ )
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_fail.user.product
)
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0, "this is the amount paid to the BP"
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE,
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 322, "after recon, we should be paid"
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 306, "this is the amount paid to the BP"
# Now reverse it back to fail
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 0
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_fail.user.product
)
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
def test_complete_already_complete(
- self, session_complete, thl_lm, task_adjustment_manager
+ self,
+ session_complete: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_complete
mid = session_complete.uuid
wall_uuid = session_complete.wall_events[-1].uuid
- ledger_manager = thl_lm
for _ in range(4):
# just run it 4 times to make sure nothing happens 4 times
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE,
)
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_complete.user.product
)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_complete.user.product
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert current_amount == 123
- assert ledger_manager.get_account_balance(commission_account) == 6
+ assert thl_ledger_manager.get_account_balance(commission_account) == 6
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 117
def test_incomplete_already_incomplete(
- self, session_fail, thl_lm, task_adjustment_manager
+ self,
+ session_fail: Session,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_fail
mid = session_fail.uuid
wall_uuid = session_fail.wall_events[-1].uuid
- ledger_manager = thl_lm
for _ in range(4):
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_fail.user.product
)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_fail.user.product
)
- current_amount = ledger_manager.get_account_filtered_balance(
+ current_amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", mid
)
assert current_amount == 0
- assert ledger_manager.get_account_balance(commission_account) == 0
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0
- current_bp_payout = ledger_manager.get_account_filtered_balance(
+ current_bp_payout = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert current_bp_payout == 0
def test_complete_to_recon_user_wallet(
self,
- session_complete_with_wallet,
- user_with_wallet,
- thl_lm,
- task_adjustment_manager,
+ session_complete_with_wallet: Session,
+ # user_with_wallet: User,
+ thl_ledger_manager: ThlLedgerManager,
+ task_adjustment_manager: TaskAdjustmentManager,
):
- s = session_complete_with_wallet
- mid = s.uuid
- wall_uuid = s.wall_events[-1].uuid
- ledger_manager = thl_lm
+ mid = session_complete_with_wallet.uuid
+ wall_uuid = session_complete_with_wallet.wall_events[-1].uuid
- revenue_account = ledger_manager.get_account_task_complete_revenue()
- amount = ledger_manager.get_account_filtered_balance(
+ revenue_account = thl_ledger_manager.get_account_task_complete_revenue()
+ amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert amount == 123, "this is the amount of revenue from this task complete"
- bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet(
- s.user.product
+ bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ session_complete_with_wallet.user.product
)
- user_wallet_account = ledger_manager.get_account_or_create_user_wallet(s.user)
- commission_account = ledger_manager.get_account_or_create_bp_commission(
- s.user.product
+ user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet(
+ session_complete_with_wallet.user
)
- amount = ledger_manager.get_account_filtered_balance(
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_complete_with_wallet.user.product
+ )
+ amount = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert amount == 70, "this is the amount paid to the BP"
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
user_wallet_account, "thl_session", mid
)
assert amount == 47, "this is the amount paid to the user"
- assert (
- ledger_manager.get_account_balance(commission_account) == 6
- ), "earned commission"
+ assert thl_ledger_manager.get_account_balance(commission_account) == 6, (
+ "earned commission"
+ )
task_adjustment_manager.handle_single_recon(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
wall_uuid=wall_uuid,
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
)
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
revenue_account, "thl_wall", wall_uuid
)
assert amount == 0
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
bp_wallet_account, "thl_session", mid
)
assert amount == 0
- amount = ledger_manager.get_account_filtered_balance(
+ amount = thl_ledger_manager.get_account_filtered_balance(
user_wallet_account, "thl_session", mid
)
assert amount == 0
- assert (
- ledger_manager.get_account_balance(commission_account) == 0
- ), "earned commission"
+ assert thl_ledger_manager.get_account_balance(commission_account) == 0, (
+ "earned commission"
+ )
diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py
index 55c89c0..11edc99 100644
--- a/tests/managers/thl/test_task_status.py
+++ b/tests/managers/thl/test_task_status.py
@@ -1,47 +1,61 @@
-import pytest
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
+from typing import TYPE_CHECKING
+
+import pytest
-from generalresearch.managers.thl.session import SessionManager
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
Status,
- WallAdjustedStatus,
StatusCode1,
+ WallAdjustedStatus,
)
from generalresearch.models.thl.product import (
PayoutConfig,
- UserWalletConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ UserWalletConfig,
)
-from generalresearch.models.thl.session import Session, WallOut
+from generalresearch.models.thl.session import WallOut
from generalresearch.models.thl.task_status import TaskStatusResponse
-from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
-start1 = datetime(2023, 2, 1, tzinfo=timezone.utc)
+start1 = datetime(2023, 2, 1, tzinfo=UTC)
finish1 = start1 + timedelta(minutes=5)
recon1 = start1 + timedelta(days=20)
-start2 = datetime(2023, 2, 2, tzinfo=timezone.utc)
+start2 = datetime(2023, 2, 2, tzinfo=UTC)
finish2 = start2 + timedelta(minutes=5)
-start3 = datetime(2023, 2, 3, tzinfo=timezone.utc)
+start3 = datetime(2023, 2, 3, tzinfo=UTC)
finish3 = start3 + timedelta(minutes=5)
-@pytest.fixture(scope="session")
-def bp1(product_manager):
+@pytest.fixture()
+def bp1(
+ product_factory: Callable[..., Product], product_manager: ProductManager
+) -> Product:
# user wallet disabled, payout xform NULL
- return product_manager.create_dummy(
+ return product_factory(
user_wallet_config=UserWalletConfig(enabled=False),
payout_config=PayoutConfig(),
)
-@pytest.fixture(scope="session")
-def bp2(product_manager):
+@pytest.fixture()
+def bp2(
+ product_factory: Callable[..., Product], product_manager: ProductManager
+) -> Product:
# user wallet disabled, payout xform 40%
- return product_manager.create_dummy(
+ return product_factory(
user_wallet_config=UserWalletConfig(enabled=False),
payout_config=PayoutConfig(
payout_transformation=PayoutTransformation(
@@ -52,10 +66,12 @@ def bp2(product_manager):
)
-@pytest.fixture(scope="session")
-def bp3(product_manager):
+@pytest.fixture()
+def bp3(
+ product_factory: Callable[..., Product], product_manager: ProductManager
+) -> Product:
# user wallet enabled, payout xform 50%
- return product_manager.create_dummy(
+ return product_factory(
user_wallet_config=UserWalletConfig(enabled=True),
payout_config=PayoutConfig(
payout_transformation=PayoutTransformation(
@@ -70,9 +86,9 @@ class TestTaskStatus:
def test_task_status_complete_1(
self,
- bp1,
- user_factory,
- finished_session_factory,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
session_manager: SessionManager,
):
# User Payout xform NULL
@@ -130,7 +146,11 @@ class TestTaskStatus:
assert tsr == expected_tsr
def test_task_status_complete_2(
- self, bp2, user_factory, finished_session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%
user2: User = user_factory(product=bp2)
@@ -197,7 +217,11 @@ class TestTaskStatus:
assert tsr == expected_tsr
def test_task_status_complete_3(
- self, bp3, user_factory, finished_session_factory, session_manager
+ self,
+ bp3: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# Wallet enabled User Payout xform 50% (the response is identical
# to the user wallet disabled w same xform)
@@ -227,12 +251,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s3.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_fail(
- self, bp1, user_factory, finished_session_factory, session_manager
+ self,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: user payout is None always
user1: User = user_factory(product=bp1)
@@ -263,12 +292,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s1.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_fail_xform(
- self, bp2, user_factory, finished_session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%: user_payout is 0 (not None)
@@ -298,12 +332,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_abandon(
- self, bp1, user_factory, session_factory, session_manager
+ self,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: all payout fields are None
user: User = user_factory(product=bp1)
@@ -332,12 +371,17 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_abandon_xform(
- self, bp2, user_factory, session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%: all payout fields are None (same as when payout xform is null)
user: User = user_factory(product=bp2)
@@ -369,17 +413,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_fail(
self,
- bp1,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# Complete -> Fail
# User Payout xform NULL: adjusted_user_* and user_* is still all None
@@ -418,17 +463,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_fail_xform(
self,
- bp2,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# Complete -> Fail
# User Payout xform 40%: adjusted_user_payout is 0 (not null)
@@ -470,17 +516,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_abandon(
self,
- bp1,
- user_factory,
- session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -524,17 +571,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_abandon_xform(
self,
- bp2,
- user_factory,
- session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -581,17 +629,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_fail(
self,
- bp1,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -635,17 +684,18 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
def test_task_status_adj_complete_from_fail_xform(
self,
- bp2,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -691,6 +741,7 @@ class TestTaskStatus:
}
)
tsr = session_manager.get_task_status_response(s.uuid)
+ assert isinstance(tsr, TaskStatusResponse)
# Not bothering with wall events ...
tsr.wall_events = None
assert tsr == expected_tsr
diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py
index 0d7ffef..5d12052 100644
--- a/tests/managers/thl/test_user_manager/test_base.py
+++ b/tests/managers/thl/test_user_manager/test_base.py
@@ -1,23 +1,37 @@
import logging
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.managers.thl.user_manager import (
- UserCreateNotAllowedError,
get_bp_user_create_limit_hourly,
)
+from generalresearch.managers.thl.user_manager.exceptions import (
+ UserCreateNotAllowedError,
+)
+from generalresearch.managers.thl.user_manager.mysql_user_manager import (
+ MysqlUserManager,
+)
from generalresearch.managers.thl.user_manager.rate_limit import (
RateLimitItemPerHourConstantKey,
+ UserManagerLimiter,
)
-from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
-)
-from generalresearch.models.thl.product import Product, UserCreateConfig
+from generalresearch.models.thl.product import UserCreateConfig
from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.managers.thl.user_manager.user_manager import (
+ UserManager,
+ )
+ from generalresearch.managers.thl.userhealth import AuditLogManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
+
logger = logging.getLogger()
@@ -83,10 +97,11 @@ class TestUserManager:
class TestBlockUserManager:
- def test_block_user(self, product, user_manager: UserManager):
+ def test_block_user(self, product: Product, user_manager: UserManager):
product_user_id = f"user-{uuid4().hex[:10]}"
# mysql_user_manager to skip user creation limit check
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
@@ -109,16 +124,19 @@ class TestBlockUserManager:
user = user_manager.get_user(user_id=user.user_id)
assert user.blocked
- def test_block_user_whitelist(self, product, user_manager, thl_web_rw):
+ def test_block_user_whitelist(
+ self, product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig
+ ):
product_user_id = f"user-{uuid4().hex[:10]}"
# mysql_user_manager to skip user creation limit check
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
assert not user.blocked
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# Adds user to whitelist
thl_web_rw.execute_write(
"""
@@ -135,8 +153,13 @@ class TestBlockUserManager:
class TestCreateUserManager:
- def test_create_user(self, product_manager, thl_web_rw, user_manager):
- product: Product = product_manager.create_dummy(
+ def test_create_user(
+ self,
+ product_factory: Callable[..., Product],
+ thl_web_rw: PostgresConfig,
+ user_manager: UserManager,
+ ):
+ product: Product = product_factory(
user_create_config=UserCreateConfig(
min_hourly_create_limit=10, max_hourly_create_limit=69
),
@@ -144,6 +167,7 @@ class TestCreateUserManager:
product_user_id = f"user-{uuid4().hex[:10]}"
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
user: User = user_manager.mysql_user_manager.create_user(
product_id=product.id, product_user_id=product_user_id
)
@@ -156,7 +180,7 @@ class TestCreateUserManager:
# make sure thl_user row is created
res_thl_user = thl_web_rw.execute_sql_query(
- query=f"""
+ query="""
SELECT *
FROM thl_user AS u
WHERE u.id = %s
@@ -172,8 +196,13 @@ class TestCreateUserManager:
assert u2.user_id == user.user_id
assert u2.uuid == user.uuid
- def test_create_user_integrity_error(self, product_manager, user_manager, caplog):
- product: Product = product_manager.create_dummy(
+ def test_create_user_integrity_error(
+ self,
+ user_manager: UserManager,
+ product_factory: Callable[..., Product],
+ caplog,
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -185,6 +214,7 @@ class TestCreateUserManager:
product_user_id = f"user-{uuid4().hex[:10]}"
rand_msg = f"log-{uuid4().hex}"
+ assert isinstance(user_manager.mysql_user_manager, MysqlUserManager)
with caplog.at_level(logging.INFO):
logger.info(rand_msg)
user1 = user_manager.mysql_user_manager.create_user(
@@ -213,9 +243,14 @@ class TestCreateUserManager:
assert user1 == user2
- def test_raise_allow_user_create(self, product_manager, user_manager):
+ def test_raise_allow_user_create(
+ self,
+ product_manager: ProductManager,
+ user_manager: UserManager,
+ product_factory: Callable[..., Product],
+ ):
rand_num = randint(25, 200)
- product: Product = product_manager.create_dummy(
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -247,10 +282,11 @@ class TestCreateUserManager:
assert key == f"LIMITER/thl-grpc/allow_user_create/{instance.id}"
# make sure we clear the key or subsequent tests will fail
+ assert isinstance(user_manager.user_manager_limiter, UserManagerLimiter)
user_manager.user_manager_limiter.storage.clear(key=key)
n = 0
- with pytest.raises(expected_exception=UserCreateNotAllowedError) as cm:
+ with pytest.raises(expected_exception=UserCreateNotAllowedError):
for n, _ in enumerate(range(rl_value + 5)):
user_manager.user_manager_limiter.raise_allow_user_create(
product=product
@@ -260,14 +296,16 @@ class TestCreateUserManager:
class TestUserManagerMethods:
- def test_audit_log(self, user_manager, user, audit_log_manager):
+ def test_audit_log(
+ self, user_manager: UserManager, user: User, audit_log_manager: AuditLogManager
+ ):
from generalresearch.models.thl.userhealth import AuditLog
res = audit_log_manager.filter_by_user_id(user_id=user.user_id)
assert len(res) == 0
msg = uuid4().hex
- user_manager.audit_log(user=user, level=30, event_type=msg)
+ user_manager.audit_log(audit_log_manager, user=user, level=30, event_type=msg)
res = audit_log_manager.filter_by_user_id(user_id=user.user_id)
assert len(res) == 1
diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py
index 0313bbf..ed7d458 100644
--- a/tests/managers/thl/test_user_manager/test_mysql.py
+++ b/tests/managers/thl/test_user_manager/test_mysql.py
@@ -1,25 +1,28 @@
-from test_utils.models.conftest import user, user_manager
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.user_manager.mysql_user_manager import (
+ MysqlUserManager,
+ )
+ from generalresearch.models.thl.user import User
class TestUserManagerMysqlNew:
- def test_get_notset(self, user_manager):
- assert (
- user_manager.mysql_user_manager.get_user_from_mysql(user_id=-3105) is None
- )
+ def test_get_notset(self, mysql_user_manager: MysqlUserManager):
+ assert mysql_user_manager.get_user_from_mysql(user_id=-3105) is None
- def test_get_user_id(self, user, user_manager):
- assert (
- user_manager.mysql_user_manager.get_user_from_mysql(user_id=user.user_id)
- == user
- )
+ def test_get_user_id(self, user: User, mysql_user_manager: MysqlUserManager):
+ assert mysql_user_manager.get_user_from_mysql(user_id=user.user_id) == user
- def test_get_uuid(self, user, user_manager):
- u = user_manager.mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid)
+ def test_get_uuid(self, user: User, mysql_user_manager: MysqlUserManager):
+ u = mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid)
assert u == user
- def test_get_ubp(self, user, user_manager):
- u = user_manager.mysql_user_manager.get_user_from_mysql(
+ def test_get_ubp(self, user: User, mysql_user_manager: MysqlUserManager):
+ u = mysql_user_manager.get_user_from_mysql(
product_id=user.product_id, product_user_id=user.product_user_id
)
assert u == user
diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py
index a69519e..f6b59c9 100644
--- a/tests/managers/thl/test_user_manager/test_redis.py
+++ b/tests/managers/thl/test_user_manager/test_redis.py
@@ -1,29 +1,41 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.base import Permission
+from generalresearch.managers.thl.user_manager.redis_user_manager import (
+ RedisUserManager,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
class TestUserManagerRedis:
- def test_get_notset(self, user_manager, user):
- user_manager.clear_user_inmemory_cache(user=user)
- assert user_manager.redis_user_manager.get_user(user_id=user.user_id) is None
+ def test_get_notset(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.clear_user(user=user)
+ assert redis_user_manager.get_user(user_id=user.user_id) is None
- def test_get_user_id(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_user_id(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
- assert user_manager.redis_user_manager.get_user(user_id=user.user_id) == user
+ assert redis_user_manager.get_user(user_id=user.user_id) == user
- def test_get_uuid(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_uuid(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
- assert user_manager.redis_user_manager.get_user(user_uuid=user.uuid) == user
+ assert redis_user_manager.get_user(user_uuid=user.uuid) == user
- def test_get_ubp(self, user_manager, user):
- user_manager.redis_user_manager.set_user(user=user)
+ def test_get_ubp(self, redis_user_manager: RedisUserManager, user: User):
+ redis_user_manager.set_user(user=user)
assert (
- user_manager.redis_user_manager.get_user(
+ redis_user_manager.get_user(
product_id=user.product_id, product_user_id=user.product_user_id
)
== user
@@ -34,7 +46,13 @@ class TestUserManagerRedis:
# I mean, the sets are implicitly tested by the get tests above. no point
pass
- def test_get_with_cache_prefix(self, settings, user, thl_web_rw, thl_web_rr):
+ def test_get_with_cache_prefix(
+ self,
+ settings: GRLBaseSettings,
+ user: User,
+ thl_web_rw: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ ):
"""
Confirm the prefix functionality is working; we do this so it
is easier to migrate between any potentially breaking versions
@@ -69,9 +87,11 @@ class TestUserManagerRedis:
product_id=user.product_id, product_user_id=user.product_user_id
)
+ assert isinstance(um1.redis_user_manager, RedisUserManager)
res1 = um1.redis_user_manager.client.get(f"user-lookup:user_id:{user.user_id}")
assert res1 is not None
+ assert isinstance(um2.redis_user_manager, RedisUserManager)
res2 = um2.redis_user_manager.client.get(
f"user-lookup-v2:user_id:{user.user_id}"
)
diff --git a/tests/managers/thl/test_user_manager/test_user_fetch.py b/tests/managers/thl/test_user_manager/test_user_fetch.py
index a4b3d57..9a279ed 100644
--- a/tests/managers/thl/test_user_manager/test_user_fetch.py
+++ b/tests/managers/thl/test_user_manager/test_user_fetch.py
@@ -1,14 +1,25 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models.thl.user import User
-from test_utils.models.conftest import product, user_manager, user_factory
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
class TestUserManagerFetch:
- def test_fetch(self, user_factory, product, user_manager):
+ def test_fetch(
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_manager: UserManager,
+ ):
user1: User = user_factory(product=product)
user2: User = user_factory(product=product)
res = user_manager.fetch_by_bpuids(
@@ -30,7 +41,7 @@ class TestUserManagerFetch:
res = user_manager.fetch(user_uuids=[uuid4().hex])
assert len(res) == 0
- def test_fetch_invalid(self, user_manager):
+ def test_fetch_invalid(self, user_manager: UserManager):
with pytest.raises(AssertionError) as e:
user_manager.fetch(user_uuids=[], user_ids=None)
assert "Must pass ONE of user_ids, user_uuids" in str(e.value)
diff --git a/tests/managers/thl/test_user_manager/test_user_metadata.py b/tests/managers/thl/test_user_manager/test_user_metadata.py
index 91dc16a..eb6a272 100644
--- a/tests/managers/thl/test_user_manager/test_user_metadata.py
+++ b/tests/managers/thl/test_user_manager/test_user_metadata.py
@@ -1,20 +1,38 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.models.thl.user_profile import UserMetadata
-from test_utils.models.conftest import user, user_manager, user_factory
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.user_manager.user_metadata_manager import (
+ UserMetadataManager,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
class TestUserMetadataManager:
- def test_get_notset(self, user, user_manager, user_metadata_manager):
+ def test_get_notset(
+ self,
+ user: User,
+ user_metadata_manager: UserMetadataManager,
+ ):
# The row in the db won't exist. It just returns the default obj with everything None (except for the user_id)
um1 = user_metadata_manager.get(user_id=user.user_id)
assert um1 == UserMetadata(user_id=user.user_id)
- def test_create(self, user_factory, product, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_create(
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_metadata_manager: UserMetadataManager,
+ ):
u1: User = user_factory(product=product)
@@ -27,8 +45,12 @@ class TestUserMetadataManager:
um2 = user_metadata_manager.get(email_address=email_address)
assert um == um2
- def test_create_no_email(self, product, user_factory, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_create_no_email(
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
+ ):
u1: User = user_factory(product=product)
um = UserMetadata(user_id=u1.user_id)
@@ -38,8 +60,12 @@ class TestUserMetadataManager:
um2 = user_metadata_manager.get(user_id=u1.user_id)
assert um == um2
- def test_update(self, product, user_factory, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_update(
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
+ ):
u: User = user_factory(product=product)
@@ -58,8 +84,9 @@ class TestUserMetadataManager:
email_address=email_address.replace("example1", "example2"),
)
- def test_filter(self, user_factory, product, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_filter(
+ self, user_factory: Callable[..., User], product: Product, user_metadata_manager
+ ):
user1: User = user_factory(product=product)
user2: User = user_factory(product=product)
diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py
index 7728f9f..d99b2b8 100644
--- a/tests/managers/thl/test_user_streak.py
+++ b/tests/managers/thl/test_user_streak.py
@@ -1,19 +1,33 @@
+from __future__ import annotations
+
import copy
-from datetime import datetime, timezone, timedelta, date
+from collections.abc import Callable
+from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
+from typing import TYPE_CHECKING
from zoneinfo import ZoneInfo
import pytest
-from generalresearch.managers.thl.user_streak import compute_streaks_from_days
-from generalresearch.models.thl.definitions import StatusCode1, Status
+from generalresearch.managers.thl.user_streak import (
+ compute_streaks_from_days,
+)
+from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.user_streak import (
- UserStreak,
- StreakState,
- StreakPeriod,
StreakFulfillment,
+ StreakPeriod,
+ StreakState,
+ UserStreak,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.user_streak import (
+ UserStreakManager,
+ )
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
+
def test_compute_streaks_from_days():
days = [
@@ -59,7 +73,7 @@ def test_compute_streaks_from_days():
@pytest.fixture
-def broken_active_streak(user):
+def broken_active_streak(user: User) -> list[UserStreak]:
return [
UserStreak(
period=StreakPeriod.DAY,
@@ -94,8 +108,12 @@ def broken_active_streak(user):
]
-def create_session_fail(session_manager, start, user):
- session = session_manager.create_dummy(started=start, country_iso="us", user=user)
+def create_session_fail(
+ session_manager: SessionManager,
+ start: datetime,
+ user: User,
+):
+ session = session_manager.create(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
finished=start + timedelta(minutes=1),
@@ -104,8 +122,12 @@ def create_session_fail(session_manager, start, user):
)
-def create_session_complete(session_manager, start, user):
- session = session_manager.create_dummy(started=start, country_iso="us", user=user)
+def create_session_complete(
+ session_manager: SessionManager,
+ start: datetime,
+ user: User,
+):
+ session = session_manager.create(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
finished=start + timedelta(minutes=1),
@@ -115,7 +137,7 @@ def create_session_complete(session_manager, start, user):
)
-def test_user_streak_empty(user_streak_manager, user):
+def test_user_streak_empty(user_streak_manager: UserStreakManager, user: User):
streaks = user_streak_manager.get_user_streaks(
user_id=user.user_id, country_iso="us"
)
@@ -123,14 +145,19 @@ def test_user_streak_empty(user_streak_manager, user):
def test_user_streaks_active_broken(
- user_streak_manager, user, session_manager, broken_active_streak
+ user_streak_manager: UserStreakManager,
+ user: User,
+ session_manager: SessionManager,
+ broken_active_streak: list[UserStreak],
+ bare_session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
):
# Testing active streak, but broken (not today or yesterday)
- start1 = datetime(2025, 2, 12, tzinfo=timezone.utc)
+ start1 = datetime(2025, 2, 12, tzinfo=UTC)
end1 = start1 + timedelta(minutes=1)
# abandon counts as inactive
- session = session_manager.create_dummy(started=start1, country_iso="us", user=user)
+ session = bare_session_factory(started=start1, country_iso="us", user=user)
streak = user_streak_manager.get_user_streaks(user_id=user.user_id)
assert streak == []
@@ -171,12 +198,14 @@ def test_user_streaks_active_broken(
assert streaks == expected_streaks
-def test_user_streak_complete_active(user_streak_manager, user, session_manager):
+def test_user_streak_complete_active(
+ user_streak_manager: UserStreakManager, user: User, session_manager: SessionManager
+):
"""Testing active streak that is today"""
# They completed yesterday NY time. Today isn't over so streak is pending
start1 = datetime.now(tz=ZoneInfo("America/New_York")) - timedelta(days=1)
- create_session_complete(session_manager, start1.astimezone(tz=timezone.utc), user)
+ create_session_complete(session_manager, start1.astimezone(tz=UTC), user)
last_complete_day = start1.date()
expected_streak = UserStreak(
@@ -192,16 +221,16 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager)
streaks = user_streak_manager.get_user_streaks(
user_id=user.user_id, country_iso="us"
)
- streak = [
+ streak = next(
s
for s in streaks
if s.fulfillment == StreakFulfillment.COMPLETE and s.period == StreakPeriod.DAY
- ][0]
+ )
assert streak == expected_streak
# And now they complete today
start2 = datetime.now(tz=ZoneInfo("America/New_York"))
- create_session_complete(session_manager, start2.astimezone(tz=timezone.utc), user)
+ create_session_complete(session_manager, start2.astimezone(tz=UTC), user)
last_complete_day = start2.date()
expected_streak = UserStreak(
longest_streak=2,
@@ -217,9 +246,9 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager)
streaks = user_streak_manager.get_user_streaks(
user_id=user.user_id, country_iso="us"
)
- streak = [
+ streak = next(
s
for s in streaks
if s.fulfillment == StreakFulfillment.COMPLETE and s.period == StreakPeriod.DAY
- ][0]
+ )
assert streak == expected_streak
diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py
index 1cda8de..268b110 100644
--- a/tests/managers/thl/test_userhealth.py
+++ b/tests/managers/thl/test_userhealth.py
@@ -1,27 +1,43 @@
-from datetime import timezone, datetime
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
from uuid import uuid4
import faker
import pytest
from generalresearch.managers.thl.userhealth import (
+ AuditLogManager,
IPRecordManager,
UserIpHistoryManager,
)
-from generalresearch.models.thl.ipinfo import GeoIPInformation
+from generalresearch.models.thl.ipinfo import (
+ GeoIPInformation,
+)
from generalresearch.models.thl.user_iphistory import (
IPRecord,
+ UserIPHistory,
)
-from generalresearch.models.thl.userhealth import AuditLogLevel, AuditLog
+from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
+
+if TYPE_CHECKING:
+ from generalresearch.models.thl.ipinfo import (
+ IPGeoname,
+ IPInformation,
+ )
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
fake = faker.Faker()
class TestAuditLog:
- def test_init(self, thl_web_rr, audit_log_manager):
- from generalresearch.managers.thl.userhealth import AuditLogManager
-
+ def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager: AuditLogManager):
alm = AuditLogManager(pg_config=thl_web_rr)
assert isinstance(alm, AuditLogManager)
@@ -33,15 +49,16 @@ class TestAuditLog:
argnames="level",
argvalues=list(AuditLogLevel),
)
- def test_create(self, audit_log_manager, user, level):
+ def test_create(
+ self, audit_log_manager: AuditLogManager, user: User, level: AuditLogLevel
+ ):
instance = audit_log_manager.create(
user_id=user.user_id, level=level, event_type=uuid4().hex
)
assert isinstance(instance, AuditLog)
assert instance.id != 1
- def test_get_by_id(self, audit_log, audit_log_manager):
- from generalresearch.models.thl.userhealth import AuditLog
+ def test_get_by_id(self, audit_log: AuditLog, audit_log_manager: AuditLogManager):
with pytest.raises(expected_exception=Exception) as cm:
audit_log_manager.get_by_id(auditlog_id=999_999_999_999)
@@ -51,14 +68,14 @@ class TestAuditLog:
res = audit_log_manager.get_by_id(auditlog_id=audit_log.id)
assert isinstance(res, AuditLog)
assert res.id == audit_log.id
- assert res.created.tzinfo == timezone.utc
+ assert res.created.tzinfo == UTC
def test_filter_by_product(
self,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -82,7 +99,11 @@ class TestAuditLog:
assert len(res) == 1
def test_filter_by_user_id(
- self, user_factory, product, audit_log_factory, audit_log_manager
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
u1 = user_factory(product=product)
u2 = user_factory(product=product)
@@ -108,10 +129,10 @@ class TestAuditLog:
def test_filter(
self,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -142,10 +163,10 @@ class TestAuditLog:
def test_filter_count(
self,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -179,7 +200,7 @@ class TestAuditLog:
res = audit_log_manager.filter_count(
user_ids=[u1.user_id, u2.user_id, u3.user_id],
- created_after=datetime.now(tz=timezone.utc),
+ created_after=datetime.now(tz=UTC),
)
assert isinstance(res, int)
assert res == 0
@@ -205,18 +226,28 @@ class TestAuditLog:
class TestIPRecordManager:
- def test_init(self, thl_web_rr, thl_redis_config, ip_record_manager):
+ def test_init(
+ self,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
+ ip_record_manager: IPRecordManager,
+ ):
instance = IPRecordManager(pg_config=thl_web_rr, redis_config=thl_redis_config)
assert isinstance(instance, IPRecordManager)
assert isinstance(ip_record_manager, IPRecordManager)
- def test_create(self, ip_record_manager, user, ip_information):
- instance = ip_record_manager.create_dummy(
- user_id=user.user_id, ip=ip_information.ip
- )
+ def test_create(
+ self,
+ ip_record_manager: IPRecordManager,
+ user: User,
+ ip_information: IPInformation,
+ ip_record_factory: Callable[..., IPRecord],
+ ):
+ instance = ip_record_factory(user_id=user.user_id, ip=ip_information.ip)
assert isinstance(instance, IPRecord)
assert isinstance(instance.forwarded_ips, list)
+ assert isinstance(instance.forwarded_ip_records, list)
assert isinstance(instance.forwarded_ip_records[0], IPRecord)
assert isinstance(instance.forwarded_ips[0], str)
@@ -228,20 +259,22 @@ class TestIPRecordManager:
def test_prefetch_info(
self,
- ip_record_factory,
- ip_information_factory,
- ip_geoname,
- user,
- thl_web_rr,
- thl_redis_config,
+ ip_record_factory: Callable[..., IPRecord],
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ user: User,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
):
ip = fake.ipv4_public()
ip_information_factory(ip=ip, geoname=ip_geoname)
ipr: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
+ assert isinstance(ipr, IPRecord)
assert ipr.information is None
assert len(ipr.forwarded_ip_records) >= 1
+ assert isinstance(ipr.forwarded_ip_records, list)
fipr = ipr.forwarded_ip_records[0]
assert fipr.information is None
@@ -265,7 +298,12 @@ class TestIPRecordManager:
@pytest.mark.usefixtures("user_iphistory_manager_clear_cache")
class TestUserIpHistoryManager:
- def test_init(self, thl_web_rr, thl_redis_config, user_iphistory_manager):
+ def test_init(
+ self,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
+ user_iphistory_manager: UserIpHistoryManager,
+ ):
instance = UserIpHistoryManager(
pg_config=thl_web_rr, redis_config=thl_redis_config
)
@@ -274,27 +312,31 @@ class TestUserIpHistoryManager:
def test_latest_record(
self,
- user_iphistory_manager,
- user,
- ip_record_factory,
- ip_information_factory,
- ip_geoname,
+ user_iphistory_manager: UserIpHistoryManager,
+ user: User,
+ ip_record_factory: Callable[..., IPRecord],
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
ip = fake.ipv4_public()
- ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
+ ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True)
ipr1: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
ipr = user_iphistory_manager.get_user_latest_ip_record(user=user)
+ assert isinstance(ipr, IPRecord)
assert ipr.ip == ipr1.ip
assert ipr.is_anonymous
+ assert isinstance(ipr.information, GeoIPInformation)
assert ipr.information.lookup_prefix == "/32"
ip = fake.ipv6()
- ip_information_factory(ip=ip, geoname=ip_geoname)
+ ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id)
ipr2: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip)
ipr = user_iphistory_manager.get_user_latest_ip_record(user=user)
+ assert isinstance(ipr, IPRecord)
assert ipr.ip == ipr2.ip
+ assert isinstance(ipr.information, GeoIPInformation)
assert ipr.information.lookup_prefix == "/64"
assert ipr.information is not None
assert not ipr.is_anonymous
@@ -303,6 +345,8 @@ class TestUserIpHistoryManager:
assert country_iso == ip_geoname.country_iso
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert iph.ips[0].information is not None
assert iph.ips[1].information is not None
assert iph.ips[0].country_iso == country_iso
@@ -310,7 +354,12 @@ class TestUserIpHistoryManager:
assert iph.ips[0].ip == ipr1.ip
assert iph.ips[1].ip == ipr2.ip
- def test_virgin(self, user, user_iphistory_manager, ip_record_factory):
+ def test_virgin(
+ self,
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_record_factory: Callable[..., IPRecord],
+ ):
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
assert len(iph.ips) == 0
@@ -320,23 +369,27 @@ class TestUserIpHistoryManager:
def test_out_of_order(
self,
- ip_record_factory,
- user,
- user_iphistory_manager,
- ip_information_factory,
- ip_geoname,
+ ip_record_factory: Callable[..., IPRecord],
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
# Create the user-ip association BEFORE the ip even exists in the ipinfo table
ip = fake.ipv4_public()
ip_record_factory(user_id=user.user_id, ip=ip)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is None
assert not ipr.is_anonymous
- ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
+ ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is not None
@@ -344,23 +397,27 @@ class TestUserIpHistoryManager:
def test_out_of_order_ipv6(
self,
- ip_record_factory,
- user,
- user_iphistory_manager,
- ip_information_factory,
- ip_geoname,
+ ip_record_factory: Callable[..., IPRecord],
+ user: User,
+ user_iphistory_manager: UserIpHistoryManager,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
):
# Create the user-ip association BEFORE the ip even exists in the ipinfo table
ip = fake.ipv6()
ip_record_factory(user_id=user.user_id, ip=ip)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is None
assert not ipr.is_anonymous
- ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True)
+ ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True)
iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id)
+ assert isinstance(iph, UserIPHistory)
+ assert isinstance(iph.ips, list)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is not None
diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py
index ee44e23..777252f 100644
--- a/tests/managers/thl/test_wall_manager.py
+++ b/tests/managers/thl/test_wall_manager.py
@@ -1,25 +1,39 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from pydantic import PositiveInt
-from generalresearch.models import Source
-from generalresearch.models.thl.session import (
+from generalresearch.models.definitions import Source
+from generalresearch.models.thl.definitions import (
ReportValue,
Status,
StatusCode1,
)
-from test_utils.models.conftest import user, session
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.wall import WallCacheManager, WallManager
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
class TestWallManager:
@pytest.mark.parametrize("wall_count", [1, 2, 5, 10, 50, 99])
def test_get_wall_events(
- self, wall_manager, session_factory, user, wall_count, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ user: User,
+ wall_count: PositiveInt,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
s1: Session = session_factory(
user=user, wall_count=wall_count, started=utc_hour_ago
@@ -63,12 +77,15 @@ class TestWallManager:
]
def test_get_wall_events_list_input(
- self, wall_manager, session_factory, user, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ user: User,
+ utc_hour_ago: datetime,
):
- from generalresearch.models.thl.session import Session
session_ids = []
- for idx in range(10):
+ for _ in range(10):
s: Session = session_factory(user=user, wall_count=5, started=utc_hour_ago)
session_ids.append(s.id)
@@ -78,21 +95,21 @@ class TestWallManager:
assert isinstance(res, list)
assert len(res) == 50
- res1 = list(set([w.session_id for w in res]))
+ res1 = list({w.session_id for w in res})
res1.sort()
assert session_ids == res1
- def test_create_wall(self, wall_manager, session_manager, user, session):
+ def test_create_wall(self, wall_manager: WallManager, user: User, session: Session):
w = wall_manager.create(
session_id=session.id,
user_id=user.user_id,
uuid_id=uuid4().hex,
- started=datetime.now(tz=timezone.utc),
+ started=datetime.now(tz=UTC),
source=Source.DYNATA,
buyer_id="123",
req_survey_id="456",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
assert w is not None
@@ -100,7 +117,11 @@ class TestWallManager:
assert w == w2
def test_report_wall_abandon(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
):
w1 = wall_manager.create(
session_id=session.id,
@@ -110,7 +131,7 @@ class TestWallManager:
source=Source.DYNATA,
buyer_id="123",
req_survey_id="456",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
wall_manager.report(
wall=w1,
@@ -141,7 +162,12 @@ class TestWallManager:
# the status and finished get updated
def test_report_wall(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
):
w1 = wall_manager.create(
session_id=session.id,
@@ -151,7 +177,7 @@ class TestWallManager:
source=Source.DYNATA,
buyer_id="123",
req_survey_id="456",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
finish_ts = utc_hour_ago + timedelta(minutes=10)
@@ -178,11 +204,15 @@ class TestWallManager:
assert "This survey blows!" == w2.report_notes
def test_filter_wall_attempts(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
):
res = wall_manager.filter_wall_attempts(user_id=user.user_id)
assert len(res) == 0
- w1 = wall_manager.create(
+ wall_manager.create(
session_id=session.id,
user_id=user.user_id,
uuid_id=uuid4().hex,
@@ -190,11 +220,11 @@ class TestWallManager:
source=Source.DYNATA,
buyer_id="123",
req_survey_id="456",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
res = wall_manager.filter_wall_attempts(user_id=user.user_id)
assert len(res) == 1
- w2 = wall_manager.create(
+ wall_manager.create(
session_id=session.id,
user_id=user.user_id,
uuid_id=uuid4().hex,
@@ -202,7 +232,7 @@ class TestWallManager:
source=Source.DYNATA,
buyer_id="123",
req_survey_id="555",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
res = wall_manager.filter_wall_attempts(user_id=user.user_id)
assert len(res) == 2
@@ -210,21 +240,25 @@ class TestWallManager:
class TestWallCacheManager:
- def test_get_attempts_none(self, wall_cache_manager, user):
+ def test_get_attempts_none(self, wall_cache_manager: WallCacheManager, user: User):
attempts = wall_cache_manager.get_attempts(user.user_id)
assert len(attempts) == 0
def test_get_wall_events(
- self, wall_cache_manager, wall_manager, session_manager, user
+ self,
+ wall_cache_manager: WallCacheManager,
+ user: User,
+ bare_session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
):
- start1 = datetime.now(timezone.utc) - timedelta(hours=3)
- start2 = datetime.now(timezone.utc) - timedelta(hours=2)
- start3 = datetime.now(timezone.utc) - timedelta(hours=1)
+ start1 = datetime.now(UTC) - timedelta(hours=3)
+ start2 = datetime.now(UTC) - timedelta(hours=2)
+ start3 = datetime.now(UTC) - timedelta(hours=1)
- session = session_manager.create_dummy(started=start1, user=user)
- wall1 = wall_manager.create_dummy(
+ session = bare_session_factory(started=start1, user=user)
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start1,
req_cpi=Decimal("1.23"),
req_survey_id="11111",
@@ -238,9 +272,9 @@ class TestWallCacheManager:
attempts = wall_cache_manager.get_attempts(user_id=user.user_id)
assert len(attempts) == 1
- wall2 = wall_manager.create_dummy(
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start2,
req_cpi=Decimal("1.23"),
req_survey_id="22222",
@@ -264,10 +298,10 @@ class TestWallCacheManager:
attempts10000 = [attempts[0]] * 6000
wall_cache_manager.update_attempts_redis_(attempts10000, user_id=user.user_id)
- session = session_manager.create_dummy(started=start3, user=user)
- wall3 = wall_manager.create_dummy(
+ session = bare_session_factory(started=start3, user=user)
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start3,
req_cpi=Decimal("1.23"),
req_survey_id="33333",
@@ -279,5 +313,5 @@ class TestWallCacheManager:
redis_key = wall_cache_manager.get_cache_key_(user_id=user.user_id)
assert wall_cache_manager.redis_client.llen(redis_key) == 5000
- assert len(attempts) == 5000
+ assert len(attempts) == 5_000
assert attempts[0].req_survey_id == "33333"
diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py
index a80afbe..5b2ff0d 100644
--- a/tests/models/admin/test_report_request.py
+++ b/tests/models/admin/test_report_request.py
@@ -1,4 +1,4 @@
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pandas as pd
import pytest
@@ -19,19 +19,19 @@ class TestReportRequest:
assert rr.report_type == ReportType.POP_SESSION
assert rr.start != rr.start_floor, "rr.start != rr.start_floor"
- assert rr.start_floor.tzinfo == timezone.utc, "rr.start_floor.tzinfo not utc"
+ assert rr.start_floor.tzinfo == UTC, "rr.start_floor.tzinfo not utc"
rr1 = ReportRequest.model_validate(
{
"start": datetime(
- year=datetime.now(tz=timezone.utc).year,
+ year=datetime.now(tz=UTC).year,
month=1,
day=1,
hour=0,
minute=30,
second=25,
microsecond=35,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
),
"interval": "1h",
}
@@ -43,14 +43,14 @@ class TestReportRequest:
rr2 = ReportRequest.model_validate(
{
"start": datetime(
- year=datetime.now(tz=timezone.utc).year,
+ year=datetime.now(tz=UTC).year,
month=1,
day=1,
hour=6,
minute=30,
second=25,
microsecond=35,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
),
"interval": "1d",
}
@@ -92,8 +92,8 @@ class TestReportRequest:
with pytest.raises(expected_exception=ValidationError):
ReportRequest.model_validate(
{
- "start": datetime(year=1990, month=1, day=1, tzinfo=timezone.utc),
- "end": datetime(year=1950, month=1, day=1, tzinfo=timezone.utc),
+ "start": datetime(year=1990, month=1, day=1, tzinfo=UTC),
+ "end": datetime(year=1950, month=1, day=1, tzinfo=UTC),
}
)
@@ -156,8 +156,8 @@ class TestReportRequest:
rr = ReportRequest.model_validate(
{
"interval": "1d",
- "start": datetime(year=2000, month=1, day=1, tzinfo=timezone.utc),
- "end": datetime(year=2000, month=1, day=10, tzinfo=timezone.utc),
+ "start": datetime(year=2000, month=1, day=1, tzinfo=UTC),
+ "end": datetime(year=2000, month=1, day=10, tzinfo=UTC),
}
)
diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py
index 530142e..e8a5aa3 100644
--- a/tests/models/custom_types/test_aware_datetime.py
+++ b/tests/models/custom_types/test_aware_datetime.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
-from datetime import datetime, timezone
+from datetime import UTC, datetime
import pytest
import pytz
@@ -27,14 +27,14 @@ class TestAwareDatetimeISO:
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
def test_dt(self):
- dt = datetime(2023, 10, 10, 1, 1, 1, tzinfo=timezone.utc)
+ dt = datetime(2023, 10, 10, 1, 1, 1, tzinfo=UTC)
t = AwareDatetimeISOModel(dt=dt, dt_optional=dt)
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
t = AwareDatetimeISOModel(dt=dt, dt_optional=None)
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
- dt = datetime(2023, 10, 10, 1, 1, 1, microsecond=123, tzinfo=timezone.utc)
+ dt = datetime(2023, 10, 10, 1, 1, 1, microsecond=123, tzinfo=UTC)
t = AwareDatetimeISOModel(dt=dt, dt_optional=dt)
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
@@ -42,7 +42,7 @@ class TestAwareDatetimeISO:
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
def test_no_tz(self):
- dt = datetime(2023, 10, 10, 1, 1, 1)
+ dt = datetime(2023, 10, 10, 1, 1, 1) # noqa
with pytest.raises(expected_exception=ValidationError):
AwareDatetimeISOModel(dt=dt, dt_optional=None)
diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py
index 16e1f83..eff02d3 100644
--- a/tests/models/custom_types/test_dsn.py
+++ b/tests/models/custom_types/test_dsn.py
@@ -1,4 +1,5 @@
-from typing import Optional
+from __future__ import annotations
+
from uuid import uuid4
import pytest
@@ -11,16 +12,15 @@ from generalresearch.models.custom_types import DaskDsn, SentryDsn
class SettingsModel(BaseModel):
- dask: Optional["DaskDsn"] = Field(default=None)
- sentry: Optional["SentryDsn"] = Field(default=None)
- db: Optional["MySQLDsn"] = Field(default=None)
+ dask: DaskDsn | None = Field(default=None)
+ sentry: SentryDsn | None = Field(default=None)
+ db: MySQLDsn | None = Field(default=None)
# --- Pytest themselves ---
class TestDaskDsn:
-
def test_base(self):
from dask.distributed import Client
diff --git a/tests/models/custom_types/test_therest.py b/tests/models/custom_types/test_therest.py
index 13e9bae..01bc644 100644
--- a/tests/models/custom_types/test_therest.py
+++ b/tests/models/custom_types/test_therest.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import json
from uuid import UUID
diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py
index 736c971..b3a9f13 100644
--- a/tests/models/dynata/test_eligbility.py
+++ b/tests/models/dynata/test_eligbility.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
class TestEligibility:
@@ -40,7 +42,7 @@ class TestEligibility:
"project_id": "p1",
"status": "OPEN",
"project_exclusions": set(),
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
"category_exclusions": set(),
"category_ids": set(),
"cpi": 1,
@@ -172,7 +174,7 @@ class TestEligibility:
"project_id": "p1",
"status": "OPEN",
"project_exclusions": set(),
- "created": datetime.now(tz=timezone.utc),
+ "created": datetime.now(tz=UTC),
"category_exclusions": set(),
"category_ids": set(),
"cpi": 1,
diff --git a/tests/models/dynata/test_survey.py b/tests/models/dynata/test_survey.py
index ad953a3..3e33897 100644
--- a/tests/models/dynata/test_survey.py
+++ b/tests/models/dynata/test_survey.py
@@ -1,3 +1,6 @@
+from __future__ import annotations
+
+
class TestDynataCondition:
def test_condition_create(self):
diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py
index 6c84a5d..21e07a4 100644
--- a/tests/models/gr/test_authentication.py
+++ b/tests/models/gr/test_authentication.py
@@ -1,21 +1,30 @@
+from __future__ import annotations
+
import binascii
import json
import os
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from random import randint
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models.gr.authentication import GRUser
-from generalresearch.models.gr.team import Membership, Team
+from generalresearch.models.gr.authentication import Claims, GRToken, GRUser
+from generalresearch.models.gr.team import Team
+
+if TYPE_CHECKING:
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.gr.team import Membership
+ from generalresearch.models.thl.product import Product
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
SSO_ISSUER = ""
class TestGRUser:
-
def test_init(self, gr_user: GRUser):
assert isinstance(gr_user, GRUser)
@@ -29,7 +38,13 @@ class TestGRUser:
def test_businesses(self):
pass
- def test_teams(self, gr_user: GRUser, membership, gr_db, gr_redis_config):
+ def test_teams(
+ self,
+ gr_user: GRUser,
+ gr_membership: Membership,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
assert gr_user.teams is None
@@ -41,18 +56,18 @@ class TestGRUser:
def test_prefetch_team_duplicates(
self,
- gr_user_token,
+ gr_user_token: GRToken,
gr_user: GRUser,
- membership: Membership,
- product_factory,
- membership_factory,
- team: Team,
- thl_web_rr,
- gr_redis_config,
- gr_db,
+ gr_membership: Membership,
+ product_factory: Callable[..., Product],
+ gr_membership_factory: Callable[..., Membership],
+ gr_team: Team,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
):
- product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ product_factory(team=gr_team)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
gr_user.prefetch_teams(
pg_config=gr_db,
@@ -64,12 +79,12 @@ class TestGRUser:
def test_products(
self,
gr_user: GRUser,
- product_factory,
- team: Team,
- membership: Membership,
- gr_db,
- thl_web_rr,
- gr_redis_config,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership: Membership,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
):
from generalresearch.models.thl.product import Product
@@ -77,13 +92,15 @@ class TestGRUser:
# Create a new Team membership, and then create a Product that
# is part of that team
- membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config)
- p: Product = product_factory(team=team)
+ gr_membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config)
+ assert isinstance(gr_membership.team, Team)
+
+ p: Product = product_factory(team=gr_team)
assert p.id_int
- assert team.uuid == membership.team.uuid
- assert p.team_id == team.uuid
- assert p.team_uuid == membership.team.uuid
- assert gr_user.id == membership.user_id
+ assert gr_team.uuid == gr_membership.team.uuid
+ assert p.team_id == gr_team.uuid
+ assert p.team_uuid == gr_membership.team.uuid
+ assert gr_user.id == gr_membership.user_id
gr_user.prefetch_products(
pg_config=gr_db,
@@ -96,8 +113,7 @@ class TestGRUser:
class TestGRUserMethods:
-
- def test_cache_key(self, gr_user, gr_redis):
+ def test_cache_key(self, gr_user: GRUser):
assert isinstance(gr_user.cache_key, str)
assert ":" in gr_user.cache_key
assert str(gr_user.id) in gr_user.cache_key
@@ -105,14 +121,13 @@ class TestGRUserMethods:
def test_to_redis(
self,
gr_user: GRUser,
- gr_redis,
- team: Team,
- business,
- product_factory,
- membership_factory: Callable[Membership],
+ gr_team: Team,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ gr_membership_factory: Callable[..., Membership],
):
- product_factory(team=team, business=business)
- membership_factory(team=team, gr_user=gr_user)
+ product_factory(team=gr_team, business=gr_business)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
res = gr_user.to_redis()
assert isinstance(res, str)
@@ -125,49 +140,50 @@ class TestGRUserMethods:
def test_set_cache(
self,
gr_user: GRUser,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- gr_redis_config,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
):
- assert gr_redis.get(name=gr_user.cache_key) is None
- assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is None
- assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is None
- assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is None
+
+ client = gr_redis_config.create_redis_client()
+
+ assert client.get(name=gr_user.cache_key) is None
+ assert client.get(name=f"{gr_user.cache_key}:team_uuids") is None
+ assert client.get(name=f"{gr_user.cache_key}:business_uuids") is None
+ assert client.get(name=f"{gr_user.cache_key}:product_uuids") is None
gr_user.set_cache(
pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
- assert gr_redis.get(name=gr_user.cache_key) is not None
- assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is not None
- assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is not None
- assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is not None
+ assert client.get(name=gr_user.cache_key) is not None
+ assert client.get(name=f"{gr_user.cache_key}:team_uuids") is not None
+ assert client.get(name=f"{gr_user.cache_key}:business_uuids") is not None
+ assert client.get(name=f"{gr_user.cache_key}:product_uuids") is not None
def test_set_cache_gr_user(
self,
gr_user: GRUser,
- gr_user_token,
- gr_redis,
- gr_redis_config,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- thl_redis_config,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership_factory: Callable[..., Membership],
+ thl_redis_config: RedisConfig,
):
from generalresearch.models.gr.authentication import GRUser
- p1 = product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ client = gr_redis_config.create_redis_client()
+
+ p1 = product_factory(team=gr_team)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
gr_user.set_cache(
pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
- res: str = gr_redis.get(name=gr_user.cache_key)
+ res: str = client.get(name=gr_user.cache_key)
gru2 = GRUser.from_redis(res)
assert gr_user.model_dump_json(
@@ -183,22 +199,21 @@ class TestGRUserMethods:
def test_set_cache_team_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- gr_redis_config,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
+ gr_membership,
):
- product_factory(team=team)
+ product_factory(team=gr_team)
+ client = gr_redis_config.create_redis_client()
gr_user.set_cache(
pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
- res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:team_uuids"))
+ res = json.loads(client.get(name=f"{gr_user.cache_key}:team_uuids"))
assert len(res) == 1
assert gr_user.team_uuids == res
@@ -206,81 +221,74 @@ class TestGRUserMethods:
def test_set_cache_business_uuids(
self,
gr_user: GRUser,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- business,
- team,
- gr_redis_config,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_business: Business,
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
):
- product_factory(team=team, business=business)
+ product_factory(team=gr_team, business=gr_business)
gr_user.set_cache(
pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
- res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:business_uuids"))
+
+ client = gr_redis_config.create_redis_client()
+ res = json.loads(client.get(name=f"{gr_user.cache_key}:business_uuids"))
assert len(res) == 1
assert gr_user.business_uuids == res
def test_set_cache_product_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- gr_redis_config,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
+ gr_membership,
):
- product_factory(team=team)
+ product_factory(team=gr_team)
gr_user.set_cache(
pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config
)
- res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:product_uuids"))
+ client = gr_redis_config.create_redis_client()
+ res = json.loads(client.get(name=f"{gr_user.cache_key}:product_uuids"))
assert len(res) == 1
assert gr_user.product_uuids == res
class TestGRToken:
-
@pytest.fixture
- def gr_token(self, gr_user):
- from generalresearch.models.gr.authentication import GRToken
-
- now = datetime.now(tz=timezone.utc)
+ def gr_token(self, gr_user: GRUser):
+ now = datetime.now(tz=UTC)
token = binascii.hexlify(os.urandom(20)).decode()
gr_token = GRToken(key=token, created=now, user_id=gr_user.id)
return gr_token
- def test_init(self, gr_token):
- from generalresearch.models.gr.authentication import GRToken
-
+ def test_init(self, gr_token: GRToken):
assert isinstance(gr_token, GRToken)
assert gr_token.created
- def test_user(self, gr_token, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRUser
-
+ def test_user(
+ self, gr_token: GRToken, gr_db: PostgresConfig, gr_redis_config: RedisConfig
+ ):
assert gr_token.user is None
gr_token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config)
assert isinstance(gr_token.user, GRUser)
- def test_auth_header(self, gr_token):
+ def test_auth_header(self, gr_token: GRToken):
assert isinstance(gr_token.auth_header, dict)
class TestClaims:
-
def test_init(self):
- from generalresearch.models.gr.authentication import Claims
d = {
"iss": SSO_ISSUER,
diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py
index a9f01a8..5ab5dff 100644
--- a/tests/models/gr/test_base.py
+++ b/tests/models/gr/test_base.py
@@ -1,16 +1,20 @@
+from __future__ import annotations
+
import subprocess
+from collections.abc import Callable
from pathlib import Path
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
from pydantic import PostgresDsn
-from generalresearch.pg_helper import PostgresConfig
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
class TestGRPostgresDjangoCreation:
- def test_git(self, git_key_path: Path, gr_repo: Callable[..., Path]):
+ def test_git(self, gr_repo: Callable[..., Path]):
repo_path = gr_repo()
try:
@@ -33,14 +37,5 @@ class TestGRPostgresDjangoCreation:
django_db_factory: Callable[..., None],
):
- dsn = django_db_factory("gr")
+ dsn = django_db_factory("gr.common")
assert isinstance(dsn, PostgresDsn)
-
- # def test_django_tables(self, thl_web_rw: PostgresConfig):
- # res = thl_web_rw.execute_sql_query(query="""
- # SELECT COUNT(*)
- # FROM information_schema.tables
- # WHERE table_schema = 'public';
- # """)
- # assert len(res) == 1
- # assert res[0]["count"] == 56
diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py
index 7a84f23..e38850d 100644
--- a/tests/models/gr/test_business.py
+++ b/tests/models/gr/test_business.py
@@ -1,7 +1,11 @@
+from __future__ import annotations
+
import os
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Optional
+from pathlib import Path
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
@@ -15,34 +19,55 @@ from distributed.utils_test import (
from pytest import approx
from generalresearch.currency import USDCent
-from generalresearch.managers.gr.business import BusinessBankAccountManager
from generalresearch.models.gr.business import (
Business,
BusinessAddress,
- BusinessBankAccount,
BusinessContact,
)
from generalresearch.models.thl.finance import (
BusinessBalances,
ProductBalances,
)
-from generalresearch.pg_helper import PostgresConfig
+from generalresearch.models.thl.product import Product
+
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+ )
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.gr.business import BusinessBankAccountManager
+ from generalresearch.managers.gr.team import TeamManager
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.payout import (
+ BusinessPayoutEventManager,
+ PayoutEventManager,
+ )
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.business import (
+ BusinessBankAccount,
+ )
+ from generalresearch.models.gr.team import Team
+ from generalresearch.models.thl.product import BrokerageProductPayoutEvent
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class TestBusinessBankAccount:
-
def test_init(
self,
- business: Business,
- business_bank_account_manager: BusinessBankAccountManager,
+ gr_business: Business,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
):
- from generalresearch.models.gr.business import (
- BusinessBankAccount,
- TransferMethod,
- )
+ from generalresearch.models.gr.business import BusinessBankAccount
+ from generalresearch.models.gr.definitions import TransferMethod
- instance = business_bank_account_manager.create(
- business_id=business.id,
+ instance = gr_business_bank_account_manager.create(
+ business_id=gr_business.id,
uuid=uuid4().hex,
transfer_method=TransferMethod.ACH,
)
@@ -50,30 +75,28 @@ class TestBusinessBankAccount:
def test_business(
self,
- business_bank_account: BusinessBankAccount,
- business: Business,
- gr_db,
- gr_redis_config,
+ gr_business_bank_account: BusinessBankAccount,
+ gr_business: Business,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
):
from generalresearch.models.gr.business import Business
- assert business_bank_account.business is None
+ assert gr_business_bank_account.business is None
- business_bank_account.prefetch_business(
+ gr_business_bank_account.prefetch_business(
pg_config=gr_db, redis_config=gr_redis_config
)
- assert isinstance(business_bank_account.business, Business)
- assert business_bank_account.business.uuid == business.uuid
+ assert isinstance(gr_business_bank_account.business, Business)
+ assert gr_business_bank_account.business.uuid == gr_business.uuid
class TestBusinessAddress:
-
- def test_init(self, business_address: BusinessAddress):
- assert isinstance(business_address, BusinessAddress)
+ def test_init(self, gr_business_address: BusinessAddress):
+ assert isinstance(gr_business_address, BusinessAddress)
class TestBusinessContact:
-
def test_init(self):
bc = BusinessContact(name="abc", email="test@abc.com")
@@ -82,346 +105,352 @@ class TestBusinessContact:
class TestBusiness:
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
- def test_init(self, business):
- from generalresearch.models.gr.business import Business
+ def test_init(self, gr_business: Business):
- assert isinstance(business, Business)
- assert isinstance(business.id, int)
- assert isinstance(business.uuid, str)
+ assert isinstance(gr_business, Business)
+ assert isinstance(gr_business.id, int)
+ assert isinstance(gr_business.uuid, str)
def test_str_and_repr(
self,
- business,
- product_factory,
- thl_web_rr,
- lm,
- thl_lm,
- business_payout_event_manager,
- bp_payout_factory,
- start,
- user_factory,
- session_with_tx_factory,
- pop_ledger_merge,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
+ product_manager: ProductManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BusinessPayoutEventManager
+ ],
+ start: datetime,
+ user_factory: Callable[..., User],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
client_no_amm: DaskClient,
ledger_collection,
- mnt_filepath,
- create_main_accounts,
+ mnt_filepath: GRLDatasets,
+ create_main_accounts: Callable[..., None],
):
create_main_accounts()
- p1 = product_factory(business=business)
+ p1 = product_factory(business=gr_business)
u1 = user_factory(product=p1)
- p2 = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
+ p2 = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
- res1 = repr(business)
+ res1 = repr(gr_business)
- assert business.uuid in res1
+ assert gr_business.uuid in res1
assert "<Business: " in res1
- res2 = str(business)
+ res2 = str(gr_business)
- assert business.uuid in res2
+ assert gr_business.uuid in res2
assert "Name:" in res2
assert "Not Loaded" in res2
- business.prefetch_products(thl_pg_config=thl_web_rr)
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- res3 = str(business)
+ gr_business.prefetch_products(product_manager=product_manager)
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ res3 = str(gr_business)
assert "Products: 2" in res3
assert "Ledger Accounts: 2" in res3
# -- need some tx to make these interesting
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
session_with_tx_factory(
user=u1,
wall_req_cpi=Decimal("2.50"),
started=start + timedelta(days=5),
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(50),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- res4 = str(business)
+ res4 = str(gr_business)
assert "Payouts: 1" in res4
assert "Available Balance: 141" in res4
- def test_addresses(self, business, business_address, gr_db):
+ def test_addresses(
+ self, gr_business: Business, gr_db: PostgresConfig, gr_business_address
+ ):
from generalresearch.models.gr.business import BusinessAddress
- assert business.addresses is None
+ assert gr_business.addresses is None
- business.prefetch_addresses(pg_config=gr_db)
- assert isinstance(business.addresses, list)
- assert len(business.addresses) == 1
- assert isinstance(business.addresses[0], BusinessAddress)
+ gr_business.prefetch_addresses(pg_config=gr_db)
+ assert isinstance(gr_business.addresses, list)
+ assert len(gr_business.addresses) == 1
+ assert isinstance(gr_business.addresses[0], BusinessAddress)
- def test_teams(self, business, team, team_manager, gr_db):
- assert business.teams is None
+ def test_teams(
+ self,
+ gr_business: Business,
+ gr_team: Team,
+ gr_team_manager: TeamManager,
+ gr_db: PostgresConfig,
+ ):
+ assert gr_business.teams is None
- business.prefetch_teams(pg_config=gr_db)
- assert isinstance(business.teams, list)
- assert len(business.teams) == 0
+ gr_business.prefetch_teams(pg_config=gr_db)
+ assert isinstance(gr_business.teams, list)
+ assert len(gr_business.teams) == 0
- team_manager.add_business(team=team, business=business)
- assert len(business.teams) == 0
- business.prefetch_teams(pg_config=gr_db)
- assert len(business.teams) == 1
+ gr_team_manager.add_business(team=gr_team, business=gr_business)
+ assert len(gr_business.teams) == 0
+ gr_business.prefetch_teams(pg_config=gr_db)
+ assert len(gr_business.teams) == 1
- def test_products(self, business, product_factory, thl_web_rr):
- from generalresearch.models.thl.product import Product
+ def test_products(
+ self,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ product_manager: ProductManager,
+ ):
- p1 = product_factory(business=business)
- assert business.products is None
+ p1 = product_factory(business=gr_business)
+ assert gr_business.products is None
- business.prefetch_products(thl_pg_config=thl_web_rr)
- assert isinstance(business.products, list)
- assert len(business.products) == 1
- assert isinstance(business.products[0], Product)
+ gr_business.prefetch_products(product_manager=product_manager)
+ assert isinstance(gr_business.products, list)
+ assert len(gr_business.products) == 1
+ assert isinstance(gr_business.products[0], Product)
- assert business.products[0].uuid == p1.uuid
+ assert gr_business.products[0].uuid == p1.uuid
# Add two more, but list is still one until we prefetch
- p2 = product_factory(business=business)
- p3 = product_factory(business=business)
- assert len(business.products) == 1
+ product_factory(business=gr_business)
+ product_factory(business=gr_business)
+ assert len(gr_business.products) == 1
- business.prefetch_products(thl_pg_config=thl_web_rr)
- assert len(business.products) == 3
+ gr_business.prefetch_products(product_manager=product_manager)
+ assert len(gr_business.products) == 3
- def test_bank_accounts(self, business, business_bank_account, gr_db):
- assert business.products is None
+ def test_bank_accounts(
+ self,
+ gr_business: Business,
+ gr_business_bank_account,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ ):
+ assert gr_business.products is None
# It's an empty list after prefetch
- business.prefetch_bank_accounts(pg_config=gr_db)
- assert isinstance(business.bank_accounts, list)
- assert len(business.bank_accounts) == 1
+ gr_business.prefetch_bank_accounts(
+ business_bank_account_manager=gr_business_bank_account_manager
+ )
+ assert isinstance(gr_business.bank_accounts, list)
+ assert len(gr_business.bank_accounts) == 1
def test_balance(
self,
- business: Business,
- mnt_filepath,
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
client_no_amm: DaskClient,
thl_web_rr: PostgresConfig,
- ledger_manager,
- pop_ledger_merge,
+ ledger_manager: LedgerManager,
+ pop_ledger_merge: PopLedgerMerge,
+ product_manager: ProductManager,
):
- assert business.balance is None
+ assert gr_business.balance is None
with pytest.raises(expected_exception=AssertionError) as cm:
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
assert "Cannot build Business Balance" in str(cm.value)
- assert business.balance is None
+ assert gr_business.balance is None
# TODO: Add parquet building so that this doesn't fail and we can
# properly assign a business.balance
def test_payouts_no_accounts(
self,
- business,
- product_factory,
- thl_web_rr,
- thl_ledger_manager,
- business_payout_event_manager,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
):
- assert business.payouts is None
+ assert gr_business.payouts is None
with pytest.raises(expected_exception=AssertionError) as cm:
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
assert "Must provide product_uuids" in str(cm.value)
- p = product_factory(business=business)
+ p = product_factory(business=gr_business)
thl_ledger_manager.get_account_or_create_bp_wallet(product=p)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert isinstance(business.payouts, list)
- assert len(business.payouts) == 0
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 0
def test_payouts(
self,
- business: Business,
- product_factory: Callable[Product],
- bp_payout_factory,
- thl_ledger_manager,
- thl_web_rr,
- business_payout_event_manager,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ create_main_accounts: Callable[..., None],
):
create_main_accounts()
- p = product_factory(business=business)
+ p = product_factory(business=gr_business)
thl_ledger_manager.get_account_or_create_bp_wallet(product=p)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
- product=p, amount=USDCent(123), skip_wallet_balance_check=True
- )
+ brokerage_product_payout_event_factory(product=p, amount=USDCent(123))
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert sum([p.amount for p in business.payouts]) == 123
+ assert len(gr_business.payouts) == 1
+ assert sum([p.amount for p in gr_business.payouts]) == 123
# Add another!
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p,
amount=USDCent(123),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
- )
- business_payout_event_manager.set_account_lookup_table(
- thl_lm=thl_ledger_manager
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 2
- assert sum([p.amount for p in business.payouts]) == 246
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 2
+ assert len(gr_business.payouts[0].bp_payouts) == 1
+ assert sum([p.amount for p in gr_business.payouts]) == 246
def test_payouts_totals(
self,
- business,
- product_factory,
- bp_payout_factory,
- thl_lm,
- thl_web_rr,
- business_payout_event_manager,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ thl_web_rr: PostgresConfig,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ create_main_accounts: Callable[..., None],
):
- from generalresearch.models.thl.product import Product
create_main_accounts()
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(1),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(25),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(50),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 3
- assert business.payouts_total == USDCent(76)
- assert business.payouts_total_str == "$0.76"
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 3
+ assert len(gr_business.payouts[0].bp_payouts) == 1
+ assert len(gr_business.payouts[1].bp_payouts) == 1
+ assert len(gr_business.payouts[2].bp_payouts) == 1
+ assert gr_business.payouts_total == USDCent(76)
+ assert gr_business.payouts_total_str == "$0.76"
def test_pop_financial(
self,
- business,
- thl_web_rr,
- thl_ledger_manager,
- mnt_filepath,
- client_no_amm,
- pop_ledger_merge,
+ gr_business: Business,
+ product_manager: ProductManager,
+ thl_ledger_manager: ThlLedgerManager,
+ mnt_filepath: GRLDatasets,
+ client_no_amm: DaskClient,
+ pop_ledger_merge: PopLedgerMerge,
):
- assert business.pop_financial is None
- business.prebuild_pop_financial(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ assert gr_business.pop_financial is None
+ gr_business.prebuild_pop_financial(
+ product_manager=product_manager,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.pop_financial == []
+ assert gr_business.pop_financial == []
- def test_bp_accounts(self, business, lm, thl_web_rr, product_factory, thl_lm):
- assert business.bp_accounts is None
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- assert business.bp_accounts == []
-
- from generalresearch.models.thl.product import Product
+ def test_bp_accounts(
+ self,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ product_manager: ProductManager,
+ ):
+ assert gr_business.bp_accounts is None
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ assert gr_business.bp_accounts == []
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- assert len(business.bp_accounts) == 1
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ assert len(gr_business.bp_accounts) == 1
class TestBusinessBalance:
-
@pytest.fixture
- def start(self) -> "datetime":
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ def start(self) -> datetime:
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
@pytest.mark.skip
@@ -432,34 +461,27 @@ class TestBusinessBalance:
def test_single_product(
self,
- business,
- product_factory,
- user_factory,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ ledger_manager: LedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ session_with_tx_factory: Callable[..., Session],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ product_manager: ProductManager,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p1)
@@ -478,57 +500,50 @@ class TestBusinessBalance:
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert isinstance(business.balance, BusinessBalances)
- assert business.balance.payout == 190
- assert business.balance.adjustment == 0
- assert business.balance.net == 190
- assert business.balance.retainer == 47
- assert business.balance.available_balance == 143
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.adjustment == 0
+ assert gr_business.balance.net == 190
+ assert gr_business.balance.retainer == 47
+ assert gr_business.balance.available_balance == 143
- assert len(business.balance.product_balances) == 1
- pb = business.balance.product_balances[0]
+ assert len(gr_business.balance.product_balances) == 1
+ pb = gr_business.balance.product_balances[0]
assert isinstance(pb, ProductBalances)
- assert pb.balance == business.balance.balance
- assert pb.available_balance == business.balance.available_balance
+ assert pb.balance == gr_business.balance.balance
+ assert pb.available_balance == gr_business.balance.available_balance
assert pb.adjustment_percent == 0.0
def test_multi_product(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- ledger_manager,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ start: datetime,
+ session_with_tx_factory: Callable[..., Session],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: User = user_factory(product=product_factory(business=business))
+ u1: User = user_factory(product=product_factory(business=gr_business))
+ u2: User = user_factory(product=product_factory(business=gr_business))
session_with_tx_factory(
user=u1,
@@ -545,33 +560,33 @@ class TestBusinessBalance:
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert isinstance(business.balance, BusinessBalances)
- assert business.balance.payout == 190
- assert business.balance.balance == 190
- assert business.balance.adjustment == 0
- assert business.balance.net == 190
- assert business.balance.retainer == 46
- assert business.balance.available_balance == 144
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.balance == 190
+ assert gr_business.balance.adjustment == 0
+ assert gr_business.balance.net == 190
+ assert gr_business.balance.retainer == 46
+ assert gr_business.balance.available_balance == 144
- assert len(business.balance.product_balances) == 2
+ assert len(gr_business.balance.product_balances) == 2
- pb1 = business.balance.product_balances[0]
- pb2 = business.balance.product_balances[1]
+ pb1 = gr_business.balance.product_balances[0]
+ pb2 = gr_business.balance.product_balances[1]
assert isinstance(pb1, ProductBalances)
assert pb1.product_id == u1.product_id
assert isinstance(pb2, ProductBalances)
assert pb2.product_id == u2.product_id
for pb in [pb1, pb2]:
- assert pb.balance != business.balance.balance
- assert pb.available_balance != business.balance.available_balance
+ assert pb.balance != gr_business.balance.balance
+ assert pb.available_balance != gr_business.balance.available_balance
assert pb.adjustment_percent == 0.0
assert pb1.product_id in [u1.product_id, u2.product_id]
@@ -592,34 +607,33 @@ class TestBusinessBalance:
def test_multi_product_multi_payout(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ product_manager: ProductManager,
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ payout_event_manager: PayoutEventManager,
+ session_with_tx_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: User = user_factory(product=product_factory(business=business))
+ u1: User = user_factory(product=product_factory(business=gr_business))
+ u2: User = user_factory(product=product_factory(business=gr_business))
session_with_tx_factory(
user=u1,
@@ -633,62 +647,58 @@ class TestBusinessBalance:
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(5),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.balance.payout == 190
- assert business.balance.net == 190
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.net == 190
- assert business.balance.balance == 135
+ assert gr_business.balance.balance == 135
def test_multi_product_multi_payout_adjustment(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ payout_event_manager: PayoutEventManager,
+ session_with_tx_factory: Callable[..., Session],
+ delete_ledger_db: Callable[..., None],
+ product_manager: ProductManager,
+ create_main_accounts: Callable[..., None],
ledger_collection,
task_adj_collection,
- pop_ledger_merge,
- wall_manager,
- session_manager,
- adj_to_fail_with_tx_factory,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""
- Product 1 $2.50 Complete
@@ -712,11 +722,9 @@ class TestBusinessBalance:
delete_df_collection(coll=ledger_collection)
delete_df_collection(coll=task_adj_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: User = user_factory(product=product_factory(business=business))
- u3: User = user_factory(product=product_factory(business=business))
+ u1: User = user_factory(product=product_factory(business=gr_business))
+ u2: User = user_factory(product=product_factory(business=gr_business))
+ u3: User = user_factory(product=product_factory(business=gr_business))
s1 = session_with_tx_factory(
user=u1,
@@ -729,22 +737,17 @@ class TestBusinessBalance:
wall_req_cpi=Decimal("2.50"),
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
adj_to_fail_with_tx_factory(session=s1, created=start + timedelta(days=5))
@@ -769,57 +772,60 @@ class TestBusinessBalance:
df = client_no_amm.compute(pop_ledger_merge.ddf(), sync=True)
assert df.shape == (20, 28)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.balance.payout == 714
- assert business.balance.adjustment == -238
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 714
+ assert gr_business.balance.adjustment == -238
- assert business.balance.product_balances[0].adjustment == -238
- assert business.balance.product_balances[1].adjustment == 0
- assert business.balance.product_balances[2].adjustment == 0
+ assert gr_business.balance.product_balances[0].adjustment == -238
+ assert gr_business.balance.product_balances[1].adjustment == 0
+ assert gr_business.balance.product_balances[2].adjustment == 0
- assert business.balance.expense == 0
- assert business.balance.net == 714 - 238
- assert business.balance.balance == business.balance.payout - (250 + 50 + 238)
+ assert gr_business.balance.expense == 0
+ assert gr_business.balance.net == 714 - 238
+ assert gr_business.balance.balance == gr_business.balance.payout - (
+ 250 + 50 + 238
+ )
predicted_retainer = sum(
[
pb.balance * 0.25
- for pb in business.balance.product_balances
+ for pb in gr_business.balance.product_balances
if pb.balance > 0
]
)
- assert business.balance.retainer == approx(predicted_retainer, rel=0.01)
+ assert gr_business.balance.retainer == approx(predicted_retainer, rel=0.01)
def test_neg_balance_cache(
self,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
payout_event_manager,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
+ product_manager: ProductManager,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
):
"""Test having a Business with two products.. one that lost money
and one that gained money. Ensure that the Business balance
@@ -830,15 +836,12 @@ class TestBusinessBalance:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- p2: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
+ p2: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p2)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
# Product 1: Complete, Payout, Recon..
s1 = session_with_tx_factory(
@@ -846,14 +849,11 @@ class TestBusinessBalance:
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
created=start + timedelta(days=1, minutes=1),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
adj_to_fail_with_tx_factory(
session=s1,
@@ -861,12 +861,12 @@ class TestBusinessBalance:
)
# Product 2: Complete, Complete.
- s2 = session_with_tx_factory(
+ session_with_tx_factory(
user=u2,
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1, minutes=3),
)
- s3 = session_with_tx_factory(
+ session_with_tx_factory(
user=u2,
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1, minutes=4),
@@ -876,16 +876,17 @@ class TestBusinessBalance:
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
# Check Product 1
- pb1 = business.balance.product_balances[0]
+ assert isinstance(gr_business.balance, BusinessBalances)
+ pb1 = gr_business.balance.product_balances[0]
assert pb1.product_id == p1.uuid
assert pb1.payout == 71
assert pb1.adjustment == -71
@@ -895,7 +896,7 @@ class TestBusinessBalance:
assert pb1.available_balance == 0
# Check Product 2
- pb2 = business.balance.product_balances[1]
+ pb2 = gr_business.balance.product_balances[1]
assert pb2.product_id == p2.uuid
assert pb2.payout == 71 * 2
assert pb2.adjustment == 0
@@ -905,7 +906,8 @@ class TestBusinessBalance:
assert pb2.available_balance == 107
# Check Business
- bb1 = business.balance
+ bb1 = gr_business.balance
+ assert isinstance(bb1, BusinessBalances)
assert bb1.payout == (71 * 3) # Raw total of completes
assert bb1.adjustment == -71 # 1 Complete >> Failure
assert bb1.expense == 0
@@ -923,29 +925,27 @@ class TestBusinessBalance:
def test_multi_product_multi_payout_adjustment_at_timestamp(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ start: datetime,
payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ session_with_tx_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
task_adj_collection,
- pop_ledger_merge,
- wall_manager,
- session_manager,
- adj_to_fail_with_tx_factory,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""
This test measures a complex Business situation, but then makes
@@ -985,11 +985,9 @@ class TestBusinessBalance:
delete_df_collection(coll=ledger_collection)
delete_df_collection(coll=task_adj_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: User = user_factory(product=product_factory(business=business))
- u3: User = user_factory(product=product_factory(business=business))
+ u1: User = user_factory(product=product_factory(business=gr_business))
+ u2: User = user_factory(product=product_factory(business=gr_business))
+ u3: User = user_factory(product=product_factory(business=gr_business))
s1 = session_with_tx_factory(
user=u1,
@@ -1002,22 +1000,17 @@ class TestBusinessBalance:
wall_req_cpi=Decimal("2.50"),
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u2.product,
amount=USDCent(50),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
session_with_tx_factory(
@@ -1042,73 +1035,80 @@ class TestBusinessBalance:
df = client_no_amm.compute(pop_ledger_merge.ddf(), sync=True)
assert df.shape == (20, 28)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=1, hours=1),
)
- day1_bal = business.balance
+ day1_bal = gr_business.balance
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=2, hours=1),
)
- day2_bal = business.balance
+ day2_bal = gr_business.balance
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=3, hours=1),
)
- day3_bal = business.balance
+ day3_bal = gr_business.balance
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=4, hours=1),
)
- day4_bal = business.balance
+ day4_bal = gr_business.balance
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=5, hours=1),
)
- day5_bal = business.balance
+ day5_bal = gr_business.balance
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
at_timestamp=start + timedelta(days=6, hours=1),
)
- day6_bal = business.balance
+ day6_bal = gr_business.balance
+
+ assert isinstance(day1_bal, BusinessBalances)
+ assert isinstance(day2_bal, BusinessBalances)
+ assert isinstance(day3_bal, BusinessBalances)
+ assert isinstance(day4_bal, BusinessBalances)
+ assert isinstance(day5_bal, BusinessBalances)
+ assert isinstance(day6_bal, BusinessBalances)
assert day1_bal.payout == 238
assert day1_bal.retainer == 59
@@ -1136,9 +1136,8 @@ class TestBusinessBalance:
class TestBusinessMethods:
-
@pytest.fixture(scope="function")
- def start(self, utc_90days_ago) -> "datetime":
+ def start(self, utc_90days_ago: datetime) -> datetime:
s = utc_90days_ago.replace(microsecond=0)
return s
@@ -1149,72 +1148,74 @@ class TestBusinessMethods:
@pytest.fixture(scope="function")
def duration(
self,
- ) -> Optional["timedelta"]:
+ ) -> timedelta | None:
return None
- def test_cache_key(self, business, gr_redis):
- assert isinstance(business.cache_key, str)
- assert ":" in business.cache_key
- assert str(business.uuid) in business.cache_key
+ def test_cache_key(self, gr_business: Business):
+ assert isinstance(gr_business.cache_key, str)
+ assert ":" in gr_business.cache_key
+ assert str(gr_business.uuid) in gr_business.cache_key
def test_set_cache(
self,
- business,
- gr_redis,
- gr_db,
- thl_web_rr,
- client_no_amm,
- mnt_filepath,
- lm,
- thl_lm,
+ gr_business: Business,
+ thl_web_rr: PostgresConfig,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
business_payout_event_manager,
- product_factory,
- membership_factory,
- team,
- session_with_tx_factory,
- user_factory,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ product_manager: ProductManager,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ session_with_tx_factory: Callable[..., Session],
+ user_factory: Callable[..., User],
ledger_collection,
- pop_ledger_merge,
- utc_60days_ago,
- delete_ledger_db,
- create_main_accounts,
- gr_redis_config,
- mnt_gr_api_dir,
+ pop_ledger_merge: PopLedgerMerge,
+ utc_60days_ago: datetime,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ gr_redis_config: RedisConfig,
+ mnt_gr_api_dir: Path,
):
- assert gr_redis.get(name=business.cache_key) is None
+ client = gr_redis_config.create_redis_client()
+ assert client.get(name=gr_business.cache_key) is None
- p1 = product_factory(team=team, business=business)
+ p1 = product_factory(team=gr_team, business=gr_business)
u1 = user_factory(product=p1)
# Business needs tx & incite to build balance
delete_ledger_db()
create_main_accounts()
- thl_lm.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
session_with_tx_factory(user=u1, started=utc_60days_ago)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.set_cache(
- pg_config=gr_db,
+ gr_business.set_cache(
+ product_manager=product_manager,
+ business_bank_account_manager=gr_business_bank_account_manager,
+ pg_config=thl_web_rr,
thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
- lm=lm,
- thl_lm=thl_lm,
+ lm=ledger_manager,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
pop_ledger=pop_ledger_merge,
mnt_gr_api=mnt_gr_api_dir,
)
- assert gr_redis.hgetall(name=business.cache_key) is not None
+ assert client.hgetall(name=gr_business.cache_key) is not None
from generalresearch.models.gr.business import Business
# We're going to pull only a specific year, but make sure that
# it's being assigned to the field regardless
- year = datetime.now(tz=timezone.utc).year
+ year = datetime.now(tz=UTC).year
res = Business.from_redis(
- uuid=business.uuid,
+ uuid=gr_business.uuid,
fields=[f"pop_financial:{year}"],
gr_redis_config=gr_redis_config,
)
@@ -1222,53 +1223,53 @@ class TestBusinessMethods:
def test_set_cache_business(
self,
- gr_user,
- business,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- client_no_amm,
- mnt_filepath,
- lm,
- thl_lm,
+ gr_business: Business,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
business_payout_event_manager,
- user_factory,
- delete_ledger_db,
- create_main_accounts,
- session_with_tx_factory,
+ product_manager: ProductManager,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ user_factory: Callable[..., User],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ session_with_tx_factory: Callable[..., Session],
ledger_collection,
- team_manager,
- pop_ledger_merge,
- gr_redis_config,
- utc_60days_ago,
- mnt_gr_api_dir,
+ gr_team_manager: TeamManager,
+ pop_ledger_merge: PopLedgerMerge,
+ gr_redis_config: RedisConfig,
+ utc_60days_ago: datetime,
+ mnt_gr_api_dir: Path,
):
from generalresearch.models.gr.business import Business
- p1 = product_factory(team=team, business=business)
+ p1 = product_factory(team=gr_team, business=gr_business)
u1 = user_factory(product=p1)
- team_manager.add_business(team=team, business=business)
+ gr_team_manager.add_business(team=gr_team, business=gr_business)
# Business needs tx & incite to build balance
delete_ledger_db()
create_main_accounts()
- thl_lm.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
session_with_tx_factory(user=u1, started=utc_60days_ago)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.set_cache(
+ gr_business.set_cache(
+ product_manager=product_manager,
+ business_bank_account_manager=gr_business_bank_account_manager,
pg_config=gr_db,
thl_web_rr=thl_web_rr,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
- lm=lm,
- thl_lm=thl_lm,
+ lm=ledger_manager,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
pop_ledger=pop_ledger_merge,
mnt_gr_api=mnt_gr_api_dir,
@@ -1276,7 +1277,7 @@ class TestBusinessMethods:
# keys: List = Business.required_fields() + ["products", "bp_accounts"]
business2 = Business.from_redis(
- uuid=business.uuid,
+ uuid=gr_business.uuid,
fields=[
"id",
"tax_number",
@@ -1295,11 +1296,16 @@ class TestBusinessMethods:
gr_redis_config=gr_redis_config,
)
- assert business.model_dump_json() == business2.model_dump_json()
+ assert isinstance(business2, Business)
+ assert gr_business.model_dump_json() == business2.model_dump_json()
+ # assert isinstance(business2.balance, BusinessBalances)
+ assert isinstance(business2.products, list)
+ assert isinstance(business2.teams, list)
assert p1.uuid in [p.uuid for p in business2.products]
assert len(business2.teams) == 1
- assert team.uuid in [t.uuid for t in business2.teams]
+ assert gr_team.uuid in [t.uuid for t in business2.teams]
+ assert isinstance(business2.balance, BusinessBalances)
assert business2.balance.payout == 48
assert business2.balance.balance == 48
assert business2.balance.net == 48
@@ -1312,39 +1318,39 @@ class TestBusinessMethods:
assert len(business2.bp_accounts) == 1
assert len(business2.bp_accounts) == len(business2.product_uuids)
+ assert isinstance(business2.pop_financial, list)
assert len(business2.pop_financial) == 1
assert business2.pop_financial[0].payout == business2.balance.payout
assert business2.pop_financial[0].net == business2.balance.net
def test_prebuild_enriched_session_parquet(
self,
- event_report_request,
enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ product_manager: ProductManager,
+ session_collection: SessionDFCollection,
+ thl_web_rr: PostgresConfig,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
- p1 = product_factory(business=business)
- p2 = product_factory(business=business)
+ p1 = product_factory(business=gr_business)
+ p2 = product_factory(business=gr_business)
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -1360,8 +1366,8 @@ class TestBusinessMethods:
pg_config=thl_web_rr,
)
- business.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_enriched_session_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -1370,40 +1376,40 @@ class TestBusinessMethods:
# Now try to read from path
df = pd.read_parquet(
- os.path.join(mnt_gr_api_dir, "pop_session", f"{business.file_key}.parquet")
+ os.path.join(
+ mnt_gr_api_dir, "pop_session", f"{gr_business.file_key}.parquet"
+ )
)
assert isinstance(df, pd.DataFrame)
def test_prebuild_enriched_wall_parquet(
self,
- event_report_request,
- enriched_session_merge,
enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ product_manager: ProductManager,
+ session_collection: SessionDFCollection,
+ thl_web_rr: PostgresConfig,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
- p1 = product_factory(business=business)
- p2 = product_factory(business=business)
+ p1 = product_factory(business=gr_business)
+ p2 = product_factory(business=gr_business)
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -1419,8 +1425,8 @@ class TestBusinessMethods:
pg_config=thl_web_rr,
)
- business.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_enriched_wall_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -1429,6 +1435,6 @@ class TestBusinessMethods:
# Now try to read from path
df = pd.read_parquet(
- os.path.join(mnt_gr_api_dir, "pop_event", f"{business.file_key}.parquet")
+ os.path.join(mnt_gr_api_dir, "pop_event", f"{gr_business.file_key}.parquet")
)
assert isinstance(df, pd.DataFrame)
diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py
index d728bbe..e853817 100644
--- a/tests/models/gr/test_team.py
+++ b/tests/models/gr/test_team.py
@@ -1,127 +1,180 @@
+from __future__ import annotations
+
import os
-from datetime import timedelta
+from collections.abc import Callable
+from datetime import datetime, timedelta
from decimal import Decimal
+from pathlib import Path
+from typing import TYPE_CHECKING
import pandas as pd
+from dask.distributed import Client as DaskClient
+from distributed.utils_test import (
+ client_no_amm,
+)
+
+from generalresearch.models.gr.business import Business
+from generalresearch.models.gr.team import Team
+from generalresearch.models.thl.product import Product
+
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
+ from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+ )
+ from generalresearch.managers.gr.authentication import GRUserManager
+ from generalresearch.managers.gr.business import BusinessManager
+ from generalresearch.managers.gr.team import MembershipManager, TeamManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.authentication import GRUser
+ from generalresearch.models.gr.team import Membership
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class TestTeam:
+ def test_init(self, gr_team: Team):
- def test_init(self, team):
- from generalresearch.models.gr.team import Team
-
- assert isinstance(team, Team)
- assert isinstance(team.id, int)
- assert isinstance(team.uuid, str)
+ assert isinstance(gr_team, Team)
+ assert isinstance(gr_team.id, int)
+ assert isinstance(gr_team.uuid, str)
- def test_memberships_none(self, team, gr_user_factory, gr_db):
- assert team.memberships is None
+ def test_memberships_none(
+ self, gr_team: Team, gr_membership_manager: MembershipManager
+ ):
+ assert gr_team.memberships is None
- team.prefetch_memberships(pg_config=gr_db)
- assert isinstance(team.memberships, list)
- assert len(team.memberships) == 0
+ gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager)
+ assert isinstance(gr_team.memberships, list)
+ assert len(gr_team.memberships) == 0
def test_memberships(
self,
- team,
- membership,
- gr_user,
- gr_user_factory,
- membership_factory,
- membership_manager,
- gr_db,
+ gr_team: Team,
+ gr_user: GRUser,
+ gr_membership,
+ gr_user_factory: Callable[..., GRUser],
+ gr_membership_manager: MembershipManager,
):
- assert team.memberships is None
+ assert gr_team.memberships is None
- team.prefetch_memberships(pg_config=gr_db)
- assert isinstance(team.memberships, list)
- assert len(team.memberships) == 1
- assert team.memberships[0].user_id == gr_user.id
+ gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager)
+ assert isinstance(gr_team.memberships, list)
+ assert len(gr_team.memberships) == 1
+ assert gr_team.memberships[0].user_id == gr_user.id
# Create another new Membership
- membership_manager.create(team=team, gr_user=gr_user_factory())
- assert len(team.memberships) == 1
- team.prefetch_memberships(pg_config=gr_db)
- assert len(team.memberships) == 2
+ gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory())
+ assert len(gr_team.memberships) == 1
+ gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager)
+ assert len(gr_team.memberships) == 2
def test_gr_users(
- self, team, gr_user_factory, membership_manager, gr_db, gr_redis_config
+ self,
+ gr_team: Team,
+ gr_user_factory: Callable[..., GRUser],
+ gr_membership_manager: MembershipManager,
+ gr_user_manager: GRUserManager,
):
- assert team.gr_users is None
+ assert gr_team.gr_users is None
- team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config)
- assert isinstance(team.gr_users, list)
- assert len(team.gr_users) == 0
+ gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager)
+ assert isinstance(gr_team.gr_users, list)
+ assert len(gr_team.gr_users) == 0
# Create a new Membership
- membership_manager.create(team=team, gr_user=gr_user_factory())
- assert len(team.gr_users) == 0
- team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config)
- assert len(team.gr_users) == 1
+ gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory())
+ assert len(gr_team.gr_users) == 0
+ gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager)
+ assert len(gr_team.gr_users) == 1
# Create another Membership
- membership_manager.create(team=team, gr_user=gr_user_factory())
- assert len(team.gr_users) == 1
- team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config)
- assert len(team.gr_users) == 2
+ gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory())
+ assert len(gr_team.gr_users) == 1
+ gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager)
+ assert len(gr_team.gr_users) == 2
- def test_businesses(self, team, business, team_manager, gr_db, gr_redis_config):
- from generalresearch.models.gr.business import Business
+ def test_businesses(
+ self,
+ gr_team: Team,
+ gr_business: Business,
+ team_manager: TeamManager,
+ gr_business_manager: BusinessManager,
+ ):
- assert team.businesses is None
+ assert gr_team.businesses is None
- team.prefetch_businesses(pg_config=gr_db, redis_config=gr_redis_config)
- assert isinstance(team.businesses, list)
- assert len(team.businesses) == 0
+ gr_team.prefetch_businesses(gr_business_manager=gr_business_manager)
+ assert isinstance(gr_team.businesses, list)
+ assert len(gr_team.businesses) == 0
- team_manager.add_business(team=team, business=business)
- assert len(team.businesses) == 0
- team.prefetch_businesses(pg_config=gr_db, redis_config=gr_redis_config)
- assert len(team.businesses) == 1
- assert isinstance(team.businesses[0], Business)
- assert team.businesses[0].uuid == business.uuid
+ team_manager.add_business(team=gr_team, business=gr_business)
+ assert len(gr_team.businesses) == 0
+ gr_team.prefetch_businesses(gr_business_manager=gr_business_manager)
+ assert len(gr_team.businesses) == 1
+ assert isinstance(gr_team.businesses[0], Business)
+ assert gr_team.businesses[0].uuid == gr_business.uuid
- def test_products(self, team, product_factory, thl_web_rr):
- from generalresearch.models.thl.product import Product
+ def test_products(
+ self,
+ gr_team: Team,
+ product_factory: Callable[..., Product],
+ thl_web_rr: PostgresConfig,
+ product_manager: ProductManager,
+ ):
- assert team.products is None
+ assert gr_team.products is None
- team.prefetch_products(thl_pg_config=thl_web_rr)
- assert isinstance(team.products, list)
- assert len(team.products) == 0
+ gr_team.prefetch_products(product_manager=product_manager)
+ assert isinstance(gr_team.products, list)
+ assert len(gr_team.products) == 0
- product_factory(team=team)
- assert len(team.products) == 0
- team.prefetch_products(thl_pg_config=thl_web_rr)
- assert len(team.products) == 1
- assert isinstance(team.products[0], Product)
+ product_factory(team=gr_team)
+ assert len(gr_team.products) == 0
+ gr_team.prefetch_products(product_manager=product_manager)
+ assert len(gr_team.products) == 1
+ assert isinstance(gr_team.products[0], Product)
class TestTeamMethods:
-
- def test_cache_key(self, team, gr_redis):
- assert isinstance(team.cache_key, str)
- assert ":" in team.cache_key
- assert str(team.uuid) in team.cache_key
+ def test_cache_key(self, gr_team: Team):
+ assert isinstance(gr_team.cache_key, str)
+ assert ":" in gr_team.cache_key
+ assert str(gr_team.uuid) in gr_team.cache_key
def test_set_cache(
self,
- team,
- gr_redis,
- gr_db,
- thl_web_rr,
- gr_redis_config,
- client_no_amm,
- mnt_filepath,
- mnt_gr_api_dir,
- enriched_wall_merge,
- enriched_session_merge,
+ gr_team: Team,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ enriched_wall_merge: EnrichedWallMerge,
+ enriched_session_merge: EnrichedSessionMerge,
+ product_manager: ProductManager,
+ gr_user_manager: GRUserManager,
+ gr_business_manager: BusinessManager,
+ gr_membership_manager: MembershipManager,
):
- assert gr_redis.get(name=team.cache_key) is None
-
- team.set_cache(
- pg_config=gr_db,
- thl_web_rr=thl_web_rr,
+ client = gr_redis_config.create_redis_client()
+ assert client.get(name=gr_team.cache_key) is None
+
+ gr_team.set_cache(
+ product_manager=product_manager,
+ gr_user_manager=gr_user_manager,
+ gr_business_manager=gr_business_manager,
+ gr_membership_manager=gr_membership_manager,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -130,33 +183,36 @@ class TestTeamMethods:
enriched_session=enriched_session_merge,
)
- assert gr_redis.hgetall(name=team.cache_key) is not None
+ assert client.hgetall(name=gr_team.cache_key) is not None
def test_set_cache_team(
self,
- gr_user,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- gr_redis_config,
- client_no_amm,
- mnt_filepath,
- mnt_gr_api_dir,
- enriched_wall_merge,
- enriched_session_merge,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership_factory: Callable[..., Membership],
+ gr_redis_config: RedisConfig,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ enriched_wall_merge: EnrichedWallMerge,
+ enriched_session_merge: EnrichedSessionMerge,
+ product_manager: ProductManager,
+ gr_user_manager: GRUserManager,
+ gr_business_manager: BusinessManager,
+ gr_membership_manager: MembershipManager,
):
from generalresearch.models.gr.team import Team
- p1 = product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ p1 = product_factory(team=gr_team)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
- team.set_cache(
- pg_config=gr_db,
- thl_web_rr=thl_web_rr,
+ gr_team.set_cache(
+ product_manager=product_manager,
+ gr_user_manager=gr_user_manager,
+ gr_business_manager=gr_business_manager,
+ gr_membership_manager=gr_membership_manager,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -166,46 +222,47 @@ class TestTeamMethods:
)
team2 = Team.from_redis(
- uuid=team.uuid,
+ uuid=gr_team.uuid,
fields=["id", "memberships", "gr_users", "businesses", "products"],
gr_redis_config=gr_redis_config,
)
- assert team.model_dump_json() == team2.model_dump_json()
+ assert isinstance(team2, Team)
+ assert isinstance(team2.products, list)
+ assert isinstance(team2.gr_users, list)
+ assert gr_team.model_dump_json() == team2.model_dump_json()
assert p1.uuid in [p.uuid for p in team2.products]
assert len(team2.gr_users) == 1
assert gr_user.id in [gru.id for gru in team2.gr_users]
def test_prebuild_enriched_session_parquet(
self,
- event_report_request,
- enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
- team,
+ enriched_session_merge: EnrichedSessionMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ thl_web_rr: PostgresConfig,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ gr_team: Team,
+ product_manager: ProductManager,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
- p1 = product_factory(team=team)
- p2 = product_factory(team=team)
+ p1 = product_factory(team=gr_team)
+ p2 = product_factory(team=gr_team)
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -221,8 +278,8 @@ class TestTeamMethods:
pg_config=thl_web_rr,
)
- team.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ gr_team.prebuild_enriched_session_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -231,41 +288,38 @@ class TestTeamMethods:
# Now try to read from path
df = pd.read_parquet(
- os.path.join(mnt_gr_api_dir, "pop_session", f"{team.file_key}.parquet")
+ os.path.join(mnt_gr_api_dir, "pop_session", f"{gr_team.file_key}.parquet")
)
assert isinstance(df, pd.DataFrame)
def test_prebuild_enriched_wall_parquet(
self,
- event_report_request,
- enriched_session_merge,
- enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
- team,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: EnrichedSessionMerge,
+ thl_web_rr: PostgresConfig,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ gr_team: Team,
+ product_manager: ProductManager,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
- p1 = product_factory(team=team)
- p2 = product_factory(team=team)
+ p1 = product_factory(team=gr_team)
+ p2 = product_factory(team=gr_team)
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -281,8 +335,8 @@ class TestTeamMethods:
pg_config=thl_web_rr,
)
- team.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ gr_team.prebuild_enriched_wall_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -291,6 +345,6 @@ class TestTeamMethods:
# Now try to read from path
df = pd.read_parquet(
- os.path.join(mnt_gr_api_dir, "pop_event", f"{team.file_key}.parquet")
+ os.path.join(mnt_gr_api_dir, "pop_event", f"{gr_team.file_key}.parquet")
)
assert isinstance(df, pd.DataFrame)
diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py
index 330f919..ea2fc8c 100644
--- a/tests/models/innovate/test_question.py
+++ b/tests/models/innovate/test_question.py
@@ -1,15 +1,17 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from generalresearch.models.definitions import Source
from generalresearch.models.innovate.question import (
InnovateQuestion,
- InnovateQuestionType,
InnovateQuestionOption,
+ InnovateQuestionType,
)
from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestionSelectorTE,
UpkQuestion,
+ UpkQuestionChoice,
UpkQuestionSelectorMC,
+ UpkQuestionSelectorTE,
UpkQuestionType,
- UpkQuestionChoice,
)
diff --git a/tests/models/legacy/test_offerwall_parse_response.py b/tests/models/legacy/test_offerwall_parse_response.py
index b1c96ad..93f5c26 100644
--- a/tests/models/legacy/test_offerwall_parse_response.py
+++ b/tests/models/legacy/test_offerwall_parse_response.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import json
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.legacy.bucket import (
BucketTask,
DurationSummary,
diff --git a/tests/models/legacy/test_profiling_questions.py b/tests/models/legacy/test_profiling_questions.py
index 1afaa6b..6f781ae 100644
--- a/tests/models/legacy/test_profiling_questions.py
+++ b/tests/models/legacy/test_profiling_questions.py
@@ -1,7 +1,11 @@
+from __future__ import annotations
+
+from generalresearch.models.legacy.questions import UpkQuestionResponse
+
+
class TestUpkQuestionResponse:
def test_init(self):
- from generalresearch.models.legacy.questions import UpkQuestionResponse
s = (
'{"status": "success", "count": 7, "questions": [{"selector": "SL", "validation": {"patterns": [{'
diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py
index 224334a..f14c1a7 100644
--- a/tests/models/legacy/test_user_question_answer_in.py
+++ b/tests/models/legacy/test_user_question_answer_in.py
@@ -1,9 +1,25 @@
+from __future__ import annotations
+
import json
+from collections.abc import Callable
+from datetime import datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from generalresearch.models.definitions import Source
+from generalresearch.models.legacy.questions import (
+ UserQuestionAnswers,
+)
+from generalresearch.models.thl.session import Session, Wall
+from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
+ from generalresearch.models.thl.product import Product
+
class TestUserQuestionAnswers:
"""This is for the GRS POST submission that may contain multiple
@@ -15,21 +31,11 @@ class TestUserQuestionAnswers:
def test_json_init(
self,
- product_manager,
- user_manager,
- session_manager,
- wall_manager,
- user_factory,
- product,
- session_factory,
- utc_hour_ago,
+ user_factory: Callable[..., User],
+ product: Product,
+ session_factory: Callable[..., Session],
+ utc_hour_ago: datetime,
):
- from generalresearch.models import Source
- from generalresearch.models.legacy.questions import (
- UserQuestionAnswers,
- )
- from generalresearch.models.thl.session import Session, Wall
- from generalresearch.models.thl.user import User
u: User = user_factory(product=product)
@@ -60,11 +66,8 @@ class TestUserQuestionAnswers:
assert isinstance(instance, UserQuestionAnswers)
def test_simple_validation_errors(
- self, product_manager, user_manager, session_manager, wall_manager
+ self,
):
- from generalresearch.models.legacy.questions import (
- UserQuestionAnswers,
- )
with pytest.raises(ValueError):
UserQuestionAnswers.model_validate(
@@ -114,7 +117,7 @@ class TestUserQuestionAnswers:
with pytest.raises(ValueError):
answers = [
- {"question_id": uuid4().hex, "answer": ["a"]} for i in range(101)
+ {"question_id": uuid4().hex, "answer": ["a"]} for _ in range(101)
]
UserQuestionAnswers.model_validate(
{
@@ -139,9 +142,6 @@ class TestUserQuestionAnswers:
# TODO: depending on if or how many of these types of errors actually
# occur, we could get fancy and just drop one of them. I don't
# think this is worth exploring yet unless we see if it's a problem.
- from generalresearch.models.legacy.questions import (
- UserQuestionAnswers,
- )
consistent_qid = uuid4().hex
with pytest.raises(ValueError) as cm:
@@ -161,11 +161,11 @@ class TestUserQuestionAnswers:
def test_allow_answer_failures_silent(
self,
- user_manager,
- product,
- user_factory,
- utc_hour_ago,
- session_factory,
+ user_manager: UserManager,
+ product: Product,
+ user_factory: Callable[..., User],
+ utc_hour_ago: datetime,
+ session_factory: Callable[..., Session],
):
"""
There are many instances where suppliers may be submitting answers
@@ -173,11 +173,6 @@ class TestUserQuestionAnswers:
that one QuestionAnswerIn without "loosing" any of the other
QuestionAnswerIn items that they provided.
"""
- from generalresearch.models.legacy.questions import (
- UserQuestionAnswers,
- )
- from generalresearch.models.thl.session import Session, Wall
- from generalresearch.models.thl.user import User
u: User = user_factory(product=product)
@@ -263,12 +258,12 @@ class TestUserQuestionAnswerIn:
UserQuestionAnswerIn,
)
- for qid in {
+ for qid in (
"2fbedb2b9f7647b09ff5e52fa119cc5e",
"4030c52371b04e80b64e058d9c5b82e9",
"a91cb1dea814480dba12d9b7b48696dd",
"1d1e2e8380ac474b87fb4e4c569b48df",
- }:
+ ):
# This is the UserAgent question which only allows a single answer
with pytest.raises(ValueError) as cm:
UserQuestionAnswerIn.model_validate(
@@ -282,7 +277,7 @@ class TestUserQuestionAnswerIn:
UserQuestionAnswerIn,
)
- answer = [uuid4().hex[:6] for i in range(11)]
+ answer = [uuid4().hex[:6] for _ in range(11)]
with pytest.raises(ValueError) as cm:
UserQuestionAnswerIn.model_validate(
{"question_id": uuid4().hex, "answer": answer}
@@ -294,8 +289,8 @@ class TestUserQuestionAnswerIn:
UserQuestionAnswerIn,
)
- answer = ["aaa" for i in range(5)]
- with pytest.raises(ValueError) as cm:
+ answer = ["aaa" for _ in range(5)]
+ with pytest.raises(ValueError):
UserQuestionAnswerIn.model_validate(
{"question_id": uuid4().hex, "answer": answer}
)
diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py
index bedf9c2..c1141fb 100644
--- a/tests/models/morning/test.py
+++ b/tests/models/morning/test.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from generalresearch.models.morning.question import MorningQuestion
@@ -163,8 +165,8 @@ bid = {
# what gets run in MorningAPI._format_bid
bid["language_isos"] = ("eng",)
bid["country_iso"] = "us"
-bid["end_date"] = datetime(2024, 7, 19, 9, 1, 13, 520243, tzinfo=timezone.utc)
-bid["published_at"] = datetime(2024, 6, 19, 9, 1, 13, 520243, tzinfo=timezone.utc)
+bid["end_date"] = datetime(2024, 7, 19, 9, 1, 13, 520243, tzinfo=UTC)
+bid["published_at"] = datetime(2024, 6, 19, 9, 1, 13, 520243, tzinfo=UTC)
bid.update(bid["statistics"])
bid["qualified_conversion"] /= 100
bid["system_conversion"] /= 100
diff --git a/tests/models/network/__init__.py b/tests/models/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/tests/models/network/__init__.py
+++ /dev/null
diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py
deleted file mode 100644
index 2965300..0000000
--- a/tests/models/network/test_mtr.py
+++ /dev/null
@@ -1,26 +0,0 @@
-from generalresearch.models.network.mtr.execute import execute_mtr
-import faker
-
-from generalresearch.models.network.tool_run import ToolName, ToolClass
-
-fake = faker.Faker()
-
-
-def test_execute_mtr(toolrun_manager):
- ip = "65.19.129.53"
-
- run = execute_mtr(ip=ip, report_cycles=3)
- assert run.tool_name == ToolName.MTR
- assert run.tool_class == ToolClass.TRACEROUTE
- assert run.ip == ip
- result = run.parsed
-
- last_hop = result.hops[-1]
- assert last_hop.asn == 6939
- assert last_hop.domain == "grlengine.com"
-
- last_hop_1 = result.hops[-2]
- assert last_hop_1.asn == 6939
- assert last_hop_1.domain == "he.net"
-
- toolrun_manager.create_mtr_run(run)
diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py
deleted file mode 100644
index a135a13..0000000
--- a/tests/models/network/test_nmap.py
+++ /dev/null
@@ -1,30 +0,0 @@
-import subprocess
-
-import faker
-
-from generalresearch.managers.network.tool_run import ToolRunManager
-from generalresearch.models.network.definitions import IPProtocol
-from generalresearch.models.network.nmap.execute import execute_nmap
-from generalresearch.models.network.nmap.result import PortState
-from generalresearch.models.network.tool_run import ToolClass, ToolName
-
-fake = faker.Faker()
-
-
-def resolve(host: str):
- return subprocess.check_output(["dig", host, "+short"]).decode().strip()
-
-
-def test_execute_nmap_scanme(toolrun_manager: ToolRunManager):
- ip = resolve("scanme.nmap.org")
-
- run = execute_nmap(ip=ip, top_ports=None, ports="20-30", enable_advanced=False)
- assert run.tool_name == ToolName.NMAP
- assert run.tool_class == ToolClass.PORT_SCAN
- assert run.ip == ip
- result = run.parsed
-
- port22 = result._port_index[(IPProtocol.TCP, 22)]
- assert port22.state == PortState.OPEN
-
- toolrun_manager.create_nmap_run(run)
diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py
deleted file mode 100644
index abc83c9..0000000
--- a/tests/models/network/test_nmap_parser.py
+++ /dev/null
@@ -1,22 +0,0 @@
-import os
-
-import pytest
-
-from generalresearch.models.network.nmap.parser import parse_nmap_xml
-
-@pytest.fixture
-def nmap_raw_output_2(request) -> str:
- fp = os.path.join(request.config.rootpath, "data/nmaprun2.xml")
- with open(fp) as f:
- data = f.read()
- return data
-
-
-def test_nmap_xml_parser(nmap_raw_output, nmap_raw_output_2):
- n = parse_nmap_xml(nmap_raw_output)
- assert n.tcp_open_ports == [61232]
- assert len(n.trace.hops) == 18
-
- n = parse_nmap_xml(nmap_raw_output_2)
- assert n.tcp_open_ports == [22, 80, 9929, 31337]
- assert n.trace is None
diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py
deleted file mode 100644
index 5c3b024..0000000
--- a/tests/models/network/test_rdns.py
+++ /dev/null
@@ -1,34 +0,0 @@
-import faker
-
-from generalresearch.managers.network.tool_run import ToolRunManager
-from generalresearch.models.network.rdns.execute import execute_rdns
-from generalresearch.models.network.tool_run import ToolClass, ToolName
-
-fake = faker.Faker()
-
-
-def test_execute_rdns_grl(toolrun_manager: ToolRunManager):
- ip = "65.19.129.53"
- run = execute_rdns(ip=ip)
- assert run.tool_name == ToolName.DIG
- assert run.tool_class == ToolClass.RDNS
- assert run.ip == ip
- result = run.parsed
- assert result.primary_hostname == "in1-smtp.grlengine.com"
- assert result.primary_domain == "grlengine.com"
- assert result.hostname_count == 1
-
- toolrun_manager.create_rdns_run(run)
-
-
-def test_execute_rdns_none(toolrun_manager: ToolRunManager):
- ip = fake.ipv6()
- run = execute_rdns(ip)
- result = run.parsed
-
- assert result.primary_hostname is None
- assert result.primary_domain is None
- assert result.hostname_count == 0
- assert result.hostnames == []
-
- toolrun_manager.create_rdns_run(run)
diff --git a/tests/models/precision/__init__.py b/tests/models/precision/__init__.py
index 8006fa3..e69de29 100644
--- a/tests/models/precision/__init__.py
+++ b/tests/models/precision/__init__.py
@@ -1,115 +0,0 @@
-survey_json = {
- "cpi": "1.44",
- "country_isos": "ca",
- "language_isos": "eng",
- "country_iso": "ca",
- "language_iso": "eng",
- "buyer_id": "7047",
- "bid_loi": 1200,
- "bid_ir": 0.45,
- "source": "e",
- "used_question_ids": ["age", "country_iso", "gender", "gender_1"],
- "survey_id": "0000",
- "group_id": "633473",
- "status": "open",
- "name": "beauty survey",
- "survey_guid": "c7f375c5077d4c6c8209ff0b539d7183",
- "category_id": "-1",
- "global_conversion": None,
- "desired_count": 96,
- "achieved_count": 0,
- "allowed_devices": "1,2,3",
- "entry_link": "https://www.opinionetwork.com/survey/entry.aspx?mid=[%MID%]&project=633473&key=%%key%%",
- "excluded_surveys": "470358,633286",
- "quotas": [
- {
- "name": "25-34,Male,Quebec",
- "id": "2324110",
- "guid": "23b5760d24994bc08de451b3e62e77c7",
- "status": "open",
- "desired_count": 48,
- "achieved_count": 0,
- "termination_count": 0,
- "overquota_count": 0,
- "condition_hashes": ["b41e1a3", "bc89ee8", "4124366", "9f32c61"],
- },
- {
- "name": "25-34,Female,Quebec",
- "id": "2324111",
- "guid": "0706f1a88d7e4f11ad847c03012e68d2",
- "status": "open",
- "desired_count": 48,
- "achieved_count": 0,
- "termination_count": 4,
- "overquota_count": 0,
- "condition_hashes": ["b41e1a3", "0cdc304", "500af2c", "9f32c61"],
- },
- ],
- "conditions": {
- "b41e1a3": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "country_iso",
- "values": ["ca"],
- "criterion_hash": "b41e1a3",
- "value_len": 1,
- "sizeof": 2,
- },
- "bc89ee8": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "gender",
- "values": ["male"],
- "criterion_hash": "bc89ee8",
- "value_len": 1,
- "sizeof": 4,
- },
- "4124366": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "gender_1",
- "values": ["male"],
- "criterion_hash": "4124366",
- "value_len": 1,
- "sizeof": 4,
- },
- "9f32c61": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "age",
- "values": ["25", "26", "27", "28", "29", "30", "31", "32", "33", "34"],
- "criterion_hash": "9f32c61",
- "value_len": 10,
- "sizeof": 20,
- },
- "0cdc304": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "gender",
- "values": ["female"],
- "criterion_hash": "0cdc304",
- "value_len": 1,
- "sizeof": 6,
- },
- "500af2c": {
- "logical_operator": "OR",
- "value_type": 1,
- "negate": False,
- "question_id": "gender_1",
- "values": ["female"],
- "criterion_hash": "500af2c",
- "value_len": 1,
- "sizeof": 6,
- },
- },
- "expected_end_date": "2024-06-28T10:40:33.000000Z",
- "created": None,
- "updated": None,
- "is_live": True,
- "all_hashes": ["0cdc304", "b41e1a3", "9f32c61", "bc89ee8", "4124366", "500af2c"],
-}
diff --git a/tests/models/precision/test_survey.py b/tests/models/precision/test_survey.py
index ff2d6d1..4d671f2 100644
--- a/tests/models/precision/test_survey.py
+++ b/tests/models/precision/test_survey.py
@@ -1,10 +1,15 @@
-class TestPrecisionQuota:
+from __future__ import annotations
+
+from typing import Any
+
+from generalresearch.models.precision import PrecisionStatus
+from generalresearch.models.precision.survey import PrecisionSurvey
- def test_quota_passes(self):
- from generalresearch.models.precision.survey import PrecisionSurvey
- from tests.models.precision import survey_json
- s = PrecisionSurvey.model_validate(survey_json)
+class TestPrecisionQuota:
+
+ def test_quota_passes(self, precision_survey_json: dict[str, Any]):
+ s = PrecisionSurvey.model_validate(precision_survey_json)
q = s.quotas[0]
ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]}
assert q.matches(ce)
@@ -16,12 +21,9 @@ class TestPrecisionQuota:
assert not q.matches(ce)
assert not q.matches({})
- def test_quota_passes_closed(self):
- from generalresearch.models.precision import PrecisionStatus
- from generalresearch.models.precision.survey import PrecisionSurvey
- from tests.models.precision import survey_json
+ def test_quota_passes_closed(self, precision_survey_json: dict[str, Any]):
- s = PrecisionSurvey.model_validate(survey_json)
+ s = PrecisionSurvey.model_validate(precision_survey_json)
q = s.quotas[0]
q.status = PrecisionStatus.CLOSED
ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]}
@@ -32,20 +34,15 @@ class TestPrecisionQuota:
class TestPrecisionSurvey:
- def test_passes(self):
- from generalresearch.models.precision.survey import PrecisionSurvey
- from tests.models.precision import survey_json
+ def test_passes(self, precision_survey_json: dict[str, Any]):
- s = PrecisionSurvey.model_validate(survey_json)
+ s = PrecisionSurvey.model_validate(precision_survey_json)
ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]}
assert s.determine_eligibility(ce)
- def test_elig_closed_quota(self):
- from generalresearch.models.precision import PrecisionStatus
- from generalresearch.models.precision.survey import PrecisionSurvey
- from tests.models.precision import survey_json
+ def test_elig_closed_quota(self, precision_survey_json: dict[str, Any]):
- s = PrecisionSurvey.model_validate(survey_json)
+ s = PrecisionSurvey.model_validate(precision_survey_json)
ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]}
q = s.quotas[0]
q.status = PrecisionStatus.CLOSED
@@ -57,12 +54,9 @@ class TestPrecisionSurvey:
# Now me match an open quota and dont match the closed quota, so we should be eligible
assert s.determine_eligibility(ce)
- def test_passes_sp(self):
- from generalresearch.models.precision import PrecisionStatus
- from generalresearch.models.precision.survey import PrecisionSurvey
- from tests.models.precision import survey_json
+ def test_passes_sp(self, precision_survey_json: dict[str, Any]):
- s = PrecisionSurvey.model_validate(survey_json)
+ s = PrecisionSurvey.model_validate(precision_survey_json)
ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]}
passes, hashes = s.determine_eligibility_soft(ce)
diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py
index 68d7838..10ce884 100644
--- a/tests/models/prodege/test_survey_participation.py
+++ b/tests/models/prodege/test_survey_participation.py
@@ -1,16 +1,19 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
+
+from generalresearch.models.prodege import ProdegePastParticipationType
+from generalresearch.models.prodege.survey import (
+ ProdegePastParticipation,
+ ProdegeUserPastParticipation,
+)
class TestProdegeParticipation:
def test_exclude(self):
- from generalresearch.models.prodege import ProdegePastParticipationType
- from generalresearch.models.prodege.survey import (
- ProdegePastParticipation,
- ProdegeUserPastParticipation,
- )
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
pp = ProdegePastParticipation.from_api(
{
"participation_project_ids": [152677146, 152803285],
@@ -84,12 +87,8 @@ class TestProdegeParticipation:
assert not pp.is_eligible(upps)
def test_include(self):
- from generalresearch.models.prodege.survey import (
- ProdegePastParticipation,
- ProdegeUserPastParticipation,
- )
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
pp = ProdegePastParticipation.from_api(
{
"participation_project_ids": [152677146, 152803285],
diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py
index ba118d7..d469530 100644
--- a/tests/models/spectrum/test_question.py
+++ b/tests/models/spectrum/test_question.py
@@ -1,17 +1,19 @@
-from datetime import datetime, timezone
+from __future__ import annotations
-from generalresearch.models import Source
+from datetime import UTC, datetime
+
+from generalresearch.models.definitions import Source
from generalresearch.models.spectrum.question import (
- SpectrumQuestionOption,
SpectrumQuestion,
- SpectrumQuestionType,
SpectrumQuestionClass,
+ SpectrumQuestionOption,
+ SpectrumQuestionType,
)
from generalresearch.models.thl.profiling.upk_question import (
UpkQuestion,
+ UpkQuestionChoice,
UpkQuestionSelectorMC,
UpkQuestionType,
- UpkQuestionChoice,
)
@@ -32,6 +34,7 @@ class TestSpectrumQuestion:
"mod_on": 1706557247467,
}
q = SpectrumQuestion.from_api(example_1, "us", "eng")
+ assert isinstance(q, SpectrumQuestion)
expected_q = SpectrumQuestion(
question_id="213",
@@ -43,7 +46,7 @@ class TestSpectrumQuestion:
tags=None,
options=None,
class_num=SpectrumQuestionClass.CORE,
- created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc),
+ created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC),
is_live=True,
source=Source.SPECTRUM,
category_id=None,
@@ -72,6 +75,8 @@ class TestSpectrumQuestion:
"mod_on": 1706557249817,
}
q = SpectrumQuestion.from_api(example_2, "us", "eng")
+ assert isinstance(q, SpectrumQuestion)
+
expected_q = SpectrumQuestion(
question_id="211",
country_iso="us",
@@ -85,7 +90,7 @@ class TestSpectrumQuestion:
SpectrumQuestionOption(id="112", text="Female", order=1),
],
class_num=SpectrumQuestionClass.CORE,
- created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc),
+ created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC),
is_live=True,
source=Source.SPECTRUM,
category_id=None,
@@ -160,7 +165,7 @@ class TestSpectrumQuestion:
SpectrumQuestionOption(id="999", text="None of the above", order=3),
],
class_num=SpectrumQuestionClass.EXTENDED,
- created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=timezone.utc),
+ created=datetime(2017, 8, 16, 7, 52, 7, 688000, tzinfo=UTC),
is_live=True,
source=Source.SPECTRUM,
category_id=None,
diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py
index b612a63..02c5d3f 100644
--- a/tests/models/spectrum/test_survey.py
+++ b/tests/models/spectrum/test_survey.py
@@ -1,15 +1,25 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from decimal import Decimal
+from generalresearch.models.definitions import (
+ LogicalOperator,
+ Source,
+ TaskCalculationType,
+)
+from generalresearch.models.spectrum import SpectrumStatus
+from generalresearch.models.spectrum.survey import (
+ SpectrumCondition,
+ SpectrumQuota,
+ SpectrumSurvey,
+)
+from generalresearch.models.thl.survey.condition import ConditionValueType
+
class TestSpectrumCondition:
def test_condition_create(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.spectrum.survey import (
- SpectrumCondition,
- )
- from generalresearch.models.thl.survey.condition import ConditionValueType
c = SpectrumCondition.from_api(
{
@@ -64,10 +74,6 @@ class TestSpectrumCondition:
class TestSpectrumQuota:
def test_quota_create(self):
- from generalresearch.models.spectrum.survey import (
- SpectrumCondition,
- SpectrumQuota,
- )
d = {
"quota_id": "a846b545-4449-4d76-93a2-f8ebdf6e711e",
@@ -84,9 +90,6 @@ class TestSpectrumQuota:
assert q.is_open
def test_quota_passes(self):
- from generalresearch.models.spectrum.survey import (
- SpectrumQuota,
- )
q = SpectrumQuota(remaining_count=57, condition_hashes=["a"])
assert q.passes({"a": True})
@@ -103,9 +106,6 @@ class TestSpectrumQuota:
assert not q.passes({"a": True})
def test_quota_passes_soft(self):
- from generalresearch.models.spectrum.survey import (
- SpectrumQuota,
- )
q = SpectrumQuota(remaining_count=57, condition_hashes=["a", "b", "c"])
# Pass if we match all
@@ -122,29 +122,17 @@ class TestSpectrumQuota:
class TestSpectrumSurvey:
def test_survey_create(self):
- from generalresearch.models import (
- LogicalOperator,
- Source,
- TaskCalculationType,
- )
- from generalresearch.models.spectrum import SpectrumStatus
- from generalresearch.models.spectrum.survey import (
- SpectrumCondition,
- SpectrumQuota,
- SpectrumSurvey,
- )
- from generalresearch.models.thl.survey.condition import ConditionValueType
# Note: d is the raw response after calling SpectrumAPI.preprocess_survey() on it!
d = {
"survey_id": 29333264,
"survey_name": "Exciting New Survey #29333264",
"survey_status": 22,
- "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc),
+ "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC),
"category": "Exciting New",
"category_code": 232,
- "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc),
- "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc),
+ "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC),
+ "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC),
"soft_launch": False,
"click_balancing": 0,
"price_type": 1,
@@ -202,6 +190,8 @@ class TestSpectrumSurvey:
"exclusion_period": 0,
}
s = SpectrumSurvey.from_api(d)
+ assert isinstance(s, SpectrumSurvey)
+
expected_survey = SpectrumSurvey(
cpi=Decimal("1.20000"),
country_isos=["fr"],
@@ -212,7 +202,7 @@ class TestSpectrumSurvey:
survey_id="29333264",
survey_name="Exciting New Survey #29333264",
status=SpectrumStatus.LIVE,
- field_end_date=datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc),
+ field_end_date=datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC),
category_code="232",
calculation_type=TaskCalculationType.COMPLETES,
requires_pii=False,
@@ -240,8 +230,8 @@ class TestSpectrumSurvey:
values=["18-64"],
)
},
- created_api=datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc),
- modified_api=datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc),
+ created_api=datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC),
+ modified_api=datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC),
updated=None,
)
assert expected_survey.model_dump_json() == s.model_dump_json()
@@ -255,11 +245,11 @@ class TestSpectrumSurvey:
"survey_id": 29333264,
"survey_name": "#29333264",
"survey_status": 22,
- "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc),
+ "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC),
"category": "Exciting New",
"category_code": 232,
- "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc),
- "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc),
+ "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC),
+ "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC),
"soft_launch": False,
"click_balancing": 0,
"price_type": 1,
@@ -303,6 +293,8 @@ class TestSpectrumSurvey:
"exclusion_period": 0,
}
s = SpectrumSurvey.from_api(d)
+ assert isinstance(s, SpectrumSurvey)
+
assert {"212", "1202", "214"} == s.used_question_ids
assert s.is_live
assert s.is_open
@@ -318,11 +310,11 @@ class TestSpectrumSurvey:
"survey_id": 29333264,
"survey_name": "#29333264",
"survey_status": 22,
- "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc),
+ "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=UTC),
"category": "Exciting New",
"category_code": 232,
- "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc),
- "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc),
+ "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=UTC),
+ "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC),
"soft_launch": False,
"click_balancing": 0,
"price_type": 1,
@@ -345,6 +337,8 @@ class TestSpectrumSurvey:
"exclusion_period": 0,
}
s = SpectrumSurvey.from_api(d)
+ assert isinstance(s, SpectrumSurvey)
+
s.qualifications = ["a", "b", "c"]
s.quotas = [
SpectrumQuota(remaining_count=10, condition_hashes=["a", "b"]),
@@ -411,3 +405,15 @@ class TestSpectrumSurvey:
assert (None, {"c", "d"}) == s.determine_eligibility_soft(
{"a": True, "b": True, "c": None, "d": None}
)
+
+
+def test_spectrum_something(
+ spectrum_conditions: list[SpectrumCondition], spectrum_api_surveys_json: list[str]
+):
+
+ c1 = spectrum_conditions[0]
+ c3 = spectrum_conditions[2]
+
+ survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0])
+ assert c1.criterion_hash in survey.qualifications
+ assert c3.criterion_hash in survey.qualifications
diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py
index 582093c..0300956 100644
--- a/tests/models/spectrum/test_survey_manager.py
+++ b/tests/models/spectrum/test_survey_manager.py
@@ -1,72 +1,36 @@
-import copy
+from __future__ import annotations
+
import logging
-from datetime import timezone, datetime
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING, Any
from pymysql import IntegrityError
+from generalresearch.config import is_debug
-logger = logging.getLogger()
+if TYPE_CHECKING:
+ from generalresearch.managers.spectrum.survey import (
+ SpectrumSurveyManager,
+ )
+ from generalresearch.sql_helper import SqlHelper
-example_survey_api_response = {
- "survey_id": 29333264,
- "survey_name": "#29333264",
- "survey_status": 22,
- "field_end_date": datetime(2024, 5, 23, 18, 18, 31, tzinfo=timezone.utc),
- "category": "Exciting New",
- "category_code": 232,
- "crtd_on": datetime(2024, 5, 20, 17, 48, 13, tzinfo=timezone.utc),
- "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=timezone.utc),
- "soft_launch": False,
- "click_balancing": 0,
- "price_type": 1,
- "pii": False,
- "buyer_message": "",
- "buyer_id": 4726,
- "incl_excl": 0,
- "cpi": Decimal("1.20"),
- "last_complete_date": None,
- "project_last_complete_date": None,
- "quotas": [
- {
- "quota_id": "c2bc961e-4f26-4223-b409-ebe9165cfdf5",
- "quantities": {"currently_open": 491, "remaining": 495, "achieved": 0},
- "criteria": [
- {
- "qualification_code": 214,
- "range_sets": [{"units": 311, "to": 64, "from": 18}],
- }
- ],
- }
- ],
- "qualifications": [
- {
- "range_sets": [{"units": 311, "to": 64, "from": 18}],
- "qualification_code": 212,
- },
- {"condition_codes": ["111", "117", "112"], "qualification_code": 1202},
- ],
- "country_iso": "fr",
- "language_iso": "fre",
- "bid_ir": 0.4,
- "bid_loi": 600,
- "overall_ir": None,
- "overall_loi": None,
- "last_block_ir": None,
- "last_block_loi": None,
- "survey_exclusions": set(),
- "exclusion_period": 0,
-}
+logger = logging.getLogger()
class TestSpectrumSurvey:
- def test_survey_create(self, settings, spectrum_manager, spectrum_rw):
+ def test_survey_create(
+ self,
+ spectrum_survey_manager: SpectrumSurveyManager,
+ spectrum_rw: SqlHelper,
+ spectrum_api_survey_json: dict[str, Any],
+ ):
from generalresearch.models.spectrum.survey import SpectrumSurvey
- assert settings.debug, "CRITICAL: Do not run this on production."
+ assert is_debug(), "CRITICAL: Do not run this on production."
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
spectrum_rw.execute_sql_query(
query=f"""
DELETE FROM `{spectrum_rw.db}`.spectrum_survey
@@ -74,26 +38,31 @@ class TestSpectrumSurvey:
commit=True,
)
- d = example_survey_api_response.copy()
- s = SpectrumSurvey.from_api(d)
- spectrum_manager.create(s)
+ s = SpectrumSurvey.from_api(spectrum_api_survey_json)
+ assert isinstance(s, SpectrumSurvey)
+ spectrum_survey_manager.create(s)
- surveys = spectrum_manager.get_survey_library(updated_since=now)
+ surveys = spectrum_survey_manager.get_survey_library(updated_since=now)
assert len(surveys) == 1
assert "29333264" == surveys[0].survey_id
assert s.is_unchanged(surveys[0])
try:
- spectrum_manager.create(s)
+ spectrum_survey_manager.create(s)
except IntegrityError as e:
print(e.args)
- def test_survey_update(self, settings, spectrum_manager, spectrum_rw):
+ def test_survey_update(
+ self,
+ spectrum_survey_manager: SpectrumSurveyManager,
+ spectrum_rw: SqlHelper,
+ spectrum_api_survey_json: dict[str, Any],
+ ):
from generalresearch.models.spectrum.survey import SpectrumSurvey
- assert settings.debug, "CRITICAL: Do not run this on production."
+ assert is_debug(), "CRITICAL: Do not run this on production."
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
spectrum_rw.execute_sql_query(
query=f"""
DELETE FROM `{spectrum_rw.db}`.spectrum_survey
@@ -101,14 +70,13 @@ class TestSpectrumSurvey:
""",
commit=True,
)
- d = copy.deepcopy(example_survey_api_response)
- s = SpectrumSurvey.from_api(d)
- print(s)
+ s = SpectrumSurvey.from_api(spectrum_api_survey_json)
+ assert isinstance(s, SpectrumSurvey)
- spectrum_manager.create(s)
+ spectrum_survey_manager.create(s)
s.cpi = Decimal("0.50")
- spectrum_manager.update([s])
- surveys = spectrum_manager.get_survey_library(updated_since=now)
+ spectrum_survey_manager.update([s])
+ surveys = spectrum_survey_manager.get_survey_library(updated_since=now)
assert len(surveys) == 1
assert "29333264" == surveys[0].survey_id
assert Decimal("0.50") == surveys[0].cpi
@@ -123,8 +91,8 @@ class TestSpectrumSurvey:
s.bid_loi = None
s.overall_loi = 1000
s.last_block_loi = 1000
- spectrum_manager.update([s])
- surveys = spectrum_manager.get_survey_library(updated_since=now)
+ spectrum_survey_manager.update([s])
+ surveys = spectrum_survey_manager.get_survey_library(updated_since=now)
assert 600 == surveys[0].bid_loi
assert 1000 == surveys[0].overall_loi
assert 1000 == surveys[0].last_block_loi
diff --git a/tests/models/test_currency.py b/tests/models/test_currency.py
index 40cff88..e946126 100644
--- a/tests/models/test_currency.py
+++ b/tests/models/test_currency.py
@@ -3,27 +3,29 @@ functionality is the same, but pasting here so the tests are in the
correct spot...
"""
+from __future__ import annotations
+
from decimal import Decimal
from random import randint
import pytest
+from generalresearch.currency import USDCent, USDMill, format_usd_cent
+
class TestUSDCentModel:
def test_construct_int(self):
- from generalresearch.currency import USDCent
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
assert int_val == instance
def test_construct_float(self):
- from generalresearch.currency import USDCent
+ float_val: float = 10.6789
with pytest.warns(expected_warning=Warning) as record:
- float_val: float = 10.6789
instance = USDCent(float_val)
assert len(record) == 1
@@ -34,10 +36,9 @@ class TestUSDCentModel:
assert instance == 10
def test_construct_decimal(self):
- from generalresearch.currency import USDCent
+ decimal_val: Decimal = Decimal("10.0")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.0")
instance = USDCent(decimal_val)
assert len(record) == 1
@@ -50,8 +51,8 @@ class TestUSDCentModel:
assert instance == 10
# Now with rounding
+ decimal_val: Decimal = Decimal("10.6789")
with pytest.warns(Warning) as record:
- decimal_val: Decimal = Decimal("10.6789")
instance = USDCent(decimal_val)
assert len(record) == 1
@@ -64,16 +65,12 @@ class TestUSDCentModel:
assert instance == 10
def test_construct_negative(self):
- from generalresearch.currency import USDCent
-
with pytest.raises(expected_exception=ValueError) as cm:
USDCent(-1)
assert "USDCent not be less than zero" in str(cm.value)
def test_operation_add(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -83,9 +80,7 @@ class TestUSDCentModel:
assert int_val1 + int_val2 == instance1 + instance2
def test_operation_subtract(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(500_000, 999_999)
int_val2 = randint(0, 499_999)
@@ -95,21 +90,17 @@ class TestUSDCentModel:
assert int_val1 - int_val2 == instance1 - instance2
def test_operation_subtract_to_neg(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
with pytest.raises(expected_exception=ValueError) as cm:
- instance - USDCent(1_000_000)
+ _ = instance - USDCent(1_000_000)
assert "USDCent not be less than zero" in str(cm.value)
def test_operation_multiply(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -119,15 +110,11 @@ class TestUSDCentModel:
assert int_val1 * int_val2 == instance1 * instance2
def test_operation_div(self):
- from generalresearch.currency import USDCent
-
with pytest.raises(ValueError) as cm:
- USDCent(10) / 2
+ _ = USDCent(10) / 2
assert "Division not allowed for USDCent" in str(cm.value)
def test_operation_result_type(self):
- from generalresearch.currency import USDCent
-
int_val = randint(1, 999_999)
instance = USDCent(int_val)
@@ -141,36 +128,30 @@ class TestUSDCentModel:
assert isinstance(res_multipy, USDCent)
def test_operation_partner_add(self):
- from generalresearch.currency import USDCent
-
int_val = randint(1, 999_999)
instance = USDCent(int_val)
with pytest.raises(expected_exception=AssertionError):
- instance + 0.10
+ _ = instance + 0.10
with pytest.raises(expected_exception=AssertionError):
- instance + Decimal(".10")
+ _ = instance + Decimal(".10")
with pytest.raises(expected_exception=AssertionError):
- instance + "9.9"
+ _ = instance + "9.9"
with pytest.raises(expected_exception=AssertionError):
- instance + True
+ _ = instance + True
def test_abs(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = abs(randint(0, 999_999))
instance = abs(USDCent(int_val))
assert int_val == instance
def test_str(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
@@ -180,8 +161,6 @@ class TestUSDCentModel:
"""There is no correct answer here, but we at least want to make sure
that a USDCent is returned
"""
- from generalresearch.currency import USDCent
-
res = USDCent(10) // 1.2
assert not isinstance(res, USDCent)
assert isinstance(res, float)
@@ -206,18 +185,14 @@ class TestUSDCentModel:
class TestUSDMillModel:
def test_construct_int(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
assert int_val == instance
def test_construct_float(self):
- from generalresearch.currency import USDMill
-
+ float_val: float = 10.6789
with pytest.warns(expected_warning=Warning) as record:
- float_val: float = 10.6789
instance = USDMill(float_val)
assert len(record) == 1
@@ -228,10 +203,8 @@ class TestUSDMillModel:
assert instance == 10
def test_construct_decimal(self):
- from generalresearch.currency import USDMill
-
+ decimal_val: Decimal = Decimal("10.0")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.0")
instance = USDMill(decimal_val)
assert len(record) == 1
@@ -244,10 +217,11 @@ class TestUSDMillModel:
assert instance == 10
# Now with rounding
+ decimal_val: Decimal = Decimal("10.6789")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.6789")
instance = USDMill(decimal_val)
+ assert isinstance(instance, USDMill)
assert len(record) == 1
assert (
"USDMill init with a Decimal. Rounding behavior may be unexpected"
@@ -258,16 +232,12 @@ class TestUSDMillModel:
assert instance == 10
def test_construct_negative(self):
- from generalresearch.currency import USDMill
-
with pytest.raises(expected_exception=ValueError) as cm:
USDMill(-1)
assert "USDMill not be less than zero" in str(cm.value)
def test_operation_add(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -277,9 +247,7 @@ class TestUSDMillModel:
assert int_val1 + int_val2 == instance1 + instance2
def test_operation_subtract(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(500_000, 999_999)
int_val2 = randint(0, 499_999)
@@ -289,21 +257,17 @@ class TestUSDMillModel:
assert int_val1 - int_val2 == instance1 - instance2
def test_operation_subtract_to_neg(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
with pytest.raises(expected_exception=ValueError) as cm:
- instance - USDMill(1_000_000)
+ _ = instance - USDMill(1_000_000)
assert "USDMill not be less than zero" in str(cm.value)
def test_operation_multiply(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -313,15 +277,11 @@ class TestUSDMillModel:
assert int_val1 * int_val2 == instance1 * instance2
def test_operation_div(self):
- from generalresearch.currency import USDMill
-
with pytest.raises(ValueError) as cm:
- USDMill(10) / 2
+ _ = USDMill(10) / 2
assert "Division not allowed for USDMill" in str(cm.value)
def test_operation_result_type(self):
- from generalresearch.currency import USDMill
-
int_val = randint(1, 999_999)
instance = USDMill(int_val)
@@ -335,36 +295,30 @@ class TestUSDMillModel:
assert isinstance(res_multipy, USDMill)
def test_operation_partner_add(self):
- from generalresearch.currency import USDMill
-
int_val = randint(1, 999_999)
instance = USDMill(int_val)
with pytest.raises(expected_exception=AssertionError):
- instance + 0.10
+ _ = instance + 0.10
with pytest.raises(expected_exception=AssertionError):
- instance + Decimal(".10")
+ _ = instance + Decimal(".10")
with pytest.raises(expected_exception=AssertionError):
- instance + "9.9"
+ _ = instance + "9.9"
with pytest.raises(expected_exception=AssertionError):
- instance + True
+ _ = instance + True
def test_abs(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = abs(randint(0, 999_999))
instance = abs(USDMill(int_val))
assert int_val == instance
def test_str(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
@@ -374,8 +328,6 @@ class TestUSDMillModel:
"""There is no correct answer here, but we at least want to make sure
that a USDMill is returned
"""
- from generalresearch.currency import USDCent, USDMill
-
res = USDMill(10) // 1.2
assert not isinstance(res, USDMill)
assert isinstance(res, float)
@@ -400,11 +352,7 @@ class TestUSDMillModel:
class TestNegativeFormatting:
def test_pos(self):
- from generalresearch.currency import format_usd_cent
-
assert "-$987.65" == format_usd_cent(-98765)
def test_neg(self):
- from generalresearch.currency import format_usd_cent
-
assert "-$123.45" == format_usd_cent(-12345)
diff --git a/tests/models/test_device.py b/tests/models/test_device.py
index bf72c81..fdbd906 100644
--- a/tests/models/test_device.py
+++ b/tests/models/test_device.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
iphone_ua_string = (
"Mozilla/5.0 (iPhone; CPU iPhone OS 5_1 like Mac OS X) AppleWebKit/534.46 (KHTML, like Gecko) "
"Version/5.1 Mobile/9B179 Safari/7534.48.3"
@@ -13,10 +15,12 @@ chromebook_ua_string = (
)
+from generalresearch.models.definitions import DeviceType
+from generalresearch.models.device import parse_device_from_useragent
+
+
class TestDeviceUA:
def test_device_ua(self):
- from generalresearch.models import DeviceType
- from generalresearch.models.device import parse_device_from_useragent
assert parse_device_from_useragent(iphone_ua_string) == DeviceType.MOBILE
assert parse_device_from_useragent(ipad_ua_string) == DeviceType.TABLET
diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py
index bd548b3..a1da961 100644
--- a/tests/models/test_finance.py
+++ b/tests/models/test_finance.py
@@ -1,44 +1,43 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from itertools import product as iter_product
from random import randint
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
+import dask.dataframe as dd
import pandas as pd
import pytest
from dask.distributed import Client as DaskClient
# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- client_no_amm,
-)
from faker import Faker
-from generalresearch.incite.collections.thl_web import LedgerDFCollection
-from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.incite.schemas.mergers.pop_ledger import (
numerical_col_names,
)
-from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.models.thl.finance import (
BusinessBalances,
POPFinancial,
ProductBalances,
)
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.session import Session
-from generalresearch.models.thl.user import User
-from test_utils.incite.collections.conftest import ledger_collection
-from test_utils.incite.mergers.conftest import pop_ledger_merge
-from test_utils.managers.ledger.conftest import (
- session_with_tx_factory,
-)
+
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import LedgerDFCollection
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.thl.ledger import LedgerAccount
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
fake = Faker()
class TestProductBalanceInitialize:
-
def test_unknown_fields(self):
with pytest.raises(expected_exception=ValueError):
ProductBalances.model_validate(
@@ -210,6 +209,8 @@ class TestProductBalanceInitialize:
# Confirm the @property computed fields show up in openapi. I don't
# know how to do that yet... so this is check to confirm they're
# known computed fields for now
+
+ assert isinstance(instance, ProductBalances)
computed_fields = list(instance.model_computed_fields.keys())
assert "payout" in computed_fields
assert "adjustment" in computed_fields
@@ -244,7 +245,6 @@ class TestProductBalanceInitialize:
class TestBusinessBalanceInitialize:
-
def test_validate_product_ids(self):
instance1 = ProductBalances.model_validate(
{"bp_payment.CREDIT": 500, "bp_adjustment.DEBIT": 40}
@@ -653,37 +653,37 @@ class TestBusinessBalanceInitialize:
@pytest.mark.parametrize(
- argnames="offset, duration",
+ argnames="duration",
argvalues=list(
iter_product(
- ["12h", "2D"],
[timedelta(days=2), timedelta(days=5)],
)
),
)
class TestProductFinanceData:
-
def test_base(
self,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge,
+ client_no_amm,
+ duration: timedelta,
product: Product,
user_factory: Callable[..., User],
start: datetime,
- duration: timedelta,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
+ session_with_tx_factory: Callable[..., None],
):
# -- Build & Setup
- # assert ledger_collection.start is None
- # assert ledger_collection.offset is None
u: User = user_factory(product=product, created=ledger_collection.start)
+ assert u.product
for item in ledger_collection.items:
-
for _ in range(3):
rand_item_time = fake.date_time_between(
start_date=item.start,
end_date=item.finish,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
)
session_with_tx_factory(started=rand_item_time, user=u)
@@ -697,10 +697,9 @@ class TestProductFinanceData:
item_finishes = [i.finish for i in ledger_collection.items]
item_finishes.sort(reverse=True)
- last_item_finish = item_finishes[0]
# --
- account = thl_lm.get_account_or_create_bp_wallet(product=u.product)
+ account = thl_ledger_manager.get_account_or_create_bp_wallet(product=u.product)
ddf = pop_ledger_merge.ddf(
force_rr_latest=False,
@@ -732,17 +731,7 @@ class TestProductFinanceData:
assert len(res) == len({i.time for i in res})
-@pytest.mark.parametrize(
- argnames="offset, duration",
- argvalues=list(
- iter_product(
- ["12h", "2D"],
- [timedelta(days=2), timedelta(days=5)],
- )
- ),
-)
class TestPOPFinancialData:
-
def test_base(
self,
client_no_amm: DaskClient,
@@ -751,19 +740,16 @@ class TestPOPFinancialData:
user_factory: Callable[..., User],
product: Product,
start: datetime,
- duration: timedelta,
- create_main_accounts,
+ create_main_accounts: Callable[..., None],
session_with_tx_factory: Callable[..., Session],
- thl_lm: ThlLedgerManager,
- delete_df_collection,
- delete_ledger_db,
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
):
# -- Build & Setup
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- # assert ledger_collection.start is None
- # assert ledger_collection.offset is None
users = []
for _ in range(5):
@@ -773,7 +759,7 @@ class TestPOPFinancialData:
rand_item_time = fake.date_time_between(
start_date=item.start,
end_date=item.finish,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
)
session_with_tx_factory(started=rand_item_time, user=u)
@@ -792,8 +778,10 @@ class TestPOPFinancialData:
last_item_finish = item_finishes[0]
accounts = []
- for user in users:
- account = thl_lm.get_account_or_create_bp_wallet(product=u.product)
+ for _u in users:
+ account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=_u.product
+ )
accounts.append(account)
account_ids = [a.uuid for a in accounts]
@@ -809,6 +797,7 @@ class TestPOPFinancialData:
("time_idx", "<", last_item_finish),
],
)
+
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
df = df.groupby([pd.Grouper(key="time_idx", freq="D"), "account_id"]).sum()
@@ -821,25 +810,13 @@ class TestPOPFinancialData:
# This does not return the AccountID, it's the Product ID
assert i.product_id in [u.product_id for u in users]
- # 1 Product, multiple Users
+ # 1 product: Product, multiple Users
assert len(users) == len(accounts)
- # We group on days, and duration is a parameter to parametrize
- assert isinstance(duration, timedelta)
-
# -- Teardown
delete_df_collection(ledger_collection)
-@pytest.mark.parametrize(
- argnames="offset, duration",
- argvalues=list(
- iter_product(
- ["12h", "1D"],
- [timedelta(days=2), timedelta(days=3)],
- )
- ),
-)
class TestBusinessBalanceData:
def test_from_pandas(
self,
@@ -848,15 +825,14 @@ class TestBusinessBalanceData:
pop_ledger_merge: PopLedgerMerge,
user_factory: Callable[..., User],
product: Product,
- create_main_accounts,
- thl_lm: ThlLedgerManager,
- thl_web_rr,
- delete_df_collection,
- delete_ledger_db,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ product_manager: ProductManager,
+ delete_df_collection: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
session_with_tx_factory: Callable[..., Session],
- rm_ledger_collection,
+ rm_ledger_collection: Callable[..., None],
):
- from generalresearch.models.thl.ledger import LedgerAccount
delete_ledger_db()
create_main_accounts()
@@ -870,7 +846,7 @@ class TestBusinessBalanceData:
item_time = fake.date_time_between(
start_date=item.start,
end_date=item.finish,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
)
session_with_tx_factory(started=item_time, user=u)
item.initial_load(overwrite=True)
@@ -880,7 +856,9 @@ class TestBusinessBalanceData:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
# assert pop_ledger_merge.progress.has_archive.eq(True).all()
- account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=product)
+ account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=product
+ )
ddf = pop_ledger_merge.ddf(
force_rr_latest=False,
@@ -888,15 +866,18 @@ class TestBusinessBalanceData:
columns=numerical_col_names + ["account_id"],
filters=[("account_id", "in", [account.uuid])],
)
+ assert isinstance(ddf, dd.DataFrame)
ddf = ddf.groupby("account_id").sum()
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
assert isinstance(df, pd.DataFrame)
instance = BusinessBalances.from_pandas(
- input_data=df, accounts=[account], thl_pg_config=thl_web_rr
+ product_manager=product_manager,
+ input_data=df,
+ accounts=[account],
)
- balance: int = thl_lm.get_account_balance(account=account)
+ balance: int = thl_ledger_manager.get_account_balance(account=account)
assert instance.balance == balance
assert instance.net == balance
diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py
index 945ee7a..af8d2b9 100644
--- a/tests/models/thl/question/test_question_info.py
+++ b/tests/models/thl/question/test_question_info.py
@@ -1,145 +1,16 @@
+from __future__ import annotations
+
from generalresearch.models.thl.profiling.upk_property import (
- UpkProperty,
ProfilingInfo,
+ UpkProperty,
)
class TestQuestionInfo:
- def test_init(self):
+ def test_init(self, profiling_info_json: str):
- s = (
- '[{"property_label": "hispanic", "cardinality": "*", "prop_type": "i", "country_iso": "us", '
- '"property_id": "05170ae296ab49178a075cab2a2073a6", "item_id": "7911ec1468b146ee870951f8ae9cbac1", '
- '"item_label": "panamanian", "gold_standard": 1, "options": [{"id": "c358c11e72c74fa2880358f1d4be85ab", '
- '"label": "not_hispanic"}, {"id": "b1d6c475770849bc8e0200054975dc9c", "label": "yes_hispanic"}, '
- '{"id": "bd1eb44495d84b029e107c188003c2bd", "label": "other_hispanic"}, '
- '{"id": "f290ad5e75bf4f4ea94dc847f57c1bd3", "label": "mexican"}, '
- '{"id": "49f50f2801bd415ea353063bfc02d252", "label": "puerto_rican"}, '
- '{"id": "dcbe005e522f4b10928773926601f8bf", "label": "cuban"}, '
- '{"id": "467ef8ddb7ac4edb88ba9ef817cbb7e9", "label": "salvadoran"}, '
- '{"id": "3c98e7250707403cba2f4dc7b877c963", "label": "dominican"}, '
- '{"id": "981ee77f6d6742609825ef54fea824a8", "label": "guatemalan"}, '
- '{"id": "81c8057b809245a7ae1b8a867ea6c91e", "label": "colombian"}, '
- '{"id": "513656d5f9e249fa955c3b527d483b93", "label": "honduran"}, '
- '{"id": "afc8cddd0c7b4581bea24ccd64db3446", "label": "ecuadorian"}, '
- '{"id": "61f34b36e80747a89d85e1eb17536f84", "label": "argentinian"}, '
- '{"id": "5330cfa681d44aa8ade3a6d0ea198e44", "label": "peruvian"}, '
- '{"id": "e7bceaffd76e486596205d8545019448", "label": "nicaraguan"}, '
- '{"id": "b7bbb2ebf8424714962e6c4f43275985", "label": "spanish"}, '
- '{"id": "8bf539785e7a487892a2f97e52b1932d", "label": "venezuelan"}, '
- '{"id": "7911ec1468b146ee870951f8ae9cbac1", "label": "panamanian"}], "category": [{"id": '
- '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
- '"adwords_vertical_id": null}]}, {"property_label": "ethnic_group", "cardinality": "*", "prop_type": '
- '"i", "country_iso": "us", "property_id": "15070958225d4132b7f6674fcfc979f6", "item_id": '
- '"64b7114cf08143949e3bcc3d00a5d8a0", "item_label": "other_ethnicity", "gold_standard": 1, "options": [{'
- '"id": "a72e97f4055e4014a22bee4632cbf573", "label": "caucasians"}, '
- '{"id": "4760353bc0654e46a928ba697b102735", "label": "black_or_african_american"}, '
- '{"id": "20ff0a2969fa4656bbda5c3e0874e63b", "label": "asian"}, '
- '{"id": "107e0a79e6b94b74926c44e70faf3793", "label": "native_hawaiian_or_other_pacific_islander"}, '
- '{"id": "900fa12691d5458c8665bf468f1c98c1", "label": "native_americans"}, '
- '{"id": "64b7114cf08143949e3bcc3d00a5d8a0", "label": "other_ethnicity"}], "category": [{"id": '
- '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
- '"adwords_vertical_id": null}]}, {"property_label": "educational_attainment", "cardinality": "?", '
- '"prop_type": "i", "country_iso": "us", "property_id": "2637783d4b2b4075b93e2a156e16e1d8", "item_id": '
- '"934e7b81d6744a1baa31bbc51f0965d5", "item_label": "other_education", "gold_standard": 1, "options": [{'
- '"id": "df35ef9e474b4bf9af520aa86630202d", "label": "3rd_grade_completion"}, '
- '{"id": "83763370a1064bd5ba76d1b68c4b8a23", "label": "8th_grade_completion"}, '
- '{"id": "f0c25a0670c340bc9250099dcce50957", "label": "not_high_school_graduate"}, '
- '{"id": "02ff74c872bd458983a83847e1a9f8fd", "label": "high_school_completion"}, '
- '{"id": "ba8beb807d56441f8fea9b490ed7561c", "label": "vocational_program_completion"}, '
- '{"id": "65373a5f348a410c923e079ddbb58e9b", "label": "some_college_completion"}, '
- '{"id": "2d15d96df85d4cc7b6f58911fdc8d5e2", "label": "associate_academic_degree_completion"}, '
- '{"id": "497b1fedec464151b063cd5367643ffa", "label": "bachelors_degree_completion"}, '
- '{"id": "295133068ac84424ae75e973dc9f2a78", "label": "some_graduate_completion"}, '
- '{"id": "e64f874faeff4062a5aa72ac483b4b9f", "label": "masters_degree_completion"}, '
- '{"id": "cbaec19a636d476385fb8e7842b044f5", "label": "doctorate_degree_completion"}, '
- '{"id": "934e7b81d6744a1baa31bbc51f0965d5", "label": "other_education"}], "category": [{"id": '
- '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
- '"adwords_vertical_id": null}]}, {"property_label": "household_spoken_language", "cardinality": "*", '
- '"prop_type": "i", "country_iso": "us", "property_id": "5a844571073d482a96853a0594859a51", "item_id": '
- '"62b39c1de141422896ad4ab3c4318209", "item_label": "dut", "gold_standard": 1, "options": [{"id": '
- '"f65cd57b79d14f0f8460761ce41ec173", "label": "ara"}, {"id": "6d49de1f8f394216821310abd29392d9", '
- '"label": "zho"}, {"id": "be6dc23c2bf34c3f81e96ddace22800d", "label": "eng"}, '
- '{"id": "ddc81f28752d47a3b1c1f3b8b01a9b07", "label": "fre"}, {"id": "2dbb67b29bd34e0eb630b1b8385542ca", '
- '"label": "ger"}, {"id": "a747f96952fc4b9d97edeeee5120091b", "label": "hat"}, '
- '{"id": "7144b04a3219433baac86273677551fa", "label": "hin"}, {"id": "e07ff3e82c7149eaab7ea2b39ee6a6dc", '
- '"label": "ita"}, {"id": "b681eff81975432ebfb9f5cc22dedaa3", "label": "jpn"}, '
- '{"id": "5cb20440a8f64c9ca62fb49c1e80cdef", "label": "kor"}, {"id": "171c4b77d4204bc6ac0c2b81e38a10ff", '
- '"label": "pan"}, {"id": "8c3ec18e6b6c4a55a00dd6052e8e84fb", "label": "pol"}, '
- '{"id": "3ce074d81d384dd5b96f1fb48f87bf01", "label": "por"}, {"id": "6138dc951990458fa88a666f6ddd907b", '
- '"label": "rus"}, {"id": "e66e5ecc07df4ebaa546e0b436f034bd", "label": "spa"}, '
- '{"id": "5a981b3d2f0d402a96dd2d0392ec2fcb", "label": "tgl"}, {"id": "b446251bd211403487806c4d0a904981", '
- '"label": "vie"}, {"id": "92fb3ee337374e2db875fb23f52eed46", "label": "xxx"}, '
- '{"id": "8b1f590f12f24cc1924d7bdcbe82081e", "label": "ind"}, {"id": "bf3f4be556a34ff4b836420149fd2037", '
- '"label": "tur"}, {"id": "87ca815c43ba4e7f98cbca98821aa508", "label": "zul"}, '
- '{"id": "0adbf915a7a64d67a87bb3ce5d39ca54", "label": "may"}, {"id": "62b39c1de141422896ad4ab3c4318209", '
- '"label": "dut"}], "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", '
- '"path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": "gender", "cardinality": '
- '"?", "prop_type": "i", "country_iso": "us", "property_id": "73175402104741549f21de2071556cd7", '
- '"item_id": "093593e316344cd3a0ac73669fca8048", "item_label": "other_gender", "gold_standard": 1, '
- '"options": [{"id": "b9fc5ea07f3a4252a792fd4a49e7b52b", "label": "male"}, '
- '{"id": "9fdb8e5e18474a0b84a0262c21e17b56", "label": "female"}, '
- '{"id": "093593e316344cd3a0ac73669fca8048", "label": "other_gender"}], "category": [{"id": '
- '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
- '"adwords_vertical_id": null}]}, {"property_label": "age_in_years", "cardinality": "?", "prop_type": '
- '"n", "country_iso": "us", "property_id": "94f7379437874076b345d76642d4ce6d", "item_id": null, '
- '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", '
- '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": '
- '"children_age_gender", "cardinality": "*", "prop_type": "i", "country_iso": "us", "property_id": '
- '"e926142fcea94b9cbbe13dc7891e1e7f", "item_id": "b7b8074e95334b008e8958ccb0a204f1", "item_label": '
- '"female_18", "gold_standard": 1, "options": [{"id": "16a6448ec24c48d4993d78ebee33f9b4", '
- '"label": "male_under_1"}, {"id": "809c04cb2e3b4a3bbd8077ab62cdc220", "label": "female_under_1"}, '
- '{"id": "295e05bb6a0843bc998890b24c99841e", "label": "no_children"}, '
- '{"id": "142cb948d98c4ae8b0ef2ef10978e023", "label": "male_0"}, '
- '{"id": "5a5c1b0e9abc48a98b3bc5f817d6e9d0", "label": "male_1"}, '
- '{"id": "286b1a9afb884bdfb676dbb855479d1e", "label": "male_2"}, '
- '{"id": "942ca3cda699453093df8cbabb890607", "label": "male_3"}, '
- '{"id": "995818d432f643ec8dd17e0809b24b56", "label": "male_4"}, '
- '{"id": "f38f8b57f25f4cdea0f270297a1e7a5c", "label": "male_5"}, '
- '{"id": "975df709e6d140d1a470db35023c432d", "label": "male_6"}, '
- '{"id": "f60bd89bbe0f4e92b90bccbc500467c2", "label": "male_7"}, '
- '{"id": "6714ceb3ed5042c0b605f00b06814207", "label": "male_8"}, '
- '{"id": "c03c2f8271d443cf9df380e84b4dea4c", "label": "male_9"}, '
- '{"id": "11690ee0f5a54cb794f7ddd010d74fa2", "label": "male_10"}, '
- '{"id": "17bef9a9d14b4197b2c5609fa94b0642", "label": "male_11"}, '
- '{"id": "e79c8338fe28454f89ccc78daf6f409a", "label": "male_12"}, '
- '{"id": "3a4f87acb3fa41f4ae08dfe2858238c1", "label": "male_13"}, '
- '{"id": "36ffb79d8b7840a7a8cb8d63bbc8df59", "label": "male_14"}, '
- '{"id": "1401a508f9664347aee927f6ec5b0a40", "label": "male_15"}, '
- '{"id": "6e0943c5ec4a4f75869eb195e3eafa50", "label": "male_16"}, '
- '{"id": "47d4b27b7b5242758a9fff13d3d324cf", "label": "male_17"}, '
- '{"id": "9ce886459dd44c9395eb77e1386ab181", "label": "female_0"}, '
- '{"id": "6499ccbf990d4be5b686aec1c7353fd8", "label": "female_1"}, '
- '{"id": "d85ceaa39f6d492abfc8da49acfd14f2", "label": "female_2"}, '
- '{"id": "18edb45c138e451d8cb428aefbb80f9c", "label": "female_3"}, '
- '{"id": "bac6f006ed9f4ccf85f48e91e99fdfd1", "label": "female_4"}, '
- '{"id": "5a6a1a8ad00c4ce8be52dcb267b034ff", "label": "female_5"}, '
- '{"id": "6bff0acbf6364c94ad89507bcd5f4f45", "label": "female_6"}, '
- '{"id": "d0d56a0a6b6f4516a366a2ce139b4411", "label": "female_7"}, '
- '{"id": "bda6028468044b659843e2bef4db2175", "label": "female_8"}, '
- '{"id": "dbb6d50325464032b456357b1a6e5e9c", "label": "female_9"}, '
- '{"id": "b87a93d7dc1348edac5e771684d63fb8", "label": "female_10"}, '
- '{"id": "11449d0d98f14e27ba47de40b18921d7", "label": "female_11"}, '
- '{"id": "16156501e97b4263962cbbb743840292", "label": "female_12"}, '
- '{"id": "04ee971c89a345cc8141a45bce96050c", "label": "female_13"}, '
- '{"id": "e818d310bfbc4faba4355e5d2ed49d4f", "label": "female_14"}, '
- '{"id": "440d25e078924ba0973163153c417ed6", "label": "female_15"}, '
- '{"id": "78ff804cc9b441c5a524bd91e3d1f8bf", "label": "female_16"}, '
- '{"id": "4b04d804d7d84786b2b1c22e4ed440f5", "label": "female_17"}, '
- '{"id": "28bc848cd3ff44c3893c76bfc9bc0c4e", "label": "male_18"}, '
- '{"id": "b7b8074e95334b008e8958ccb0a204f1", "label": "female_18"}], "category": [{"id": '
- '"e18ba6e9d51e482cbb19acf2e6f505ce", "label": "Parenting", "path": "/People & Society/Family & '
- 'Relationships/Family/Parenting", "adwords_vertical_id": "58"}]}, {"property_label": "home_postal_code", '
- '"cardinality": "?", "prop_type": "x", "country_iso": "us", "property_id": '
- '"f3b32ebe78014fbeb1ed6ff77d6338bf", "item_id": null, "item_label": null, "gold_standard": 1, '
- '"category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", '
- '"adwords_vertical_id": null}]}, {"property_label": "household_income", "cardinality": "?", "prop_type": '
- '"n", "country_iso": "us", "property_id": "ff5b1d4501d5478f98de8c90ef996ac1", "item_id": null, '
- '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", '
- '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}]'
- )
- instance_list = ProfilingInfo.validate_json(s)
+ instance_list = ProfilingInfo.validate_json(profiling_info_json)
assert isinstance(instance_list, list)
for i in instance_list:
diff --git a/tests/models/thl/question/test_user_info.py b/tests/models/thl/question/test_user_info.py
index 0bbbc78..5410d35 100644
--- a/tests/models/thl/question/test_user_info.py
+++ b/tests/models/thl/question/test_user_info.py
@@ -1,32 +1,11 @@
+from __future__ import annotations
+
from generalresearch.models.thl.profiling.user_info import UserInfo
class TestUserInfo:
- def test_init(self):
+ def test_init(self, profiling_user_info_json: str):
- s = (
- '{"user_profile_knowledge": [], "marketplace_profile_knowledge": [{"source": "d", "question_id": '
- '"1", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "pr", '
- '"question_id": "3", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": '
- '"h", "question_id": "60", "answer": ["58"], "created": "2023-11-07T16:41:05.234096Z"}, '
- '{"source": "c", "question_id": "43", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
- '{"source": "s", "question_id": "211", "answer": ["111"], "created": '
- '"2023-11-07T16:41:05.234096Z"}, {"source": "s", "question_id": "1843", "answer": ["111"], '
- '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "h", "question_id": "13959", "answer": ['
- '"244155"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "33092", '
- '"answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender", '
- '"answer": ["10682"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "e", "question_id": '
- '"gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "f", '
- '"question_id": "gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": '
- '"i", "question_id": "gender", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
- '{"source": "c", "question_id": "137510", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, '
- '{"source": "m", "question_id": "gender", "answer": ["1"], "created": '
- '"2023-11-07T16:41:05.234096Z"}, {"source": "o", "question_id": "gender", "answer": ["male"], '
- '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender_plus", "answer": ['
- '"7657644"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "i", "question_id": '
- '"gender_plus", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", '
- '"question_id": "income_level", "answer": ["9071"], "created": "2023-11-07T16:41:05.234096Z"}]}'
- )
- instance = UserInfo.model_validate_json(s)
+ instance = UserInfo.model_validate_json(profiling_user_info_json)
assert isinstance(instance, UserInfo)
diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py
index 27091bb..5d8605a 100644
--- a/tests/models/thl/test_adjustments.py
+++ b/tests/models/thl/test_adjustments.py
@@ -1,36 +1,42 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.product import Product
from generalresearch.models.thl.session import (
- Session,
SessionAdjustedStatus,
Status,
StatusCode1,
- Wall,
WallAdjustedStatus,
+ Session,
+ Wall,
)
-from generalresearch.models.thl.user import User
-started1 = datetime(2023, 1, 1, tzinfo=timezone.utc)
-started2 = datetime(2023, 1, 1, 0, 10, 0, tzinfo=timezone.utc)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.thl.user import User
+
+started1 = datetime(2023, 1, 1, tzinfo=UTC)
+started2 = datetime(2023, 1, 1, 0, 10, 0, tzinfo=UTC)
finished1 = started1 + timedelta(minutes=10)
finished2 = started2 + timedelta(minutes=10)
-adj_ts = datetime(2023, 2, 2, tzinfo=timezone.utc)
-adj_ts2 = datetime(2023, 2, 3, tzinfo=timezone.utc)
-adj_ts3 = datetime(2023, 2, 4, tzinfo=timezone.utc)
+adj_ts = datetime(2023, 2, 2, tzinfo=UTC)
+adj_ts2 = datetime(2023, 2, 3, tzinfo=UTC)
+adj_ts3 = datetime(2023, 2, 4, tzinfo=UTC)
class TestProductAdjustments:
-
@pytest.mark.parametrize("payout", [".6", "1", "1.8", "2", "500.0000"])
def test_determine_bp_payment_no_rounding(
- self, product_factory: Callable[..., Product], payout
+ self, product_factory: Callable[..., Product], payout: str
):
p1 = product_factory(commission_pct=Decimal("0.05"))
res = p1.determine_bp_payment(thl_net=Decimal(payout))
@@ -39,7 +45,7 @@ class TestProductAdjustments:
@pytest.mark.parametrize("payout", [".01", ".05", ".5"])
def test_determine_bp_payment_rounding(
- self, product_factory: Callable[..., Product], payout
+ self, product_factory: Callable[..., Product], payout: str
):
p1 = product_factory(commission_pct=Decimal("0.05"))
res = p1.determine_bp_payment(thl_net=Decimal(payout))
@@ -48,7 +54,6 @@ class TestProductAdjustments:
class TestSessionAdjustments:
-
def test_status_complete(self, session_factory: Callable[..., Session], user: User):
# Completed Session with 2 wall events
s1 = session_factory(
@@ -60,7 +65,7 @@ class TestSessionAdjustments:
)
# Confirm only the last Wall Event is a complete
- assert not s1.wall_events[0].status == Status.COMPLETE
+ assert s1.wall_events[0].status != Status.COMPLETE
assert s1.wall_events[1].status == Status.COMPLETE
# Confirm the Session is marked as finished and the simple brokerage
@@ -71,9 +76,11 @@ class TestSessionAdjustments:
class TestAdjustments:
-
def test_finish_with_status(
- self, session_factory: Callable[..., Session], user: User, session_manager
+ self,
+ session_factory: Callable[..., Session],
+ user: User,
+ session_manager: SessionManager,
):
# Completed Session with 2 wall events
s1 = session_factory(
@@ -85,6 +92,7 @@ class TestAdjustments:
)
status, status_code_1 = s1.determine_session_status()
+ assert isinstance(user.product, Product)
payout = user.product.determine_bp_payment(Decimal(1))
session_manager.finish_with_status(
session=s1,
@@ -97,7 +105,10 @@ class TestAdjustments:
assert Decimal("0.95") == payout
def test_never_adjusted(
- self, session_factory: Callable[..., Session], user: User, session_manager
+ self,
+ session_factory: Callable[..., Session],
+ user: User,
+ session_manager: SessionManager,
):
s1 = session_factory(
user=user,
@@ -130,8 +141,8 @@ class TestAdjustments:
self,
session_factory: Callable[..., Session],
user: User,
- session_manager,
- wall_manager,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
):
# Completed Session with 2 wall events
s1 = session_factory(
@@ -174,13 +185,14 @@ class TestAdjustments:
# Because the Product doesn't have the Wallet mode enabled, the
# user_payout fields should always be None
+ assert isinstance(user.product, Product)
assert not user.product.user_wallet_config.enabled
assert s1.adjusted_user_payout is None
def test_adjustment_session_values(
self,
- wall_manager,
- session_manager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
session_factory: Callable[..., Session],
user: User,
):
@@ -218,13 +230,14 @@ class TestAdjustments:
# Because the Product doesn't have the Wallet mode enabled, the
# user_payout fields should always be None
+ assert isinstance(user.product, Product)
assert not user.product.user_wallet_config.enabled
assert s1.adjusted_user_payout is None
def test_double_adjustment_session_values(
self,
- wall_manager,
- session_manager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
session_factory: Callable[..., Session],
user: User,
):
@@ -276,8 +289,8 @@ class TestAdjustments:
def test_double_adjustment_sm_vs_db_values(
self,
- wall_manager,
- session_manager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
session_factory: Callable[..., Session],
user: User,
):
@@ -343,8 +356,8 @@ class TestAdjustments:
def test_double_adjustment_double_completes(
self,
- wall_manager,
- session_manager,
+ wall_manager: WallManager,
+ session_manager: SessionManager,
session_factory: Callable[..., Session],
user: User,
):
@@ -419,14 +432,14 @@ class TestAdjustments:
self,
session_factory: Callable[..., Session],
user: User,
- session_manager,
- wall_manager,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
utc_hour_ago: datetime,
):
s1 = session_factory(
user=user,
wall_count=1,
- wall_req_cpi=Decimal("1"),
+ wall_req_cpi=Decimal(1),
final_status=Status.COMPLETE,
started=utc_hour_ago,
)
@@ -435,6 +448,7 @@ class TestAdjustments:
assert status == Status.COMPLETE
thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete()))
+ assert isinstance(user.product, Product)
payout = user.product.determine_bp_payment(thl_net=thl_net)
session_manager.finish_with_status(
@@ -525,22 +539,20 @@ class TestAdjustments:
s1 = session_factory(
user=user,
wall_count=1,
- wall_req_cpi=Decimal("1"),
+ wall_req_cpi=Decimal(1),
final_status=Status.COMPLETE,
started=utc_hour_ago,
)
w1 = s1.wall_events[0]
status, status_code_1 = s1.determine_session_status()
- thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments()
+ _, _, bp_pay, user_pay = s1.determine_payments()
s1.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": utc_hour_ago + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
)
w1.update(
adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL,
@@ -562,6 +574,7 @@ class TestAdjustments:
new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts()
assert Status.COMPLETE == new_status
assert Decimal("0.95") == new_payout
+ assert isinstance(user.product, Product)
assert not user.product.user_wallet_config.enabled
# assert Decimal("0.48") == new_user_payout
assert new_user_payout is None
@@ -590,15 +603,14 @@ class TestAdjustments:
status, status_code_1 = s1.determine_session_status()
thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete()))
+ assert isinstance(user.product, Product)
payout = user.product.determine_bp_payment(thl_net=thl_net)
s1.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": utc_hour_ago + timedelta(minutes=25),
- "payout": payout,
- "user_payout": None,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=25),
+ payout=payout,
+ user_payout=None,
)
# Test. Adjust first fail to complete. Now we have 2 completes.
@@ -628,7 +640,10 @@ class TestAdjustments:
assert s1.adjusted_user_payout is None
def test_complete_to_fail_to_complete_adj1(
- self, user, session_factory, utc_hour_ago
+ self,
+ user: User,
+ session_factory: Callable[..., Session],
+ utc_hour_ago: datetime,
):
# Same as test_complete_to_fail_to_complete_adj but in opposite order
s1 = session_factory(
@@ -644,15 +659,14 @@ class TestAdjustments:
status, status_code_1 = s1.determine_session_status()
thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete()))
+ assert isinstance(user.product, Product)
payout = user.product.determine_bp_payment(thl_net)
s1.update(
- **{
- "status": status,
- "status_code_1": status_code_1,
- "finished": utc_hour_ago + timedelta(minutes=25),
- "payout": payout,
- "user_payout": None,
- }
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=25),
+ payout=payout,
+ user_payout=None,
)
# Test. Adjust complete to fail. Now we have 2 fails.
@@ -664,6 +678,7 @@ class TestAdjustments:
s1.adjust_status()
assert SessionAdjustedStatus.ADJUSTED_TO_FAIL == s1.adjusted_status
assert Decimal(0) == s1.adjusted_payout
+ assert isinstance(user.product, Product)
assert not user.product.user_wallet_config.enabled
# assert Decimal(0) == s.adjusted_user_payout
assert s1.adjusted_user_payout is None
@@ -708,6 +723,7 @@ class TestAdjustments:
s1.adjust_status()
assert SessionAdjustedStatus.ADJUSTED_TO_COMPLETE == s1.adjusted_status
assert Decimal("1.90") == s1.adjusted_payout
+ assert isinstance(user.product, Product)
assert not user.product.user_wallet_config.enabled
# assert Decimal("0.95") == s1.adjusted_user_payout
assert s1.adjusted_user_payout is None
diff --git a/tests/models/thl/test_bucket.py b/tests/models/thl/test_bucket.py
index 0aa5843..8d2f728 100644
--- a/tests/models/thl/test_bucket.py
+++ b/tests/models/thl/test_bucket.py
@@ -1,14 +1,17 @@
+from __future__ import annotations
+
from datetime import timedelta
from decimal import Decimal
import pytest
from pydantic import ValidationError
+from generalresearch.models.legacy.bucket import Bucket
+
class TestBucket:
def test_raises_payout(self):
- from generalresearch.models.legacy.bucket import Bucket
with pytest.raises(expected_exception=ValidationError) as e:
Bucket(user_payout_min=123)
@@ -27,7 +30,6 @@ class TestBucket:
assert "user_payout_min should be <= user_payout_max" in str(e.value)
def test_raises_loi(self):
- from generalresearch.models.legacy.bucket import Bucket
with pytest.raises(expected_exception=ValidationError) as e:
Bucket(loi_min=123)
@@ -63,7 +65,6 @@ class TestBucket:
assert "loi_q1 should be <= loi_q2" in str(e.value)
def test_parse_1(self):
- from generalresearch.models.legacy.bucket import Bucket
b1 = Bucket.parse_from_offerwall({"payout": {"min": 123}})
b_exp = Bucket(
@@ -180,7 +181,6 @@ class TestBucket:
assert b_exp == b4
def test_parse_3(self):
- from generalresearch.models.legacy.bucket import Bucket
b1 = Bucket.parse_from_offerwall({"payout": 123})
b_exp = Bucket(
diff --git a/tests/models/thl/test_buyer.py b/tests/models/thl/test_buyer.py
index eebb828..ef97166 100644
--- a/tests/models/thl/test_buyer.py
+++ b/tests/models/thl/test_buyer.py
@@ -1,4 +1,6 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.survey.buyer import BuyerCountryStat
diff --git a/tests/models/thl/test_contest/test_contest.py b/tests/models/thl/test_contest/test_contest.py
index 0fbd4cc..ed8477b 100644
--- a/tests/models/thl/test_contest/test_contest.py
+++ b/tests/models/thl/test_contest/test_contest.py
@@ -1,9 +1,13 @@
-from typing import Callable
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.models.thl.product import Product
-from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
class TestContest:
diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py
index 8b714ee..a639261 100644
--- a/tests/models/thl/test_contest/test_leaderboard_contest.py
+++ b/tests/models/thl/test_contest/test_leaderboard_contest.py
@@ -1,7 +1,11 @@
-from datetime import timezone
+from __future__ import annotations
+
+from datetime import UTC
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from redis import Redis
from generalresearch.currency import USDCent
from generalresearch.managers.leaderboard.manager import LeaderboardManager
@@ -17,16 +21,20 @@ from generalresearch.models.thl.contest.utils import (
distribute_leaderboard_prizes,
)
from generalresearch.models.thl.leaderboard import LeaderboardRow
-from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
from tests.models.thl.test_contest.test_contest import TestContest
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.user_manager.user_manager import UserManager
+ from generalresearch.models.thl.product import Product
+
class TestLeaderboardContest(TestContest):
@pytest.fixture
def leaderboard_contest(
- self, product: Product, thl_redis, user_manager
- ) -> "LeaderboardContest":
+ self, product: Product, thl_redis_client: Redis, user_manager: UserManager
+ ) -> LeaderboardContest:
board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count"
c = LeaderboardContest(
@@ -59,16 +67,22 @@ class TestLeaderboardContest(TestContest):
),
],
)
- c._redis_client = thl_redis
+ c._redis_client = thl_redis_client
c._user_manager = user_manager
return c
- def test_init(self, leaderboard_contest, thl_redis, user_1, user_2):
+ def test_init(
+ self,
+ leaderboard_contest: LeaderboardContest,
+ thl_redis_client: Redis,
+ user_1: User,
+ user_2: User,
+ ):
model = leaderboard_contest.leaderboard_model
assert leaderboard_contest.end_condition.ends_at is not None
lbm = LeaderboardManager(
- redis_client=thl_redis,
+ redis_client=thl_redis_client,
board_code=model.board_code,
country_iso=model.country_iso,
freq=model.freq,
@@ -83,15 +97,22 @@ class TestLeaderboardContest(TestContest):
lb = leaderboard_contest.get_leaderboard()
print(lb)
- def test_win(self, leaderboard_contest, thl_redis, user_1, user_2, user_3):
+ def test_win(
+ self,
+ leaderboard_contest: LeaderboardContest,
+ thl_redis_client: Redis,
+ user_1: User,
+ user_2: User,
+ user_3: User,
+ ):
model = leaderboard_contest.leaderboard_model
lbm = LeaderboardManager(
- redis_client=thl_redis,
+ redis_client=thl_redis_client,
board_code=model.board_code,
country_iso=model.country_iso,
freq=model.freq,
product_id=leaderboard_contest.product_id,
- within_time=model.period_start_local.astimezone(tz=timezone.utc),
+ within_time=model.period_start_local.astimezone(tz=UTC),
)
lbm.hit_complete_count(product_user_id=user_1.product_user_id)
@@ -102,10 +123,13 @@ class TestLeaderboardContest(TestContest):
lbm.hit_complete_count(product_user_id=user_3.product_user_id)
leaderboard_contest.end_contest()
+ assert isinstance(leaderboard_contest.all_winners, list)
assert len(leaderboard_contest.all_winners) == 3
# Prizes are $15, $10, $5. user 2 and 3 ties for 2nd place, so they split (10 + 5)
assert leaderboard_contest.all_winners[0].awarded_cash_amount == USDCent(15_00)
+
+ assert isinstance(leaderboard_contest.all_winners[0].user, User)
assert (
leaderboard_contest.all_winners[0].user.product_user_id
== user_1.product_user_id
diff --git a/tests/models/thl/test_contest/test_raffle_contest.py b/tests/models/thl/test_contest/test_raffle_contest.py
index d7920f0..e71851e 100644
--- a/tests/models/thl/test_contest/test_raffle_contest.py
+++ b/tests/models/thl/test_contest/test_raffle_contest.py
@@ -1,4 +1,8 @@
+from __future__ import annotations
+
from collections import Counter
+from datetime import datetime
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -18,9 +22,12 @@ from generalresearch.models.thl.contest.definitions import (
ContestType,
)
from generalresearch.models.thl.contest.raffle import RaffleContest
-from generalresearch.models.thl.product import Product
from tests.models.thl.test_contest.test_contest import TestContest
+if TYPE_CHECKING:
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.user import User
+
class TestRaffleContest(TestContest):
@@ -42,7 +49,9 @@ class TestRaffleContest(TestContest):
)
@pytest.fixture(scope="function")
- def ended_raffle_contest(self, raffle_contest, utc_now) -> RaffleContest:
+ def ended_raffle_contest(
+ self, raffle_contest: RaffleContest, utc_now: datetime
+ ) -> RaffleContest:
# Fake ending the contest
raffle_contest = raffle_contest.model_copy()
raffle_contest.update(
@@ -55,7 +64,7 @@ class TestRaffleContest(TestContest):
class TestRaffleContestUserView(TestRaffleContest):
- def test_user_view(self, raffle_contest, user):
+ def test_user_view(self, raffle_contest: RaffleContest, user: User):
from generalresearch.models.thl.contest.raffle import RaffleUserView
data = {
@@ -78,7 +87,7 @@ class TestRaffleContestUserView(TestRaffleContest):
assert res["current_win_probability"] == approx(0.0099, rel=0.001)
assert res["projected_win_probability"] == approx(0.0099, rel=0.001)
- def test_win_pct(self, raffle_contest, user):
+ def test_win_pct(self, raffle_contest: RaffleContest, user: User):
from generalresearch.models.thl.contest.raffle import RaffleUserView
data = {
@@ -124,7 +133,9 @@ class TestRaffleContestUserView(TestRaffleContest):
class TestRaffleContestWinners(TestRaffleContest):
- def test_winners_1_prize(self, ended_raffle_contest, user_1, user_2, user_3):
+ def test_winners_1_prize(
+ self, ended_raffle_contest, user_1: User, user_2: User, user_3: User
+ ):
ended_raffle_contest.entries = [
ContestEntry(
user=user_1,
@@ -160,7 +171,13 @@ class TestRaffleContestWinners(TestRaffleContest):
assert c[user_2.user_id] == approx(10000 * 2 / 6, rel=0.1)
assert c[user_3.user_id] == approx(10000 * 3 / 6, rel=0.1)
- def test_winners_2_prizes(self, ended_raffle_contest, user_1, user_2, user_3):
+ def test_winners_2_prizes(
+ self,
+ ended_raffle_contest: RaffleContest,
+ user_1: User,
+ user_2: User,
+ user_3: User,
+ ):
ended_raffle_contest.prizes.append(
ContestPrize(
name="iPod 64GB Black",
@@ -193,7 +210,9 @@ class TestRaffleContestWinners(TestRaffleContest):
# Same user
assert all(w.user.user_id == user_1.user_id for w in winners)
- def test_winners_2_prizes_1_entry(self, ended_raffle_contest, user_3):
+ def test_winners_2_prizes_1_entry(
+ self, ended_raffle_contest: RaffleContest, user_3: User
+ ):
ended_raffle_contest.prizes = [
ContestPrize(
name="iPod 64GB White",
@@ -218,7 +237,9 @@ class TestRaffleContestWinners(TestRaffleContest):
winners = ended_raffle_contest.select_winners()
assert len(winners) == 1
- def test_winners_2_prizes_1_entry_2_pennies(self, ended_raffle_contest, user_3):
+ def test_winners_2_prizes_1_entry_2_pennies(
+ self, ended_raffle_contest: RaffleContest, user_3: User
+ ):
ended_raffle_contest.prizes = [
ContestPrize(
name="iPod 64GB White",
@@ -243,7 +264,12 @@ class TestRaffleContestWinners(TestRaffleContest):
assert len(winners) == 2
def test_winners_3_prizes_3_entries(
- self, ended_raffle_contest, product, user_1, user_2, user_3
+ self,
+ ended_raffle_contest: RaffleContest,
+ product: Product,
+ user_1: User,
+ user_2: User,
+ user_3: User,
):
ended_raffle_contest.prizes = [
ContestPrize(
diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py
index 257de3c..7c48dbd 100644
--- a/tests/models/thl/test_ledger.py
+++ b/tests/models/thl/test_ledger.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from uuid import uuid4
import pytest
@@ -21,7 +23,7 @@ class TestLedgerTransaction:
assert [] == t.entries
assert {} == t.metadata
t = LedgerTransaction(
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
metadata={"a": "b", "user": "1234"},
ext_description="foo",
)
diff --git a/tests/models/thl/test_marketplace_condition.py b/tests/models/thl/test_marketplace_condition.py
index 8a4b25c..6936a7c 100644
--- a/tests/models/thl/test_marketplace_condition.py
+++ b/tests/models/thl/test_marketplace_condition.py
@@ -1,15 +1,18 @@
+from __future__ import annotations
+
import pytest
from pydantic import ValidationError
+from generalresearch.models.definitions import LogicalOperator
+from generalresearch.models.thl.survey.condition import (
+ ConditionValueType,
+ MarketplaceCondition,
+)
+
class TestMarketplaceCondition:
def test_list_or(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
user_qas = {"q1": {"a2"}}
c = MarketplaceCondition(
@@ -46,11 +49,6 @@ class TestMarketplaceCondition:
assert c.evaluate_criterion(user_qas) is None
def test_list_or_negate(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
user_qas = {"q1": {"a2"}}
c = MarketplaceCondition(
@@ -87,11 +85,6 @@ class TestMarketplaceCondition:
assert c.evaluate_criterion(user_qas) is None
def test_list_and(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
user_qas = {"q1": {"a1", "a2"}}
c = MarketplaceCondition(
@@ -137,7 +130,7 @@ class TestMarketplaceCondition:
assert c.evaluate_criterion(user_qas) is None
def test_list_and_negate(self):
- from generalresearch.models import LogicalOperator
+ from generalresearch.models.definitions import LogicalOperator
from generalresearch.models.thl.survey.condition import (
ConditionValueType,
MarketplaceCondition,
@@ -178,11 +171,6 @@ class TestMarketplaceCondition:
assert c.evaluate_criterion(user_qas) is None
def test_ranges(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
user_qas = {"q1": {"2", "50"}}
c = MarketplaceCondition(
@@ -245,12 +233,6 @@ class TestMarketplaceCondition:
)
def test_ranges_to_list(self):
- from generalresearch.models import LogicalOperator
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
-
user_qas = {"q1": {"2", "50"}}
MarketplaceCondition._CONVERT_LIST_TO_RANGE = ["q1"]
c = MarketplaceCondition(
@@ -265,7 +247,7 @@ class TestMarketplaceCondition:
assert ["1", "10", "11", "12", "2", "3", "4", "5"] == c.values
def test_ranges_infinity(self):
- from generalresearch.models import LogicalOperator
+ from generalresearch.models.definitions import LogicalOperator
from generalresearch.models.thl.survey.condition import (
ConditionValueType,
MarketplaceCondition,
@@ -309,10 +291,6 @@ class TestMarketplaceCondition:
assert not c.evaluate_criterion({"q1": {"50"}})
def test_answered(self):
- from generalresearch.models.thl.survey.condition import (
- ConditionValueType,
- MarketplaceCondition,
- )
user_qas = {"q1": {"a2"}}
c = MarketplaceCondition(
diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py
index 3a51328..cc00f33 100644
--- a/tests/models/thl/test_payout.py
+++ b/tests/models/thl/test_payout.py
@@ -1,10 +1,122 @@
+from __future__ import annotations
+
+from uuid import uuid4
+
+import pytest
+from pydantic import ValidationError
+
+from generalresearch.currency import USDCent
+from generalresearch.models.gr import Team
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+)
+from generalresearch.models.gr.definitions import BusinessType
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+ BusinessPayoutEvent,
+)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+
class TestBusinessPayoutEvent:
def test_validate(self):
- from generalresearch.models.gr.business import Business
- instance = Business.model_validate_json(
- json_data='{"id":123,"uuid":"947f6ba5250d442b9a66cde9ee33605a","name":"Example » Demo","kind":"c","tax_number":null,"contact":null,"addresses":[],"teams":[{"id":53,"uuid":"8e4197dcaefe4f1f831a02b212e6b44a","name":"Example » Demo","memberships":null,"gr_users":null,"businesses":null,"products":null}],"products":[{"id":"fc23e741b5004581b30e6478363525df","id_int":1234,"name":"Example","enabled":true,"payments_enabled":true,"created":"2025-04-14T13:25:37.279403Z","team_id":"9e4197dcaefe4f1f831a02b212e6b44a","business_id":"857f6ba6160d442b9a66cde9ee33605a","tags":[],"commission_pct":"0.050000","redirect_url":"https://pam-api-us.reppublika.com/v2/public/4970ef00-0ef7-11f0-9962-05cb6323c84c/grl/status","harmonizer_domain":"https://talk.generalresearch.com/","sources_config":{"user_defined":[{"name":"w","active":false,"banned_countries":[],"allow_mobile_ip":true,"supplier_id":null,"allow_pii_only_buyers":false,"allow_unhashed_buyers":false,"withhold_profiling":false,"pass_unconditional_eligible_unknowns":true,"address":null,"allow_vpn":null,"distribute_harmonizer_active":null}]},"session_config":{"max_session_len":600,"max_session_hard_retry":5,"min_payout":"0.14"},"payout_config":{"payout_format":null,"payout_transformation":null},"user_wallet_config":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null},"user_create_config":{"min_hourly_create_limit":0,"max_hourly_create_limit":null},"offerwall_config":{},"profiling_config":{"enabled":true,"grs_enabled":true,"n_questions":null,"max_questions":10,"avg_question_count":5.0,"task_injection_freq_mult":1.0,"non_us_mult":2.0,"hidden_questions_expiration_hours":168},"user_health_config":{"banned_countries":[],"allow_ban_iphist":true},"yield_man_config":{},"balance":null,"payouts_total_str":null,"payouts_total":null,"payouts":null,"user_wallet":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null}}],"bank_accounts":[],"balance":{"product_balances":[{"product_id":"fc14e741b5004581b30e6478363414df","last_event":null,"bp_payment_credit":780251,"adjustment_credit":4678,"adjustment_debit":26446,"supplier_credit":0,"supplier_debit":451513,"user_bonus_credit":0,"user_bonus_debit":0,"issued_payment":0,"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","recoup":0,"recoup_usd_str":"$0.00","adjustment_percent":0.027898714644390074}],"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"net_usd_str":"$7,584.83","payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"balance_usd_str":"$3,069.70","retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","adjustment_percent":0.027898714644390074,"recoup":0,"recoup_usd_str":"$0.00"},"payouts_total_str":"$4,515.13","payouts_total":451513,"payouts":[{"bp_payouts":[{"uuid":"40cf2c3c341e4f9d985be4bca43e6116","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-08-02T09:18:20.433329Z","amount":345735,"status":"COMPLETE","ext_ref_id":null,"payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":345735,"amount_usd_str":"$3,457.35"}],"amount":345735,"amount_usd_str":"$3,457.35","created":"2025-08-02T09:18:20.433329Z","line_items":1,"ext_ref_id":null},{"bp_payouts":[{"uuid":"63ce1787087248978919015c8fcd5ab9","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-06-10T22:16:18.765668Z","amount":105778,"status":"COMPLETE","ext_ref_id":"11175997868","payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":105778,"amount_usd_str":"$1,057.78"}],"amount":105778,"amount_usd_str":"$1,057.78","created":"2025-06-10T22:16:18.765668Z","line_items":1,"ext_ref_id":"11175997868"}]}'
+ # Doesn't validate anymore
+ # instance = Business.model_validate_json(
+ # json_data='{"id":123,"uuid":"947f6ba5250d442b9a66cde9ee33605a","name":"Example » Demo","kind":"c","tax_number":null,"contact":null,"addresses":[],"teams":[{"id":53,"uuid":"8e4197dcaefe4f1f831a02b212e6b44a","name":"Example » Demo","memberships":null,"gr_users":null,"businesses":null,"products":null}],"products":[{"id":"fc23e741b5004581b30e6478363525df","id_int":1234,"name":"Example","enabled":true,"payments_enabled":true,"created":"2025-04-14T13:25:37.279403Z","team_id":"9e4197dcaefe4f1f831a02b212e6b44a","business_id":"857f6ba6160d442b9a66cde9ee33605a","tags":[],"commission_pct":"0.050000","redirect_url":"https://pam-api-us.reppublika.com/v2/public/4970ef00-0ef7-11f0-9962-05cb6323c84c/grl/status","harmonizer_domain":"https://talk.generalresearch.com/","sources_config":{"user_defined":[{"name":"w","active":false,"banned_countries":[],"allow_mobile_ip":true,"supplier_id":null,"allow_pii_only_buyers":false,"allow_unhashed_buyers":false,"withhold_profiling":false,"pass_unconditional_eligible_unknowns":true,"address":null,"allow_vpn":null,"distribute_harmonizer_active":null}]},"session_config":{"max_session_len":600,"max_session_hard_retry":5,"min_payout":"0.14"},"payout_config":{"payout_format":null,"payout_transformation":null},"user_wallet_config":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null},"user_create_config":{"min_hourly_create_limit":0,"max_hourly_create_limit":null},"offerwall_config":{},"profiling_config":{"enabled":true,"grs_enabled":true,"n_questions":null,"max_questions":10,"avg_question_count":5.0,"task_injection_freq_mult":1.0,"non_us_mult":2.0,"hidden_questions_expiration_hours":168},"user_health_config":{"banned_countries":[],"allow_ban_iphist":true},"yield_man_config":{},"balance":null,"payouts_total_str":null,"payouts_total":null,"payouts":null,"user_wallet":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null}}],"bank_accounts":[],"balance":{"product_balances":[{"product_id":"fc14e741b5004581b30e6478363414df","last_event":null,"bp_payment_credit":780251,"adjustment_credit":4678,"adjustment_debit":26446,"supplier_credit":0,"supplier_debit":451513,"user_bonus_credit":0,"user_bonus_debit":0,"issued_payment":0,"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","recoup":0,"recoup_usd_str":"$0.00","adjustment_percent":0.027898714644390074}],"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"net_usd_str":"$7,584.83","payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"balance_usd_str":"$3,069.70","retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","adjustment_percent":0.027898714644390074,"recoup":0,"recoup_usd_str":"$0.00"},"payouts_total_str":"$4,515.13","payouts_total":451513,"payouts":[{"bp_payouts":[{"uuid":"40cf2c3c341e4f9d985be4bca43e6116","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-08-02T09:18:20.433329Z","amount":345735,"status":"COMPLETE","ext_ref_id":null,"payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":345735,"amount_usd_str":"$3,457.35"}],"amount":345735,"amount_usd_str":"$3,457.35","created":"2025-08-02T09:18:20.433329Z","line_items":1,"ext_ref_id":null},{"bp_payouts":[{"uuid":"63ce1787087248978919015c8fcd5ab9","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-06-10T22:16:18.765668Z","amount":105778,"status":"COMPLETE","ext_ref_id":"11175997868","payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":105778,"amount_usd_str":"$1,057.78"}],"amount":105778,"amount_usd_str":"$1,057.78","created":"2025-06-10T22:16:18.765668Z","line_items":1,"ext_ref_id":"11175997868"}]}'
+ # )
+ # assert isinstance(instance, Business)
+
+ # Make manually
+ b = Business(
+ id=123,
+ uuid=uuid4().hex,
+ name="Example",
+ addresses=[
+ BusinessAddress(
+ uuid=uuid4().hex,
+ city="xxx",
+ line_1="xxx",
+ state="fl",
+ business_id=123,
+ )
+ ],
+ kind=BusinessType.COMPANY,
+ teams=[Team(uuid=uuid4().hex, name="Example » Demo")],
+ products=[],
+ bank_accounts=[],
+ )
+ assert isinstance(b, Business)
+
+ ext_ref_id = uuid4().hex
+ bpe = BusinessPayoutEvent(
+ business_id=uuid4().hex,
+ amount=USDCent(100_00),
+ payout_type=PayoutType.ACH,
+ ext_ref_id=ext_ref_id,
)
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(47_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ ),
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(53_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ ),
+ ]
+
+ # Test validations (amount sum)
+ with pytest.raises(
+ ValidationError,
+ match="BusinessPayoutEvent.amount must equal the sum of bp_payouts amounts",
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(47_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ )
+ ]
+
+ with pytest.raises(
+ ValidationError,
+ match="All BrokerageProductPayoutEvent.ext_ref_id values must equal",
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(100_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id="a different value",
+ )
+ ]
- assert isinstance(instance, Business)
+ with pytest.raises(
+ ValidationError, match="All BrokerageProductPayoutEvent.payout_type values"
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.PAYPAL,
+ amount=USDCent(100_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ )
+ ]
diff --git a/tests/models/thl/test_payout_format.py b/tests/models/thl/test_payout_format.py
index 83fde25..fe7aea5 100644
--- a/tests/models/thl/test_payout_format.py
+++ b/tests/models/thl/test_payout_format.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import pytest
from pydantic import BaseModel
diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py
index 39469dc..97abf0c 100644
--- a/tests/models/thl/test_product.py
+++ b/tests/models/thl/test_product.py
@@ -2,9 +2,10 @@ from __future__ import annotations
import os
import shutil
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -12,18 +13,9 @@ from dask.distributed import Client as DaskClient
from pydantic import ValidationError
from generalresearch.currency import USDCent
-from generalresearch.incite.base import GRLDatasets
-from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.product import ProductManager
-from generalresearch.models import Source
-from generalresearch.models.gr.business import Business
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.finance import ProductBalances
from generalresearch.models.thl.product import (
- BrokerageProductPayoutEvent,
- BrokerageProductPayoutEventManager,
IntegrationMode,
PayoutConfig,
PayoutTransformation,
@@ -35,12 +27,27 @@ from generalresearch.models.thl.product import (
SupplyConfig,
SupplyPolicy,
)
-from generalresearch.models.thl.session import Session
-from generalresearch.models.thl.user import User
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.collections.thl_web import LedgerDFCollection
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import (
+ ThlLedgerManager,
+ )
+ from generalresearch.managers.thl.payout import PayoutEventManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.gr.business import Business
+ from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+ )
+ from generalresearch.models.thl.product import BrokerageProductPayoutEventManager
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
+ from generalresearch.redis_helper import RedisConfig
-class TestProduct:
+class TestProduct:
def test_init(self):
# By default, just a Pydantic instance doesn't have an id_int
instance = Product.model_validate(
@@ -56,17 +63,19 @@ class TestProduct:
# We're not excluding anything here, only in the "*Out" variants
assert "id_int" in res
- def test_init_db(self, product_manager: ProductManager):
+ def test_init_db(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
# By default, just a Pydantic instance doesn't have an id_int
- instance = product_manager.create_dummy()
+ instance = product_factory()
assert isinstance(instance.id_int, int)
+ assert isinstance(instance, Product)
res = instance.model_dump_json()
- assert isinstance(res, Product)
# we json skip & exclude
- res = instance.model_dump()
- assert isinstance(res, Product)
+ p = Product.model_validate_json(res)
+ assert isinstance(p, Product)
def test_redirect_url(self):
p = Product.model_validate(
@@ -140,12 +149,6 @@ class TestProduct:
redirect_url="https://www.google.com/hey",
)
- assert isinstance(p.payout_config.payout_transformation, PayoutTransformation)
- assert isinstance(
- p.payout_config.payout_transformation.kwargs,
- PayoutTransformationPercentArgs,
- )
-
p.payout_config.payout_transformation = PayoutTransformation.model_validate(
{
"f": "payout_transformation_percent",
@@ -156,6 +159,11 @@ class TestProduct:
assert (
"payout_transformation_percent" == p.payout_config.payout_transformation.f
)
+
+ assert isinstance(
+ p.payout_config.payout_transformation.kwargs,
+ PayoutTransformationPercentArgs,
+ )
assert 0.5 == p.payout_config.payout_transformation.kwargs.pct
assert (
Decimal("0.10") == p.payout_config.payout_transformation.kwargs.min_payout
@@ -287,10 +295,10 @@ class TestProduct:
p.profiling_config = ProfilingConfig(max_questions=1)
assert p.profiling_config.max_questions == 1
- def test_bp_account(self, product, thl_lm):
+ def test_bp_account(self, product: Product, thl_ledger_manager: ThlLedgerManager):
assert product.bp_account is None
- product.prefetch_bp_account(thl_lm=thl_lm)
+ product.prefetch_bp_account(thl_lm=thl_ledger_manager)
from generalresearch.models.thl.ledger import LedgerAccount
@@ -583,14 +591,13 @@ class TestGlobalProductConfigFor:
class TestProductFinancials:
-
@pytest.fixture
def start(self) -> datetime:
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -598,55 +605,75 @@ class TestProductFinancials:
def test_balance(
self,
- business: Business,
+ gr_business: Business,
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
mnt_filepath: GRLDatasets,
- bp_payout_factory: Callable[..., BrokerageProductPayoutEvent],
- thl_lm: ThlLedgerManager,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
start: datetime,
brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
session_with_tx_factory: Callable[..., Session],
- delete_ledger_db,
- create_main_accounts,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
client_no_amm: DaskClient,
- ledger_collection,
+ ledger_collection: LedgerDFCollection,
pop_ledger_merge: PopLedgerMerge,
- delete_df_collection,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.currency import USDCent
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
- bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_user_wallet(user=u1)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_user_wallet(user=u1)
- assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 0
+ assert (
+ len(
+ thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet.uuid
+ )
+ )
+ == 0
+ )
session_with_tx_factory(
user=u1,
wall_req_cpi=Decimal(".50"),
started=start + timedelta(days=1),
)
- assert thl_lm.get_account_balance(account=bp_wallet) == 48
- assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 1
+ assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 48
+ assert (
+ len(
+ thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet.uuid
+ )
+ )
+ == 1
+ )
session_with_tx_factory(
user=u1,
wall_req_cpi=Decimal("1.00"),
started=start + timedelta(days=2),
)
- assert thl_lm.get_account_balance(account=bp_wallet) == 143
- assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 2
+ assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 143
+ assert (
+ len(
+ thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet.uuid
+ )
+ )
+ == 2
+ )
with pytest.raises(expected_exception=AssertionError) as cm:
p1.prebuild_balance(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
)
@@ -656,7 +683,7 @@ class TestProductFinancials:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
p1.prebuild_balance(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
)
@@ -670,7 +697,7 @@ class TestProductFinancials:
assert p1.balance.available_balance == 108
p1.prebuild_payouts(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bp_pem=brokerage_product_payout_event_manager,
)
assert p1.payouts is not None
@@ -680,14 +707,21 @@ class TestProductFinancials:
# -- Now pay them out...
- bp_payout_factory(
+ from generalresearch.currency import USDCent
+
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(50),
created=start + timedelta(days=3),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 3
+ assert (
+ len(
+ thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet.uuid
+ )
+ )
+ == 3
+ )
# RM the entire directories
shutil.rmtree(ledger_collection.archive_path)
@@ -699,7 +733,7 @@ class TestProductFinancials:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
p1.prebuild_balance(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
)
@@ -713,24 +747,29 @@ class TestProductFinancials:
assert p1.balance.available_balance == 70
p1.prebuild_payouts(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bp_pem=brokerage_product_payout_event_manager,
)
assert p1.payouts is not None
assert len(p1.payouts) == 1
- assert p1.payouts_total == 50
+ assert p1.payouts_total == USDCent(50)
assert p1.payouts_total_str == "$0.50"
# -- Now pay ou another!.
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(5),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
- assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 4
+ assert (
+ len(
+ thl_ledger_manager.get_tx_filtered_by_account(
+ account_uuid=bp_wallet.uuid
+ )
+ )
+ == 4
+ )
# RM the entire directories
shutil.rmtree(ledger_collection.archive_path)
@@ -742,7 +781,7 @@ class TestProductFinancials:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
p1.prebuild_balance(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
)
@@ -756,7 +795,7 @@ class TestProductFinancials:
assert p1.balance.available_balance == 66
p1.prebuild_payouts(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bp_pem=brokerage_product_payout_event_manager,
)
assert p1.payouts is not None
@@ -766,14 +805,13 @@ class TestProductFinancials:
class TestProductBalance:
-
@pytest.fixture
def start(self) -> datetime:
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -783,18 +821,20 @@ class TestProductBalance:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
session_with_tx_factory: Callable[..., Session],
- pop_ledger_merge,
+ pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory,
- payout_event_manager,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ payout_event_manager: PayoutEventManager,
):
# Now let's load it up and actually test some things
delete_ledger_db()
@@ -813,21 +853,18 @@ class TestProductBalance:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
# 2. Payout and build Parquets 2nd time
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
created=start + timedelta(days=1, minutes=1),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
with pytest.raises(expected_exception=AssertionError) as cm:
product.prebuild_balance(
- thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm
+ thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm
)
assert "Sql and Parquet Balance inconsistent" in str(cm)
@@ -835,18 +872,20 @@ class TestProductBalance:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory,
- payout_event_manager,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ payout_event_manager: PayoutEventManager,
):
# This is very similar to the test_complete_payout_pq_inconsistent
# test, however this time we're only going to assign the payout
@@ -872,31 +911,29 @@ class TestProductBalance:
# 2. Payout and build Parquets 2nd time but this payout is "now"
# so it hasn't already been archived
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ created=datetime.now(tz=UTC),
)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
# We just want to call this to confirm it doesn't raise.
- product.prebuild_balance(thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm)
+ product.prebuild_balance(
+ thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm
+ )
class TestProductPOPFinancial:
-
@pytest.fixture
def start(self) -> datetime:
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -906,14 +943,14 @@ class TestProductPOPFinancial:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
):
@@ -942,7 +979,7 @@ class TestProductPOPFinancial:
# --- test ---
assert product.pop_financial is None
product.prebuild_pop_financial(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
@@ -962,14 +999,13 @@ class TestProductPOPFinancial:
class TestProductCache:
-
@pytest.fixture
def start(self) -> datetime:
- return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
+ return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC)
@pytest.fixture
def offset(self) -> str:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -978,17 +1014,17 @@ class TestProductCache:
def test_basic(
self,
product: Product,
- mnt_filepath,
- thl_lm,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ thl_redis_config: RedisConfig,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
):
@@ -1003,7 +1039,7 @@ class TestProductCache:
assert res is None
with pytest.raises(expected_exception=AssertionError):
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
@@ -1025,7 +1061,7 @@ class TestProductCache:
# Now try again with everything in place
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
@@ -1050,21 +1086,23 @@ class TestProductCache:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ thl_redis_config: RedisConfig,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory,
- payout_event_manager,
- adj_to_fail_with_tx_factory,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ payout_event_manager: PayoutEventManager,
+ adj_to_fail_with_tx_factory: Callable[..., None],
):
# Now let's load it up and actually test some things
delete_ledger_db()
@@ -1083,14 +1121,11 @@ class TestProductCache:
)
# 2. Payout
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
created=start + timedelta(days=1, minutes=1),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
# 3. Recon
@@ -1104,7 +1139,7 @@ class TestProductCache:
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
diff --git a/tests/models/thl/test_product_userwalletconfig.py b/tests/models/thl/test_product_userwalletconfig.py
index 614fc0a..9fb9c73 100644
--- a/tests/models/thl/test_product_userwalletconfig.py
+++ b/tests/models/thl/test_product_userwalletconfig.py
@@ -1,13 +1,15 @@
+from __future__ import annotations
+
from itertools import groupby
from random import shuffle as rshuffle
from generalresearch.models.thl.product import (
UserWalletConfig,
)
-from generalresearch.models.thl.wallet import PayoutType
+from generalresearch.models.thl.wallet.definitions import PayoutType
-def all_equal(iterable):
+def all_equal(iterable: list[str]) -> bool:
g = groupby(iterable)
return next(g, True) and not next(g, False)
@@ -41,13 +43,13 @@ class TestProductUserWalletConfig:
# in the same order because they're the same
assert isinstance(instance.model_dump_json(), str)
res = []
- for idx in range(100):
+ for _ in range(100):
res.append(instance.model_dump_json())
assert all_equal(res)
def test_model_dump_payout_types(self):
res = []
- for idx in range(100):
+ for _ in range(100):
# Generate a random order of PayoutTypes each time
payout_types = [e for e in PayoutType]
diff --git a/tests/models/thl/test_soft_pair.py b/tests/models/thl/test_soft_pair.py
index 588847e..34902e2 100644
--- a/tests/models/thl/test_soft_pair.py
+++ b/tests/models/thl/test_soft_pair.py
@@ -1,12 +1,14 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from generalresearch.models.definitions import Source
+from generalresearch.models.dynata.survey import (
+ ConditionValueType,
+ DynataCondition,
+)
from generalresearch.models.thl.soft_pair import SoftPairResult, SoftPairResultType
def test_model():
- from generalresearch.models.dynata.survey import (
- ConditionValueType,
- DynataCondition,
- )
c1 = DynataCondition(
question_id="1", value_type=ConditionValueType.LIST, values=["a", "b"]
diff --git a/tests/models/thl/test_upkquestion.py b/tests/models/thl/test_upkquestion.py
index d32875c..719fcff 100644
--- a/tests/models/thl/test_upkquestion.py
+++ b/tests/models/thl/test_upkquestion.py
@@ -1,13 +1,30 @@
+from __future__ import annotations
+
import pytest
from pydantic import ValidationError
+from generalresearch.models.morning.question import (
+ MorningQuestion,
+ MorningQuestionType,
+)
+from generalresearch.models.thl.profiling.upk_question import (
+ PatternValidation,
+ UPKImportance,
+ UpkQuestion,
+ UpkQuestionChoice,
+ UpkQuestionConfigurationMC,
+ UpkQuestionConfigurationTE,
+ UpkQuestionSelectorMC,
+ UpkQuestionSelectorTE,
+ UpkQuestionType,
+ UpkQuestionValidation,
+ order_exclusive_options,
+)
+
class TestUpkQuestion:
def test_importance(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UPKImportance,
- )
res = UPKImportance(task_score=1, task_count=None)
assert isinstance(res, UPKImportance)
@@ -20,9 +37,6 @@ class TestUpkQuestion:
assert "Input should be greater than or equal to 0" in str(e.value)
def test_pattern(self):
- from generalresearch.models.thl.profiling.upk_question import (
- PatternValidation,
- )
s = PatternValidation(message="hi", pattern="x")
with pytest.raises(ValidationError) as e:
@@ -30,13 +44,6 @@ class TestUpkQuestion:
assert "Instance is frozen" in str(e.value)
def test_mc(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- UpkQuestionChoice,
- UpkQuestionConfigurationMC,
- UpkQuestionSelectorMC,
- UpkQuestionType,
- )
q = UpkQuestion(
id="601377a0d4c74529afc6293a8e5c3b5e",
@@ -126,14 +133,6 @@ class TestUpkQuestion:
assert "Extra inputs are not permitted" in str(e.value)
def test_te(self):
- from generalresearch.models.thl.profiling.upk_question import (
- PatternValidation,
- UpkQuestion,
- UpkQuestionConfigurationTE,
- UpkQuestionSelectorTE,
- UpkQuestionType,
- UpkQuestionValidation,
- )
q = UpkQuestion(
id="601377a0d4c74529afc6293a8e5c3b5e",
@@ -152,9 +151,6 @@ class TestUpkQuestion:
assert q.choices is None
def test_deserialization(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
q = UpkQuestion.model_validate(
{
@@ -195,24 +191,18 @@ class TestUpkQuestion:
assert q == UpkQuestion.model_validate(q.model_dump(mode="json"))
def test_from_morning(self):
- from generalresearch.models.morning.question import (
- MorningQuestion,
- MorningQuestionType,
- )
q = MorningQuestion(
- **{
- "id": "gender",
- "country_iso": "us",
- "language_iso": "eng",
- "name": "Gender",
- "text": "What is your gender?",
- "type": "s",
- "options": [
- {"id": "1", "text": "yes", "order": 1},
- {"id": "2", "text": "no", "order": 2},
- ],
- }
+ id="gender",
+ country_iso="us",
+ language_iso="eng",
+ name="Gender",
+ text="What is your gender?",
+ type="s",
+ options=[
+ {"id": "1", "text": "yes", "order": 1},
+ {"id": "2", "text": "no", "order": 2},
+ ],
)
q.to_upk_question()
q = MorningQuestion(
@@ -226,13 +216,6 @@ class TestUpkQuestion:
q.to_upk_question()
def test_order(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- UpkQuestionChoice,
- UpkQuestionSelectorMC,
- UpkQuestionType,
- order_exclusive_options,
- )
q = UpkQuestion(
country_iso="us",
@@ -266,9 +249,6 @@ class TestUpkQuestion:
class TestUpkQuestionValidateAnswer:
def test_validate_answer_SA(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
question = UpkQuestion.model_validate(
{
@@ -304,9 +284,6 @@ class TestUpkQuestionValidateAnswer:
)
def test_validate_answer_MA(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
question = UpkQuestion.model_validate(
{
@@ -376,9 +353,6 @@ class TestUpkQuestionValidateAnswer:
)
def test_validate_answer_TE(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
question = UpkQuestion.model_validate(
{
diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py
index 943ae8e..68b413c 100644
--- a/tests/models/thl/test_user.py
+++ b/tests/models/thl/test_user.py
@@ -1,25 +1,35 @@
+from __future__ import annotations
+
import json
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta, timezone
from decimal import Decimal
from random import choice as rand_choice
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from pydantic import ValidationError
+from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.userhealth import AuditLogManager
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.userhealth import AuditLog
+
class TestUserUserID:
def test_valid(self):
- from generalresearch.models.thl.user import User
val = randint(1, 2**30)
user = User(user_id=val)
assert user.user_id == val
def test_type(self):
- from generalresearch.models.thl.user import User
# It will cast str to int
assert User(user_id="1").user_id == 1
@@ -44,7 +54,6 @@ class TestUserUserID:
assert "Input should be a valid integer," in str(cm.value)
def test_zero(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValidationError) as cm:
User(user_id=0)
@@ -52,7 +61,6 @@ class TestUserUserID:
assert "Input should be greater than 0" in str(cm.value)
def test_negative(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValidationError) as cm:
User(user_id=-1)
@@ -60,7 +68,6 @@ class TestUserUserID:
assert "Input should be greater than 0" in str(cm.value)
def test_too_big(self):
- from generalresearch.models.thl.user import User
val = 2**31
with pytest.raises(expected_exception=ValidationError) as cm:
@@ -69,7 +76,6 @@ class TestUserUserID:
assert "Input should be less than 2147483648" in str(cm.value)
def test_identifiable(self):
- from generalresearch.models.thl.user import User
val = randint(1, 2**30)
user = User(user_id=val)
@@ -80,7 +86,6 @@ class TestUserProductID:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
@@ -89,7 +94,6 @@ class TestUserProductID:
assert user.product_id == product_id
def test_type(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValueError) as cm:
User(user_id=self.user_id, product_id=0)
@@ -102,7 +106,6 @@ class TestUserProductID:
assert "Input should be a valid string" in str(cm.value)
def test_empty(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValueError) as cm:
User(user_id=self.user_id, product_id="")
@@ -110,7 +113,6 @@ class TestUserProductID:
assert "String should have at least 32 characters" in str(cm.value)
def test_invalid_len(self):
- from generalresearch.models.thl.user import User
# Valid uuid4s are 32 char long
product_id = uuid4().hex[:31]
@@ -133,7 +135,6 @@ class TestUserProductID:
assert "String should have at most 32 characters" in str(cm.value)
def test_invalid_uuid(self):
- from generalresearch.models.thl.user import User
# Modify the UUID to break it
product_id = uuid4().hex[:31] + "x"
@@ -144,7 +145,6 @@ class TestUserProductID:
assert "Invalid UUID" in str(cm.value)
def test_invalid_hex_form(self):
- from generalresearch.models.thl.user import User
# Sure not in hex form, but it'll get caught for being the
# wrong length before anything else
@@ -157,7 +157,6 @@ class TestUserProductID:
def test_identifiable(self):
"""Can't create a User with only a product_id because it also
needs to the product_user_id"""
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
with pytest.raises(expected_exception=ValueError) as cm:
@@ -172,10 +171,9 @@ class TestUserProductUserID:
def randomword(self, length: int = 50):
# Raw so nothing is escaped to add additional backslashes
_bpuid_allowed = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[]^_{|}~"
- return "".join(rand_choice(_bpuid_allowed) for i in range(length))
+ return "".join(rand_choice(_bpuid_allowed) for _ in range(length))
def test_valid(self):
- from generalresearch.models.thl.user import User
product_user_id = uuid4().hex[:12]
user = User(user_id=self.user_id, product_user_id=product_user_id)
@@ -184,7 +182,6 @@ class TestUserProductUserID:
assert user.product_user_id == product_user_id
def test_type(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValueError) as cm:
User(user_id=self.user_id, product_user_id=0)
@@ -197,12 +194,11 @@ class TestUserProductUserID:
assert "Input should be a valid string" in str(cm.value)
with pytest.raises(ValueError) as cm:
- User(user_id=self.user_id, product_user_id=Decimal("0"))
+ User(user_id=self.user_id, product_user_id=Decimal(0))
assert "1 validation error for User" in str(cm.value)
assert "Input should be a valid string" in str(cm.value)
def test_empty(self):
- from generalresearch.models.thl.user import User
with pytest.raises(expected_exception=ValueError) as cm:
User(user_id=self.user_id, product_user_id="")
@@ -210,7 +206,6 @@ class TestUserProductUserID:
assert "String should have at least 3 characters" in str(cm.value)
def test_invalid_len(self):
- from generalresearch.models.thl.user import User
product_user_id = self.randomword(251)
with pytest.raises(expected_exception=ValueError) as cm:
@@ -225,7 +220,6 @@ class TestUserProductUserID:
assert "String should have at least 3 characters" in str(cm.value)
def test_invalid_chars_space(self):
- from generalresearch.models.thl.user import User
product_user_id = f"{self.randomword(50)} {self.randomword(50)}"
with pytest.raises(expected_exception=ValueError) as cm:
@@ -234,9 +228,8 @@ class TestUserProductUserID:
assert "String cannot contain spaces" in str(cm.value)
def test_invalid_chars_slash(self):
- from generalresearch.models.thl.user import User
- product_user_id = f"{self.randomword(50)}\{self.randomword(50)}"
+ product_user_id = rf"{self.randomword(50)}\{self.randomword(50)}"
with pytest.raises(expected_exception=ValueError) as cm:
User(user_id=self.user_id, product_user_id=product_user_id)
assert "1 validation error for User" in str(cm.value)
@@ -253,7 +246,6 @@ class TestUserProductUserID:
I wanted a test that made sure the regex was hit. I do not know
how we want to provide with the level of specific String checks
we do in here for specific error messages."""
- from generalresearch.models.thl.user import User
product_user_id = f"{self.randomword(50)}`{self.randomword(50)}"
with pytest.raises(expected_exception=ValueError) as cm:
@@ -275,7 +267,6 @@ class TestUserProductUserID:
def test_identifiable(self):
"""Can't create a User with only a product_user_id because it also
needs to the product_id"""
- from generalresearch.models.thl.user import User
product_user_id = uuid4().hex
with pytest.raises(ValueError) as cm:
@@ -288,7 +279,6 @@ class TestUserUUID:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
uuid_pk = uuid4().hex
@@ -297,7 +287,6 @@ class TestUserUUID:
assert user.uuid == uuid_pk
def test_type(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, uuid=0)
@@ -310,12 +299,11 @@ class TestUserUUID:
assert "Input should be a valid string" in str(cm.value)
with pytest.raises(ValueError) as cm:
- User(user_id=self.user_id, uuid=Decimal("0"))
+ User(user_id=self.user_id, uuid=Decimal(0))
assert "1 validation error for User" in str(cm.value)
assert "Input should be a valid string" in str(cm.value)
def test_empty(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, uuid="")
@@ -323,7 +311,6 @@ class TestUserUUID:
assert "String should have at least 32 characters" in str(cm.value)
def test_invalid_len(self):
- from generalresearch.models.thl.user import User
# Valid uuid4s are 32 char long
uuid_pk = uuid4().hex[:31]
@@ -341,7 +328,6 @@ class TestUserUUID:
assert "String should have at most 32 characters" in str(cm.value)
def test_invalid_uuid(self):
- from generalresearch.models.thl.user import User
# Modify the UUID to break it
uuid_pk = uuid4().hex[:31] + "x"
@@ -352,7 +338,6 @@ class TestUserUUID:
assert "Invalid UUID" in str(cm.value)
def test_invalid_hex_form(self):
- from generalresearch.models.thl.user import User
# Sure not in hex form, but it'll get caught for being the
# wrong length before anything else
@@ -369,7 +354,6 @@ class TestUserUUID:
assert "Invalid UUID" in str(cm.value)
def test_identifiable(self):
- from generalresearch.models.thl.user import User
user_uuid = uuid4().hex
user = User(uuid=user_uuid)
@@ -380,33 +364,29 @@ class TestUserCreated:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
user = User(user_id=self.user_id)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
user.created = dt
assert user.created == dt
def test_tz_naive_throws_init(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
- User(user_id=self.user_id, created=datetime.now(tz=None))
+ User(user_id=self.user_id, created=datetime.now(tz=None)) # noqa
assert "1 validation error for User" in str(cm.value)
assert "Input should have timezone info" in str(cm.value)
def test_tz_naive_throws_setter(self):
- from generalresearch.models.thl.user import User
user = User(user_id=self.user_id)
with pytest.raises(ValueError) as cm:
- user.created = datetime.now(tz=None)
+ user.created = datetime.now(tz=None) # noqa
assert "1 validation error for User" in str(cm.value)
assert "Input should have timezone info" in str(cm.value)
def test_tz_utc(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
User(
@@ -417,20 +397,18 @@ class TestUserCreated:
assert "Timezone is not UTC" in str(cm.value)
def test_not_in_future(self):
- from generalresearch.models.thl.user import User
- the_future = datetime.now(tz=timezone.utc) + timedelta(minutes=1)
+ the_future = datetime.now(tz=UTC) + timedelta(minutes=1)
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, created=the_future)
assert "1 validation error for User" in str(cm.value)
assert "Input is in the future" in str(cm.value)
def test_after_anno_domini(self):
- from generalresearch.models.thl.user import User
- before_ad = datetime(
- year=2015, month=1, day=1, tzinfo=timezone.utc
- ) + timedelta(minutes=1)
+ before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta(
+ minutes=1
+ )
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, created=before_ad)
assert "1 validation error for User" in str(cm.value)
@@ -441,33 +419,29 @@ class TestUserLastSeen:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
user = User(user_id=self.user_id)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
user.last_seen = dt
assert user.last_seen == dt
def test_tz_naive_throws_init(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
- User(user_id=self.user_id, last_seen=datetime.now(tz=None))
+ User(user_id=self.user_id, last_seen=datetime.now(tz=None)) # noqa
assert "1 validation error for User" in str(cm.value)
assert "Input should have timezone info" in str(cm.value)
def test_tz_naive_throws_setter(self):
- from generalresearch.models.thl.user import User
user = User(user_id=self.user_id)
with pytest.raises(ValueError) as cm:
- user.last_seen = datetime.now(tz=None)
+ user.last_seen = datetime.now(tz=None) # noqa
assert "1 validation error for User" in str(cm.value)
assert "Input should have timezone info" in str(cm.value)
def test_tz_utc(self):
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
User(
@@ -478,20 +452,18 @@ class TestUserLastSeen:
assert "Timezone is not UTC" in str(cm.value)
def test_not_in_future(self):
- from generalresearch.models.thl.user import User
- the_future = datetime.now(tz=timezone.utc) + timedelta(minutes=1)
+ the_future = datetime.now(tz=UTC) + timedelta(minutes=1)
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, last_seen=the_future)
assert "1 validation error for User" in str(cm.value)
assert "Input is in the future" in str(cm.value)
def test_after_anno_domini(self):
- from generalresearch.models.thl.user import User
- before_ad = datetime(
- year=2015, month=1, day=1, tzinfo=timezone.utc
- ) + timedelta(minutes=1)
+ before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta(
+ minutes=1
+ )
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, last_seen=before_ad)
assert "1 validation error for User" in str(cm.value)
@@ -502,7 +474,6 @@ class TestUserBlocked:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
user = User(user_id=self.user_id, blocked=True)
assert user.blocked
@@ -510,7 +481,6 @@ class TestUserBlocked:
def test_str_casting(self):
"""We don't want any of these to work, and that's why
we set strict=True on the column"""
- from generalresearch.models.thl.user import User
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, blocked="true")
@@ -547,20 +517,18 @@ class TestUserTiming:
user_id = randint(1, 2**30)
def test_valid(self):
- from generalresearch.models.thl.user import User
- created = datetime.now(tz=timezone.utc) - timedelta(minutes=60)
- last_seen = datetime.now(tz=timezone.utc) - timedelta(minutes=59)
+ created = datetime.now(tz=UTC) - timedelta(minutes=60)
+ last_seen = datetime.now(tz=UTC) - timedelta(minutes=59)
user = User(user_id=self.user_id, created=created, last_seen=last_seen)
assert user.created == created
assert user.last_seen == last_seen
def test_created_first(self):
- from generalresearch.models.thl.user import User
- created = datetime.now(tz=timezone.utc) - timedelta(minutes=60)
- last_seen = datetime.now(tz=timezone.utc) - timedelta(minutes=59)
+ created = datetime.now(tz=UTC) - timedelta(minutes=60)
+ last_seen = datetime.now(tz=UTC) - timedelta(minutes=59)
with pytest.raises(ValueError) as cm:
User(user_id=self.user_id, created=last_seen, last_seen=created)
@@ -572,7 +540,6 @@ class TestUserModelVerification:
"""Tests that may be dependent on more than 1 attribute"""
def test_identifiable(self):
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
product_user_id = uuid4().hex
@@ -580,7 +547,6 @@ class TestUserModelVerification:
assert user.is_identifiable
def test_valid_helper(self):
- from generalresearch.models.thl.user import User
user_bool = User.is_valid_ubp(
product_id=uuid4().hex, product_user_id=uuid4().hex
@@ -594,7 +560,6 @@ class TestUserModelVerification:
class TestUserSerialization:
def test_basic_json(self):
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
product_user_id = uuid4().hex
@@ -602,7 +567,7 @@ class TestUserSerialization:
user = User(
product_id=product_id,
product_user_id=product_user_id,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
blocked=False,
)
@@ -615,7 +580,6 @@ class TestUserSerialization:
assert d.get("created").endswith("Z")
def test_basic_dict(self):
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
product_user_id = uuid4().hex
@@ -623,7 +587,7 @@ class TestUserSerialization:
user = User(
product_id=product_id,
product_user_id=product_user_id,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
blocked=False,
)
@@ -633,10 +597,11 @@ class TestUserSerialization:
assert not d.get("blocked")
assert d.get("product") is None
- assert d.get("created").tzinfo == timezone.utc
+ created = d.get("created")
+ assert isinstance(created, datetime)
+ assert created.tzinfo == UTC
def test_from_json(self):
- from generalresearch.models.thl.user import User
product_id = uuid4().hex
product_user_id = uuid4().hex
@@ -644,41 +609,51 @@ class TestUserSerialization:
user = User(
product_id=product_id,
product_user_id=product_user_id,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
blocked=False,
)
u = User.model_validate_json(user.to_json())
assert u.product_id == product_id
assert u.product is None
- assert u.created.tzinfo == timezone.utc
+ assert isinstance(u.created, datetime)
+ assert u.created.tzinfo == UTC
class TestUserMethods:
- def test_audit_log(self, user, audit_log_manager):
+ def test_audit_log(
+ self,
+ audit_log_factory: Callable[..., AuditLog],
+ user: User,
+ audit_log_manager: AuditLogManager,
+ ):
assert user.audit_log is None
user.prefetch_audit_log(audit_log_manager=audit_log_manager)
assert user.audit_log == []
- audit_log_manager.create_dummy(user_id=user.user_id)
+ audit_log_factory(user_id=user.user_id)
user.prefetch_audit_log(audit_log_manager=audit_log_manager)
assert len(user.audit_log) == 1
def test_transactions(
- self, user_factory, thl_lm, session_with_tx_factory, product_user_wallet_yes
+ self,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
+ session_with_tx_factory: Callable[..., None],
+ product_user_wallet_yes: Product,
):
u1 = user_factory(product=product_user_wallet_yes)
assert u1.transactions is None
- u1.prefetch_transactions(thl_lm=thl_lm)
+ u1.prefetch_transactions(thl_lm=thl_ledger_manager)
assert u1.transactions == []
session_with_tx_factory(user=u1)
- u1.prefetch_transactions(thl_lm=thl_lm)
+ u1.prefetch_transactions(thl_lm=thl_ledger_manager)
assert len(u1.transactions) == 1
@pytest.mark.skip(reason="TODO")
- def test_location_history(self, user):
+ def test_location_history(self, user: User):
assert user.location_history is None
diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py
index 596849c..b8a0be3 100644
--- a/tests/models/thl/test_user_iphistory.py
+++ b/tests/models/thl/test_user_iphistory.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from generalresearch.models.thl.user_iphistory import (
UserIPHistory,
@@ -8,7 +10,7 @@ from generalresearch.models.thl.user_iphistory import (
def test_collapse_ip_records():
# This does not exist in a db, so we do not need fixtures/ real user ids, whatever
- now = datetime.now(tz=timezone.utc) - timedelta(days=1)
+ now = datetime.now(tz=UTC) - timedelta(days=1)
# Gets stored most recent first. This is reversed, but the validator will order it
records = [
UserIPRecord(ip="1.2.3.5", created=now + timedelta(minutes=1)),
diff --git a/tests/models/thl/test_user_metadata.py b/tests/models/thl/test_user_metadata.py
index 3d851dc..7e84f3e 100644
--- a/tests/models/thl/test_user_metadata.py
+++ b/tests/models/thl/test_user_metadata.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import pytest
-from generalresearch.models import MAX_INT32
+from generalresearch.models.definitions import MAX_INT32
from generalresearch.models.thl.user_profile import UserMetadata
diff --git a/tests/models/thl/test_user_streak.py b/tests/models/thl/test_user_streak.py
index 0cacd3e..8300474 100644
--- a/tests/models/thl/test_user_streak.py
+++ b/tests/models/thl/test_user_streak.py
@@ -1,8 +1,8 @@
from datetime import datetime, timedelta
+from zoneinfo import ZoneInfo
import pytest
from pydantic import ValidationError
-from zoneinfo import ZoneInfo
from generalresearch.models.thl.user_streak import (
StreakFulfillment,
@@ -71,6 +71,7 @@ def test_user_streak_remaining():
)
print(f"{now.isoformat()=}, {end_of_today.isoformat()=}")
expected = (end_of_today - now).total_seconds()
+ assert isinstance(us.time_remaining_in_period, timedelta)
assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1)
@@ -92,5 +93,6 @@ def test_user_streak_remaining_month():
).replace(day=1)
print(f"{now.isoformat()=}, {end_of_month.isoformat()=}")
expected = (end_of_month - now).total_seconds()
+ assert isinstance(us.time_remaining_in_period, timedelta)
assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1)
print(us.time_remaining_in_period)
diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py
index 8398c81..61ca11d 100644
--- a/tests/models/thl/test_wall.py
+++ b/tests/models/thl/test_wall.py
@@ -1,11 +1,13 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from uuid import uuid4
import pytest
from pydantic import ValidationError
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
Status,
StatusCode1,
@@ -27,8 +29,8 @@ class TestWall:
ext_status_code_1="1.0",
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
- started=datetime(2023, 1, 1, 0, 0, 1, tzinfo=timezone.utc),
- finished=datetime(2023, 1, 1, 0, 10, 1, tzinfo=timezone.utc),
+ started=datetime(2023, 1, 1, 0, 0, 1, tzinfo=UTC),
+ finished=datetime(2023, 1, 1, 0, 10, 1, tzinfo=UTC),
)
s = w.to_json()
w2 = Wall.from_json(s)
@@ -45,8 +47,8 @@ class TestWall:
survey_id="yyy",
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
Wall(
user_id=1,
@@ -58,8 +60,8 @@ class TestWall:
status=Status.FAIL,
status_code_1=StatusCode1.MARKETPLACE_FAIL,
status_code_2=WallStatusCode2.COMPLETE_TOO_FAST,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
with pytest.raises(expected_exception=ValidationError) as e:
Wall(
@@ -71,8 +73,8 @@ class TestWall:
survey_id="yyy",
status=Status.FAIL,
status_code_1=StatusCode1.GRS_ABANDON,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
assert "If status is f, status_code_1 should be in" in str(e.value)
@@ -87,8 +89,8 @@ class TestWall:
status=Status.FAIL,
status_code_1=StatusCode1.GRS_ABANDON,
status_code_2=WallStatusCode2.COMPLETE_TOO_FAST,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
assert "If status is f, status_code_1 should be in" in str(e.value)
@@ -104,8 +106,8 @@ class TestWall:
status=Status.FAIL,
status_code_1=StatusCode1.MARKETPLACE_FAIL,
status_code_2=WallStatusCode2.COMPLETE_TOO_FAST,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
Wall(
user_id=1,
@@ -117,8 +119,8 @@ class TestWall:
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
status_code_2=None,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
Wall(
user_id=1,
@@ -130,8 +132,8 @@ class TestWall:
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
status_code_2=None,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
with pytest.raises(expected_exception=ValidationError) as e:
@@ -145,8 +147,8 @@ class TestWall:
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
status_code_2=WallStatusCode2.COMPLETE_TOO_FAST,
- started=datetime.now(timezone.utc),
- finished=datetime.now(timezone.utc) + timedelta(seconds=1),
+ started=datetime.now(UTC),
+ finished=datetime.now(UTC) + timedelta(seconds=1),
)
assert "If status_code_1 is 1, status_code_2 should be in" in str(e.value)
diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py
index 1208c56..40d3619 100644
--- a/tests/models/thl/test_wall_session.py
+++ b/tests/models/thl/test_wall_session.py
@@ -1,9 +1,11 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.session import Session, Wall
from generalresearch.models.thl.user import User
@@ -12,7 +14,7 @@ from generalresearch.models.thl.user import User
class TestWallSession:
def test_session_with_no_wall_events(self):
- started = datetime(2023, 1, 1, tzinfo=timezone.utc)
+ started = datetime(2023, 1, 1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
assert s.status is None
assert s.status_code_1 is None
@@ -24,7 +26,7 @@ class TestWallSession:
# assert s.status_code_1 == StatusCode1.SESSION_START_FAIL
def test_session_timeout_with_only_grs(self):
- started = datetime(2023, 1, 1, tzinfo=timezone.utc)
+ started = datetime(2023, 1, 1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
user_id=1,
@@ -53,7 +55,7 @@ class TestWallSession:
# assert s.status_code_1 == StatusCode1.GRS_FAIL
def test_session_with_only_grs_complete(self):
- started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc)
+ started = datetime(year=2023, month=1, day=1, tzinfo=UTC)
# A Session is started
s = Session(user=User(user_id=1), started=started)
@@ -98,7 +100,7 @@ class TestWallSession:
# assert s.status_code_1 is None
def test_session_with_only_non_grs_fail(self):
- started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc)
+ started = datetime(year=2023, month=1, day=1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
@@ -119,7 +121,7 @@ class TestWallSession:
assert s.payout is None
def test_session_with_only_non_grs_timeout(self):
- started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc)
+ started = datetime(year=2023, month=1, day=1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
@@ -139,7 +141,7 @@ class TestWallSession:
assert s.payout is None
def test_session_with_grs_and_external(self):
- started = datetime(year=2023, month=1, day=1, tzinfo=timezone.utc)
+ started = datetime(year=2023, month=1, day=1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
@@ -168,7 +170,7 @@ class TestWallSession:
s.append_wall_event(w)
w.finish(
status=Status.ABANDON,
- finished=datetime.now(tz=timezone.utc) + timedelta(minutes=10),
+ finished=datetime.now(tz=UTC) + timedelta(minutes=10),
status_code_1=StatusCode1.BUYER_ABANDON,
)
status, status_code_1 = s.determine_session_status()
@@ -206,7 +208,7 @@ class TestWallSession:
assert s.payout is None
def test_session_marketplace_fail(self):
- started = datetime(2023, 1, 1, tzinfo=timezone.utc)
+ started = datetime(2023, 1, 1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
@@ -229,7 +231,7 @@ class TestWallSession:
assert StatusCode1.SESSION_CONTINUE_QUALITY_FAIL == s.status_code_1
def test_session_unknown(self):
- started = datetime(2023, 1, 1, tzinfo=timezone.utc)
+ started = datetime(2023, 1, 1, tzinfo=UTC)
s = Session(user=User(user_id=1), started=started)
w = Wall(
diff --git a/tests/sql_helper.py b/tests/sql_helper.py
index c4cc2ca..8ab7bdb 100644
--- a/tests/sql_helper.py
+++ b/tests/sql_helper.py
@@ -19,7 +19,7 @@ class TestSqlHelper:
def test_scheme(self):
from generalresearch.sql_helper import SqlHelper
- dsn = MySQLDsn(f"mysql://root@localhost/test")
+ dsn = MySQLDsn("mysql://root@localhost/test")
instance = SqlHelper(dsn=dsn)
assert instance.is_mysql()
@@ -30,7 +30,7 @@ class TestSqlHelper:
# self.assertTrue(instance.is_postgresql())
with pytest.raises(ValidationError):
- SqlHelper(dsn=MariaDBDsn(f"maria://root@localhost/test"))
+ SqlHelper(dsn=MariaDBDsn("maria://root@localhost/test"))
def test_row_decode(self):
from generalresearch.sql_helper import decode_uuids
diff --git a/tests/test_postgres.py b/tests/test_postgres.py
index 3b3ddd0..de3f5d8 100644
--- a/tests/test_postgres.py
+++ b/tests/test_postgres.py
@@ -1,18 +1,21 @@
import socket
import subprocess
-from typing import Callable
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from pydantic import PostgresDsn
-from generalresearch.models.custom_types import InternalHostname, PostgresDict
from generalresearch.pg_helper import PostgresConfig
+if TYPE_CHECKING:
+ from generalresearch.models.custom_types import InternalHostname, PostgresDict
+
def is_port_open(host: InternalHostname, port: int = 5432, timeout: int = 3):
try:
with socket.create_connection((host, port), timeout=timeout):
return True
- except (socket.timeout, ConnectionRefusedError, OSError):
+ except (TimeoutError, ConnectionRefusedError, OSError):
return False
@@ -65,4 +68,32 @@ class TestPostgresDjangoCreation:
WHERE table_schema = 'public';
""")
assert len(res) == 1
- assert res[0]["count"] == 56
+ assert res[0]["count"] == 57
+
+ def test_django_tables_only_gr(self, gr_db: PostgresConfig):
+ """
+ IMPORTANT: This can't really run with only the GR tables,
+ that's because we have most of the database init fixtures
+ as session scoped; and we can't ensure that this will
+ run before any test that depends on the core thl
+ migrations
+ """
+
+ res = gr_db.execute_sql_query(query="""
+ SELECT COUNT(*)
+ FROM information_schema.tables
+ WHERE table_schema = 'public';
+ """)
+ assert len(res) == 1
+ assert res[0]["count"] == 65
+
+ def test_django_tables_with_gr(
+ self, thl_web_rw: PostgresConfig, gr_db: PostgresConfig
+ ):
+ res = thl_web_rw.execute_sql_query(query="""
+ SELECT COUNT(*)
+ FROM information_schema.tables
+ WHERE table_schema = 'public';
+ """)
+ assert len(res) == 1
+ assert res[0]["count"] == 65
diff --git a/tests/wall_status_codes/test_analyze.py b/tests/wall_status_codes/test_analyze.py
index fa53dbb..e36ca3d 100644
--- a/tests/wall_status_codes/test_analyze.py
+++ b/tests/wall_status_codes/test_analyze.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.wall_status_codes import innovate
diff --git a/tests/wxet/models/test_definitions.py b/tests/wxet/models/test_definitions.py
index 543b9f1..a3616dd 100644
--- a/tests/wxet/models/test_definitions.py
+++ b/tests/wxet/models/test_definitions.py
@@ -1,12 +1,18 @@
+from __future__ import annotations
+
import pytest
+from generalresearch.wxet.models.definitions import (
+ WXETStatus,
+ WXETStatusCode1,
+ WXETStatusCode2,
+ check_wxet_status_consistent,
+)
+
class TestWXETStatusCode1:
def test_is_pre_task_entry_fail_pre(self):
- from generalresearch.wxet.models.definitions import (
- WXETStatusCode1,
- )
assert WXETStatusCode1.UNKNOWN.is_pre_task_entry_fail
assert WXETStatusCode1.WXET_FAIL.is_pre_task_entry_fail
@@ -32,12 +38,6 @@ class TestCheckWXETStatusConsistent:
def test_completes(self):
- from generalresearch.wxet.models.definitions import (
- WXETStatus,
- WXETStatusCode1,
- check_wxet_status_consistent,
- )
-
with pytest.raises(AssertionError) as cm:
check_wxet_status_consistent(
status=WXETStatus.COMPLETE,
@@ -52,12 +52,6 @@ class TestCheckWXETStatusConsistent:
def test_abandon(self):
- from generalresearch.wxet.models.definitions import (
- WXETStatus,
- WXETStatusCode1,
- check_wxet_status_consistent,
- )
-
with pytest.raises(AssertionError) as cm:
check_wxet_status_consistent(
status=WXETStatus.ABANDON,
@@ -71,12 +65,6 @@ class TestCheckWXETStatusConsistent:
def test_fail(self):
- from generalresearch.wxet.models.definitions import (
- WXETStatus,
- WXETStatusCode1,
- check_wxet_status_consistent,
- )
-
for sc1 in [
WXETStatusCode1.COMPLETE,
WXETStatusCode1.WXET_ABANDON,
@@ -95,13 +83,6 @@ class TestCheckWXETStatusConsistent:
StatusCode1.WXET_FAIL
"""
- from generalresearch.wxet.models.definitions import (
- WXETStatus,
- WXETStatusCode1,
- WXETStatusCode2,
- check_wxet_status_consistent,
- )
-
for sc2 in WXETStatusCode2:
with pytest.raises(AssertionError) as cm:
check_wxet_status_consistent(
diff --git a/tests/wxet/models/test_finish_type.py b/tests/wxet/models/test_finish_type.py
index 7bdeea7..afa3c76 100644
--- a/tests/wxet/models/test_finish_type.py
+++ b/tests/wxet/models/test_finish_type.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import pytest
from generalresearch.wxet.models.definitions import WXETStatus, WXETStatusCode1