aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--generalresearch/cacheing.py2
-rw-r--r--generalresearch/config.py57
-rw-r--r--generalresearch/grliq/models/__init__.py5
-rw-r--r--generalresearch/grliq/models/custom_types.py2
-rw-r--r--generalresearch/grliq/models/decider.py7
-rw-r--r--generalresearch/grliq/models/forensic_data.py145
-rw-r--r--generalresearch/grliq/models/forensic_result.py43
-rw-r--r--generalresearch/grliq/models/forensic_summary.py29
-rw-r--r--generalresearch/grliq/models/useragents.py23
-rw-r--r--generalresearch/grliq/utils.py7
-rw-r--r--generalresearch/grpc.py17
-rw-r--r--generalresearch/healing_ppe.py7
-rw-r--r--generalresearch/incite/collections/__init__.py52
-rw-r--r--generalresearch/incite/mergers/__init__.py6
-rw-r--r--generalresearch/incite/mergers/foundations/__init__.py15
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_session.py22
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_task_adjust.py8
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_wall.py20
-rw-r--r--generalresearch/incite/mergers/foundations/user_id_product.py10
-rw-r--r--generalresearch/incite/mergers/nginx_core.py146
-rw-r--r--generalresearch/incite/mergers/nginx_fsb.py150
-rw-r--r--generalresearch/incite/mergers/nginx_grs.py141
-rw-r--r--generalresearch/incite/mergers/pop_ledger.py12
-rw-r--r--generalresearch/incite/mergers/ym_survey_wall.py10
-rw-r--r--generalresearch/incite/mergers/ym_wall_summary.py10
-rw-r--r--generalresearch/incite/schemas/mergers/nginx.py140
-rw-r--r--generalresearch/incite/schemas/mergers/ym_wall_summary.py5
-rw-r--r--generalresearch/managers/base.py18
-rw-r--r--generalresearch/managers/cint/profiling.py18
-rw-r--r--generalresearch/managers/cint/survey.py20
-rw-r--r--generalresearch/managers/criteria.py8
-rw-r--r--generalresearch/managers/dynata/profiling.py18
-rw-r--r--generalresearch/managers/dynata/survey.py18
-rw-r--r--generalresearch/managers/events.py43
-rw-r--r--generalresearch/managers/gr/authentication.py58
-rw-r--r--generalresearch/managers/gr/business.py162
-rw-r--r--generalresearch/managers/gr/team.py84
-rw-r--r--generalresearch/managers/innovate/profiling.py18
-rw-r--r--generalresearch/managers/innovate/survey.py22
-rw-r--r--generalresearch/managers/leaderboard/manager.py24
-rw-r--r--generalresearch/managers/lucid/profiling.py14
-rw-r--r--generalresearch/managers/marketplace/user_pid.py23
-rw-r--r--generalresearch/managers/morning/profiling.py20
-rw-r--r--generalresearch/managers/morning/survey.py20
-rw-r--r--generalresearch/managers/network/label.py48
-rw-r--r--generalresearch/managers/network/mtr.py16
-rw-r--r--generalresearch/managers/network/nmap.py18
-rw-r--r--generalresearch/managers/network/rdns.py4
-rw-r--r--generalresearch/managers/network/tool_run.py23
-rw-r--r--generalresearch/managers/pollfish/profiling.py18
-rw-r--r--generalresearch/managers/precision/profiling.py18
-rw-r--r--generalresearch/managers/precision/survey.py18
-rw-r--r--generalresearch/managers/prodege/profiling.py18
-rw-r--r--generalresearch/managers/prodege/survey.py16
-rw-r--r--generalresearch/managers/repdata/profiling.py18
-rw-r--r--generalresearch/managers/repdata/survey.py16
-rw-r--r--generalresearch/managers/sago/profiling.py18
-rw-r--r--generalresearch/managers/sago/survey.py22
-rw-r--r--generalresearch/managers/spectrum/profiling.py18
-rw-r--r--generalresearch/managers/spectrum/survey.py24
-rw-r--r--generalresearch/managers/survey.py5
-rw-r--r--generalresearch/managers/thl/buyer.py16
-rw-r--r--generalresearch/managers/thl/cashout_method.py45
-rw-r--r--generalresearch/managers/thl/category.py10
-rw-r--r--generalresearch/managers/thl/contest_manager.py82
-rw-r--r--generalresearch/managers/thl/ipinfo.py178
-rw-r--r--generalresearch/managers/thl/ledger_manager/conditions.py14
-rw-r--r--generalresearch/managers/thl/ledger_manager/ledger.py147
-rw-r--r--generalresearch/managers/thl/ledger_manager/thl_ledger.py91
-rw-r--r--generalresearch/managers/thl/maxmind/__init__.py162
-rw-r--r--generalresearch/managers/thl/maxmind/basic.py133
-rw-r--r--generalresearch/managers/thl/maxmind/insights.py50
-rw-r--r--generalresearch/managers/thl/payout.py177
-rw-r--r--generalresearch/managers/thl/product.py93
-rw-r--r--generalresearch/managers/thl/profiling/question.py26
-rw-r--r--generalresearch/managers/thl/profiling/schema.py7
-rw-r--r--generalresearch/managers/thl/profiling/uqa.py30
-rw-r--r--generalresearch/managers/thl/profiling/user_upk.py34
-rw-r--r--generalresearch/managers/thl/session.py155
-rw-r--r--generalresearch/managers/thl/survey.py129
-rw-r--r--generalresearch/managers/thl/survey_penalty.py11
-rw-r--r--generalresearch/managers/thl/tango_api.py4
-rw-r--r--generalresearch/managers/thl/task_adjustment.py17
-rw-r--r--generalresearch/managers/thl/user_compensate.py9
-rw-r--r--generalresearch/managers/thl/user_manager/__init__.py6
-rw-r--r--generalresearch/managers/thl/user_manager/mysql_user_manager.py30
-rw-r--r--generalresearch/managers/thl/user_manager/redis_user_manager.py20
-rw-r--r--generalresearch/managers/thl/user_manager/user_manager.py70
-rw-r--r--generalresearch/managers/thl/user_manager/user_metadata_manager.py38
-rw-r--r--generalresearch/managers/thl/user_streak.py15
-rw-r--r--generalresearch/managers/thl/userhealth.py114
-rw-r--r--generalresearch/managers/thl/wall.py99
-rw-r--r--generalresearch/models/admin/__init__.py5
-rw-r--r--generalresearch/models/admin/request.py12
-rw-r--r--generalresearch/models/cint/__init__.py1
-rw-r--r--generalresearch/models/cint/question.py26
-rw-r--r--generalresearch/models/cint/survey.py84
-rw-r--r--generalresearch/models/cint/task_collection.py10
-rw-r--r--generalresearch/models/dynata/question.py22
-rw-r--r--generalresearch/models/dynata/survey.py94
-rw-r--r--generalresearch/models/dynata/task_collection.py10
-rw-r--r--generalresearch/models/gr/authentication.py62
-rw-r--r--generalresearch/models/gr/business.py108
-rw-r--r--generalresearch/models/gr/team.py48
-rw-r--r--generalresearch/models/innovate/question.py16
-rw-r--r--generalresearch/models/innovate/survey.py85
-rw-r--r--generalresearch/models/innovate/task_collection.py12
-rw-r--r--generalresearch/pg_helper.py5
-rw-r--r--generalresearch/priority_thread_pool.py4
-rw-r--r--generalresearch/sql_helper.py28
-rw-r--r--tests/sql_helper.py2
111 files changed, 1845 insertions, 2798 deletions
diff --git a/generalresearch/cacheing.py b/generalresearch/cacheing.py
index de000cf..34df267 100644
--- a/generalresearch/cacheing.py
+++ b/generalresearch/cacheing.py
@@ -4,7 +4,7 @@ from generalresearch import retry
class RetryCache:
# Simple pylibmc.Client wrapper that implements a retry on each method
- def __init__(self, client, tries=4, delay=1, backoff=1.5):
+ def __init__(self, client, tries: int = 4, delay: int = 1, backoff: float = 1.5):
import pylibmc
self.client = client
diff --git a/generalresearch/config.py b/generalresearch/config.py
index 92eacc2..551f75a 100644
--- a/generalresearch/config.py
+++ b/generalresearch/config.py
@@ -1,7 +1,8 @@
+from __future__ import annotations
+
import os
from datetime import datetime, timezone
from pathlib import Path
-from typing import Optional
from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn
from pydantic_settings import BaseSettings
@@ -40,67 +41,67 @@ def is_debug() -> bool:
class GRLBaseSettings(BaseSettings):
debug: bool = Field(default=True)
- redis: Optional[RedisDsn] = Field(default=None)
+ redis: RedisDsn | None = Field(default=None)
redis_timeout: float = Field(default=0.10)
- thl_redis: Optional[RedisDsn] = Field(default=None)
+ thl_redis: RedisDsn | None = Field(default=None)
- dask: Optional[DaskDsn] = Field(default=None, description="")
+ dask: DaskDsn | None = Field(default=None, description="")
- sentry: Optional[SentryDsn] = Field(
+ sentry: SentryDsn | None = Field(
default=None, description="The sentry.io DSN for connecting to a project"
)
- thl_mkpl_rw_db: Optional[MariaDBDsn] = Field(default=None)
- thl_mkpl_rr_db: Optional[MariaDBDsn] = Field(default=None)
+ thl_mkpl_rw_db: MariaDBDsn | None = Field(default=None)
+ thl_mkpl_rr_db: MariaDBDsn | None = Field(default=None)
# Primary DB, SELECT permissions
- thl_web_ro_db: Optional[PostgresDsn] = Field(default=None)
+ thl_web_ro_db: PostgresDsn | None = Field(default=None)
# Primary DB, SELECT, INSERT, UPDATE permissions
- thl_web_rw_db: Optional[PostgresDsn] = Field(default=None)
+ thl_web_rw_db: PostgresDsn | None = Field(default=None)
# Primary DB, SELECT, INSERT, UPDATE, DELETE permissions
- thl_web_rwd_db: Optional[PostgresDsn] = Field(default=None)
+ thl_web_rwd_db: PostgresDsn | None = Field(default=None)
# Slave/secondary/read-replica SELECT permission only
- thl_web_rr_db: Optional[PostgresDsn] = Field(default=None)
+ thl_web_rr_db: PostgresDsn | None = Field(default=None)
tmp_dir: DirectoryPath = Field(default=Path("/tmp"))
- spectrum_rw_db: Optional[MariaDBDsn] = Field(default=None)
- spectrum_rr_db: Optional[MariaDBDsn] = Field(default=None)
+ spectrum_rw_db: MariaDBDsn | None = Field(default=None)
+ spectrum_rr_db: MariaDBDsn | None = Field(default=None)
- precision_rw_db: Optional[MariaDBDsn] = Field(default=None)
- precision_rr_db: Optional[MariaDBDsn] = Field(default=None)
+ precision_rw_db: MariaDBDsn | None = Field(default=None)
+ precision_rr_db: MariaDBDsn | None = Field(default=None)
# --- GR ----
- gr_db: Optional[PostgresDsn] = Field(default=None)
- gr_redis: Optional[RedisDsn] = Field(default=None)
+ gr_db: PostgresDsn | None = Field(default=None)
+ gr_redis: RedisDsn | None = Field(default=None)
# --- GRL IQ ---
- grliq_db: Optional[PostgresDsn] = Field(default=None)
- mnt_grliq_archive_dir: Optional[str] = Field(
+ grliq_db: PostgresDsn | None = Field(default=None)
+ mnt_grliq_archive_dir: str | None = Field(
default=None,
description="Where gr-api can pull GRL-IQ Forensic archive items like"
"the captured screenshots.",
)
- mnt_gr_api_dir: Optional[str] = Field(
+ mnt_gr_api_dir: str | None = Field(
default=None,
description="Where gr-api can pull parquet files from.",
)
# --- TangoCard Configuration ---
- tango_platform_name: Optional[str] = Field(default=None)
- tango_platform_key: Optional[str] = Field(default=None)
- tango_account_id: Optional[str] = Field(default=None)
- tango_customer_id: Optional[str] = Field(default=None)
+ tango_platform_name: str | None = Field(default=None)
+ tango_platform_key: str | None = Field(default=None)
+ tango_account_id: str | None = Field(default=None)
+ 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: Optional[str] = Field(default=None)
- amt_assignment_cashout_method_id: Optional[str] = Field(default=None)
+ amt_bonus_cashout_method_id: str | None = Field(default=None)
+ amt_assignment_cashout_method_id: str | None = Field(default=None)
# --- Maxmind Configuration ---
- maxmind_account_id: Optional[str] = Field(default=None)
- maxmind_license_key: Optional[str] = Field(default=None)
+ maxmind_account_id: str | None = Field(default=None)
+ maxmind_license_key: str | None = Field(default=None)
EXAMPLE_PRODUCT_ID = "1108d053e4fa47c5b0dbdcd03a7981e7"
diff --git a/generalresearch/grliq/models/__init__.py b/generalresearch/grliq/models/__init__.py
index 957e64c..998de2d 100644
--- a/generalresearch/grliq/models/__init__.py
+++ b/generalresearch/grliq/models/__init__.py
@@ -1,6 +1,7 @@
+from __future__ import annotations
+
import json
from enum import Enum
-from typing import List
class RiskWeighting(str, Enum):
@@ -63,4 +64,4 @@ AUDIO_CODEC_NAMES = [
]
font_str = '[".Aqua Kana",".Helvetica LT MM",".Times LT MM","18thCentury","8514oem","AR BERKLEY","AR JULIAN","AR PL UKai CN","AR PL UMing CN","AR PL UMing HK","AR PL UMing TW","AR PL UMing TW MBE","Aakar","Abadi MT Condensed Extra Bold","Abadi MT Condensed Light","Abyssinica SIL","AcmeFont","Adobe Arabic","Agency FB","Aharoni","Aharoni Bold","Al Bayan","Al Bayan Bold","Al Bayan Plain","Al Nile","Al Tarikh","Aldhabi","Alfredo","Algerian","Alien Encounters","Almonte Snow","American Typewriter","American Typewriter Bold","American Typewriter Condensed","American Typewriter Light","Amethyst","Andale Mono","Andale Mono Version","Andalus","Angsana New","AngsanaUPC","Ani","AnjaliOldLipi","Aparajita","Apple Braille","Apple Braille Outline 6 Dot","Apple Braille Outline 8 Dot","Apple Braille Pinpoint 6 Dot","Apple Braille Pinpoint 8 Dot","Apple Chancery","Apple Color Emoji","Apple LiGothic Medium","Apple LiSung Light","Apple SD Gothic Neo","Apple SD Gothic Neo Regular","Apple SD GothicNeo ExtraBold","Apple Symbols","AppleGothic","AppleGothic Regular","AppleMyungjo","AppleMyungjo Regular","AquaKana","Arabic Transparent","Arabic Typesetting","Arial","Arial Baltic","Arial Black","Arial Bold","Arial Bold Italic","Arial CE","Arial CYR","Arial Greek","Arial Hebrew","Arial Hebrew Bold","Arial Italic","Arial Narrow","Arial Narrow Bold","Arial Narrow Bold Italic","Arial Narrow Italic","Arial Rounded Bold","Arial Rounded MT Bold","Arial TUR","Arial Unicode MS","ArialHB","Arimo","Asimov","Autumn","Avenir","Avenir Black","Avenir Book","Avenir Next","Avenir Next Bold","Avenir Next Condensed","Avenir Next Condensed Bold","Avenir Next Demi Bold","Avenir Next Heavy","Avenir Next Regular","Avenir Roman","Ayuthaya","BN Jinx","BN Machine","BOUTON International Symbols","Baby Kruffy","Baghdad","Bahnschrift","Balthazar","Bangla MN","Bangla MN Bold","Bangla Sangam MN","Bangla Sangam MN Bold","Baskerville","Baskerville Bold","Baskerville Bold Italic","Baskerville Old Face","Baskerville SemiBold","Baskerville SemiBold Italic","Bastion","Batang","BatangChe","Bauhaus 93","Beirut","Bell MT","Bell MT Bold","Bell MT Italic","Bellerose","Berlin Sans FB","Berlin Sans FB Demi","Bernard MT Condensed","BiauKai","Big Caslon","Big Caslon Medium","Birch Std","Bitstream Charter","Bitstream Vera Sans","Blackadder ITC","Blackoak Std","Bobcat","Bodoni 72","Bodoni MT","Bodoni MT Black","Bodoni MT Poster Compressed","Bodoni Ornaments","BolsterBold","Book Antiqua","Book Antiqua Bold","Bookman Old Style","Bookman Old Style Bold","Bookshelf Symbol 7","Borealis","Bradley Hand","Bradley Hand ITC","Braggadocio","Brandish","Britannic Bold","Broadway","Browallia New","BrowalliaUPC","Brush Script","Brush Script MT","Brush Script MT Italic","Brush Script Std","Brussels","Calibri","Calibri Bold","Calibri Light","Californian FB","Calisto MT","Calisto MT Bold","Calligraphic","Calvin","Cambria","Cambria Bold","Cambria Math","Candara","Candara Bold","Candles","Carrois Gothic SC","Castellar","Centaur","Century","Century Gothic","Century Gothic Bold","Century Schoolbook","Century Schoolbook Bold","Century Schoolbook L","Chalkboard","Chalkboard Bold","Chalkboard SE","Chalkboard SE Bold","ChalkboardBold","Chalkduster","Chandas","Chaparral Pro","Chaparral Pro Light","Charlemagne Std","Charter","Chilanka","Chiller","Chinyen","Clarendon","Cochin","Cochin Bold","Colbert","Colonna MT","Comic Sans MS","Comic Sans MS Bold","Commons","Consolas","Consolas Bold","Constantia","Constantia Bold","Coolsville","Cooper Black","Cooper Std Black","Copperplate","Copperplate Bold","Copperplate Gothic Bold","Copperplate Light","Corbel","Corbel Bold","Cordia New","CordiaUPC","Corporate","Corsiva","Corsiva Hebrew","Corsiva Hebrew Bold","Courier","Courier 10 Pitch","Courier Bold","Courier New","Courier New Baltic","Courier New Bold","Courier New CE","Courier New Italic","Courier Oblique","Cracked Johnnie","Creepygirl","Curlz MT","Cursor","Cutive Mono","DFKai-SB","DIN Alternate","DIN Condensed","Damascus","Damascus Bold","Dancing Script","DaunPenh","David","Dayton","DecoType Naskh","Deja Vu","DejaVu LGC Sans","DejaVu Sans","DejaVu Sans Mono","DejaVu Serif","Deneane","Desdemona","Detente","Devanagari MT","Devanagari MT Bold","Devanagari Sangam MN","Didot","Didot Bold","Digifit","DilleniaUPC","Dingbats","Distant Galaxy","Diwan Kufi","Diwan Kufi Regular","Diwan Thuluth","Diwan Thuluth Regular","DokChampa","Dominican","Dotum","DotumChe","Droid Sans","Droid Sans Fallback","Droid Sans Mono","Dyuthi","Ebrima","Edwardian Script ITC","Elephant","Emmett","Engravers MT","Engravers MT Bold","Enliven","Eras Bold ITC","Estrangelo Edessa","Ethnocentric","EucrosiaUPC","Euphemia","Euphemia UCAS","Euphemia UCAS Bold","Eurostile","Eurostile Bold","Expressway Rg","FangSong","Farah","Farisi","Felix Titling","Fingerpop","Fixedsys","Flubber","Footlight MT Light","Forte","FrankRuehl","Frankfurter Venetian TT","Franklin Gothic Book","Franklin Gothic Book Italic","Franklin Gothic Medium","Franklin Gothic Medium Cond","Franklin Gothic Medium Italic","FreeMono","FreeSans","FreeSerif","FreesiaUPC","Freestyle Script","French Script MT","Futura","Futura Condensed ExtraBold","Futura Medium","GB18030 Bitmap","Gabriola","Gadugi","Garamond","Garamond Bold","Gargi","Garuda","Gautami","Gazzarelli","Geeza Pro","Geeza Pro Bold","Geneva","GenevaCY","Gentium","Gentium Basic","Gentium Book Basic","GentiumAlt","Georgia","Georgia Bold","Geotype TT","Giddyup Std","Gigi","Gill","Gill Sans","Gill Sans Bold","Gill Sans MT","Gill Sans MT Bold","Gill Sans MT Condensed","Gill Sans MT Ext Condensed Bold","Gill Sans MT Italic","Gill Sans Ultra Bold","Gill Sans Ultra Bold Condensed","Gisha","Glockenspiel","Gloucester MT Extra Condensed","Good Times","Goudy","Goudy Old Style","Goudy Old Style Bold","Goudy Stout","Greek Diner Inline TT","Gubbi","Gujarati MT","Gujarati MT Bold","Gujarati Sangam MN","Gujarati Sangam MN Bold","Gulim","GulimChe","GungSeo Regular","Gungseouche","Gungsuh","GungsuhChe","Gurmukhi","Gurmukhi MN","Gurmukhi MN Bold","Gurmukhi MT","Gurmukhi Sangam MN","Gurmukhi Sangam MN Bold","Haettenschweiler","Hand Me Down S (BRK)","Hansen","Harlow Solid Italic","Harrington","Harvest","HarvestItal","Haxton Logos TT","HeadLineA Regular","HeadlineA","Heavy Heap","Hei","Hei Regular","Heiti SC","Heiti SC Light","Heiti SC Medium","Heiti TC","Heiti TC Light","Heiti TC Medium","Helvetica","Helvetica Bold","Helvetica CY Bold","Helvetica CY Plain","Helvetica LT Std","Helvetica Light","Helvetica Neue","Helvetica Neue Bold","Helvetica Neue Medium","Helvetica Oblique","HelveticaCY","HelveticaNeueLT Com 107 XBlkCn","Herculanum","High Tower Text","Highboot","Hiragino Kaku Gothic Pro W3","Hiragino Kaku Gothic Pro W6","Hiragino Kaku Gothic ProN W3","Hiragino Kaku Gothic ProN W6","Hiragino Kaku Gothic Std W8","Hiragino Kaku Gothic StdN W8","Hiragino Maru Gothic Pro W4","Hiragino Maru Gothic ProN W4","Hiragino Mincho Pro W3","Hiragino Mincho Pro W6","Hiragino Mincho ProN W3","Hiragino Mincho ProN W6","Hiragino Sans GB W3","Hiragino Sans GB W6","Hiragino Sans W0","Hiragino Sans W1","Hiragino Sans W2","Hiragino Sans W3","Hiragino Sans W4","Hiragino Sans W5","Hiragino Sans W6","Hiragino Sans W7","Hiragino Sans W8","Hiragino Sans W9","Hobo Std","Hoefler Text","Hoefler Text Black","Hoefler Text Ornaments","Hollywood Hills","Hombre","Huxley Titling","ITC Stone Serif","ITF Devanagari","ITF Devanagari Marathi","ITF Devanagari Medium","Impact","Imprint MT Shadow","InaiMathi","Induction","Informal Roman","Ink Free","IrisUPC","Iskoola Pota","Italianate","Jamrul","JasmineUPC","Javanese Text","Jokerman","Juice ITC","KacstArt","KacstBook","KacstDecorative","KacstDigital","KacstFarsi","KacstLetter","KacstNaskh","KacstOffice","KacstOne","KacstPen","KacstPoster","KacstQurn","KacstScreen","KacstTitle","KacstTitleL","Kai","Kai Regular","KaiTi","Kailasa","Kailasa Regular","Kaiti SC","Kaiti SC Black","Kalapi","Kalimati","Kalinga","Kannada MN","Kannada MN Bold","Kannada Sangam MN","Kannada Sangam MN Bold","Kartika","Karumbi","Kedage","Kefa","Kefa Bold","Keraleeyam","Keyboard","Khmer MN","Khmer MN Bold","Khmer OS","Khmer OS System","Khmer Sangam MN","Khmer UI","Kinnari","Kino MT","KodchiangUPC","Kohinoor Bangla","Kohinoor Devanagari","Kohinoor Telugu","Kokila","Kokonor","Kokonor Regular","Kozuka Gothic Pr6N B","Kristen ITC","Krungthep","KufiStandardGK","KufiStandardGK Regular","Kunstler Script","Laksaman","Lao MN","Lao Sangam MN","Lao UI","LastResort","Latha","Leelawadee","Letter Gothic Std","LetterOMatic!","Levenim MT","LiHei Pro","LiSong Pro","Liberation Mono","Liberation Sans","Liberation Sans Narrow","Liberation Serif","Likhan","LilyUPC","Limousine","Lithos Pro Regular","LittleLordFontleroy","Lohit Assamese","Lohit Bengali","Lohit Devanagari","Lohit Gujarati","Lohit Gurmukhi","Lohit Hindi","Lohit Kannada","Lohit Malayalam","Lohit Odia","Lohit Punjabi","Lohit Tamil","Lohit Tamil Classical","Lohit Telugu","Loma","Lucida Blackletter","Lucida Bright","Lucida Bright Demibold","Lucida Bright Demibold Italic","Lucida Bright Italic","Lucida Calligraphy","Lucida Calligraphy Italic","Lucida Console","Lucida Fax","Lucida Fax Demibold","Lucida Fax Regular","Lucida Grande","Lucida Grande Bold","Lucida Handwriting","Lucida Handwriting Italic","Lucida Sans","Lucida Sans Demibold Italic","Lucida Sans Typewriter","Lucida Sans Typewriter Bold","Lucida Sans Unicode","Luminari","Luxi Mono","MS Gothic","MS Mincho","MS Outlook","MS PGothic","MS PMincho","MS Reference Sans Serif","MS Reference Specialty","MS Sans Serif","MS Serif","MS UI Gothic","MT Extra","MV Boli","Mael","Magneto","Maiandra GD","Malayalam MN","Malayalam MN Bold","Malayalam Sangam MN","Malayalam Sangam MN Bold","Malgun Gothic","Mallige","Mangal","Manorly","Marion","Marion Bold","Marker Felt","Marker Felt Thin","Marlett","Martina","Matura MT Script Capitals","Meera","Meiryo","Meiryo Bold","Meiryo UI","MelodBold","Menlo","Menlo Bold","Mesquite Std","Microsoft","Microsoft Himalaya","Microsoft JhengHei","Microsoft JhengHei UI","Microsoft New Tai Lue","Microsoft PhagsPa","Microsoft Sans Serif","Microsoft Tai Le","Microsoft Tai Le Bold","Microsoft Uighur","Microsoft YaHei","Microsoft YaHei UI","Microsoft Yi Baiti","Minerva","MingLiU","MingLiU-ExtB","MingLiU_HKSCS","Minion Pro","Miriam","Mishafi","Mishafi Gold","Mistral","Modern","Modern No. 20","Monaco","Mongolian Baiti","Monospace","Monotype Corsiva","Monotype Sorts","MoolBoran","Moonbeam","MotoyaLMaru","Mshtakan","Mshtakan Bold","Mukti Narrow","Muna","Myanmar MN","Myanmar MN Bold","Myanmar Sangam MN","Myanmar Text","Mycalc","Myriad Arabic","Myriad Hebrew","Myriad Pro","NISC18030","NSimSun","Nadeem","Nadeem Regular","Nakula","Nanum Barun Gothic","Nanum Gothic","Nanum Myeongjo","NanumBarunGothic","NanumGothic","NanumGothic Bold","NanumGothicCoding","NanumMyeongjo","NanumMyeongjo Bold","Narkisim","Nasalization","Navilu","Neon Lights","New Peninim MT","New Peninim MT Bold","News Gothic MT","News Gothic MT Bold","Niagara Engraved","Niagara Solid","Nimbus Mono L","Nimbus Roman No9 L","Nimbus Sans L","Nimbus Sans L Condensed","Nina","Nirmala UI","Nirmala.ttf","Norasi","Noteworthy","Noteworthy Bold","Noto Color Emoji","Noto Emoji","Noto Mono","Noto Naskh Arabic","Noto Nastaliq Urdu","Noto Sans","Noto Sans Armenian","Noto Sans Bengali","Noto Sans CJK","Noto Sans Canadian Aboriginal","Noto Sans Cherokee","Noto Sans Devanagari","Noto Sans Ethiopic","Noto Sans Georgian","Noto Sans Gujarati","Noto Sans Gurmukhi","Noto Sans Hebrew","Noto Sans JP","Noto Sans KR","Noto Sans Kannada","Noto Sans Khmer","Noto Sans Lao","Noto Sans Malayalam","Noto Sans Myanmar","Noto Sans Oriya","Noto Sans SC","Noto Sans Sinhala","Noto Sans Symbols","Noto Sans TC","Noto Sans Tamil","Noto Sans Telugu","Noto Sans Thai","Noto Sans Yi","Noto Serif","Notram","November","Nueva Std","Nueva Std Cond","Nyala","OCR A Extended","OCR A Std","Old English Text MT","OldeEnglish","Onyx","OpenSymbol","OpineHeavy","Optima","Optima Bold","Optima Regular","Orator Std","Oriya MN","Oriya MN Bold","Oriya Sangam MN","Oriya Sangam MN Bold","Osaka","Osaka-Mono","OsakaMono","PCMyungjo Regular","PCmyoungjo","PMingLiU","PMingLiU-ExtB","PR Celtic Narrow","PT Mono","PT Sans","PT Sans Bold","PT Sans Caption Bold","PT Sans Narrow Bold","PT Serif","Padauk","Padauk Book","Padmaa","Pagul","Palace Script MT","Palatino","Palatino Bold","Palatino Linotype","Palatino Linotype Bold","Papyrus","Papyrus Condensed","Parchment","Parry Hotter","PenultimateLight","Perpetua","Perpetua Bold","Perpetua Titling MT","Perpetua Titling MT Bold","Phetsarath OT","Phosphate","Phosphate Inline","Phosphate Solid","PhrasticMedium","PilGi Regular","Pilgiche","PingFang HK","PingFang SC","PingFang TC","Pirate","Plantagenet Cherokee","Playbill","Poor Richard","Poplar Std","Pothana2000","Prestige Elite Std","Pristina","Purisa","QuiverItal","Raanana","Raanana Bold","Raavi","Rachana","Rage Italic","RaghuMalayalam","Ravie","Rekha","Roboto","Rockwell","Rockwell Bold","Rockwell Condensed","Rockwell Extra Bold","Rockwell Italic","Rod","Roland","Rondalo","Rosewood Std Regular","RowdyHeavy","Russel Write TT","SF Movie Poster","STFangsong","STHeiti","STIXGeneral","STIXGeneral-Bold","STIXGeneral-Regular","STIXIntegralsD","STIXIntegralsD-Bold","STIXIntegralsSm","STIXIntegralsSm-Bold","STIXIntegralsUp","STIXIntegralsUp-Bold","STIXIntegralsUp-Regular","STIXIntegralsUpD","STIXIntegralsUpD-Bold","STIXIntegralsUpD-Regular","STIXIntegralsUpSm","STIXIntegralsUpSm-Bold","STIXNonUnicode","STIXNonUnicode-Bold","STIXSizeFiveSym","STIXSizeFiveSym-Regular","STIXSizeFourSym","STIXSizeFourSym-Bold","STIXSizeOneSym","STIXSizeOneSym-Bold","STIXSizeThreeSym","STIXSizeThreeSym-Bold","STIXSizeTwoSym","STIXSizeTwoSym-Bold","STIXVariants","STIXVariants-Bold","STKaiti","STSong","STXihei","SWGamekeys MT","Saab","Sahadeva","Sakkal Majalla","Salina","Samanata","Samyak Devanagari","Samyak Gujarati","Samyak Malayalam","Samyak Tamil","Sana","Sana Regular","Sans","Sarai","Sathu","Savoye LET Plain:1.0","Sawasdee","Script","Script MT Bold","Segoe MDL2 Assets","Segoe Print","Segoe Pseudo","Segoe Script","Segoe UI","Segoe UI Emoji","Segoe UI Historic","Segoe UI Semilight","Segoe UI Symbol","Serif","Shonar Bangla","Showcard Gothic","Shree Devanagari 714","Shruti","SignPainter-HouseScript","Silom","SimHei","SimSun","SimSun-ExtB","Simplified Arabic","Simplified Arabic Fixed","Sinhala MN","Sinhala MN Bold","Sinhala Sangam MN","Sinhala Sangam MN Bold","Sitka","Skia","Skia Regular","Skinny","Small Fonts","Snap ITC","Snell Roundhand","Snowdrift","Songti SC","Songti SC Black","Songti TC","Source Code Pro","Splash","Standard Symbols L","Stencil","Stencil Std","Stephen","Sukhumvit Set","Suruma","Sylfaen","Symbol","Symbole","System","System Font","TAMu_Kadambri","TAMu_Kalyani","TAMu_Maduram","TSCu_Comic","TSCu_Paranar","TSCu_Times","Tahoma","Tahoma Negreta","TakaoExGothic","TakaoExMincho","TakaoGothic","TakaoMincho","TakaoPGothic","TakaoPMincho","Tamil MN","Tamil MN Bold","Tamil Sangam MN","Tamil Sangam MN Bold","Tarzan","Tekton Pro","Tekton Pro Cond","Tekton Pro Ext","Telugu MN","Telugu MN Bold","Telugu Sangam MN","Telugu Sangam MN Bold","Tempus Sans ITC","Terminal","Terminator Two","Thonburi","Thonburi Bold","Tibetan Machine Uni","Times","Times Bold","Times New Roman","Times New Roman Baltic","Times New Roman Bold","Times New Roman Italic","Times Roman","Tlwg Mono","Tlwg Typewriter","Tlwg Typist","Tlwg Typo","TlwgMono","TlwgTypewriter","Toledo","Traditional Arabic","Trajan Pro","Trattatello","Trebuchet MS","Trebuchet MS Bold","Tunga","Tw Cen MT","Tw Cen MT Bold","Tw Cen MT Italic","URW Bookman L","URW Chancery L","URW Gothic L","URW Palladio L","Ubuntu","Ubuntu Condensed","Ubuntu Mono","Ukai","Ume Gothic","Ume Mincho","Ume P Gothic","Ume P Mincho","Ume UI Gothic","Uming","Umpush","UnBatang","UnDinaru","UnDotum","UnGraphic","UnGungseo","UnPilgi","Untitled1","Urdu Typesetting","Uroob","Utkal","Utopia","Utsaah","Valken","Vani","Vemana2000","Verdana","Verdana Bold","Vijaya","Viner Hand ITC","Vivaldi","Vivian","Vladimir Script","Vrinda","Waree","Waseem","Waverly","Webdings","WenQuanYi Bitmap Song","WenQuanYi Micro Hei","WenQuanYi Micro Hei Mono","WenQuanYi Zen Hei","Whimsy TT","Wide Latin","Wingdings","Wingdings 2","Wingdings 3","Woodcut","X-Files","Year supply of fairy cakes","Yu Gothic","Yu Mincho","Yuppy SC","Yuppy SC Regular","Yuppy TC","Yuppy TC Regular","Zapf Dingbats","Zapfino","Zawgyi-One","gargi","lklug","mry_KacstQurn","ori1Uni"]'
-SUPPORTED_FONTS: List[str] = json.loads(font_str)
+SUPPORTED_FONTS: list[str] = json.loads(font_str)
diff --git a/generalresearch/grliq/models/custom_types.py b/generalresearch/grliq/models/custom_types.py
index 1b2c9de..c5eb93c 100644
--- a/generalresearch/grliq/models/custom_types.py
+++ b/generalresearch/grliq/models/custom_types.py
@@ -1,5 +1,5 @@
-from typing_extensions 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 94802cc..4464a7f 100644
--- a/generalresearch/grliq/models/decider.py
+++ b/generalresearch/grliq/models/decider.py
@@ -1,6 +1,7 @@
+from __future__ import annotations
+
from datetime import datetime, timezone
from enum import Enum
-from typing import Optional
from pydantic import BaseModel, ConfigDict, Field
@@ -41,13 +42,13 @@ class GrlIqAttemptResult(BaseModel):
description="Whether an attempt should be allowed to continue, based on the evidence"
"available to the decider at this point in time"
)
- fraud_score: Optional[int] = Field(
+ fraud_score: int | None = Field(
ge=0,
le=100,
description="Higher equals more likely to be fraudulent",
default=None,
)
- fingerprint: Optional[str] = Field(
+ fingerprint: str | None = Field(
default=None,
description="Fingerprint that should be unique to this particular device",
)
diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py
index 70cc94d..f8bdd98 100644
--- a/generalresearch/grliq/models/forensic_data.py
+++ b/generalresearch/grliq/models/forensic_data.py
@@ -1,10 +1,12 @@
+from __future__ import annotations
+
import hashlib
import re
from collections import Counter
from datetime import datetime, timedelta, timezone
from enum import Enum
from functools import cached_property
-from typing import Any, Dict, List, Literal, Optional, Set
+from typing import Any, Literal
from uuid import uuid4
import pycountry
@@ -40,10 +42,10 @@ from generalresearch.grliq.models.forensic_result import (
Phase,
)
from generalresearch.grliq.models.useragents import (
+ BrowserFamily,
GrlUserAgent,
OSFamily,
UserAgentHints,
- BrowserFamily,
)
from generalresearch.models.custom_types import (
AwareDatetimeISO,
@@ -130,28 +132,28 @@ class GrlIqData(BaseModel):
# --- Attributes on the db table directly ---
- id: Optional[BigAutoInteger] = Field(default=None, exclude=True)
+ id: BigAutoInteger | None = Field(default=None, exclude=True)
uuid: UUIDStr = Field(
default_factory=lambda: uuid4().hex,
description="A unique identifier for this data object",
examples=[uuid4().hex],
)
- mid: Optional[UUIDStr] = Field(
+ mid: UUIDStr | None = Field(
description="The mid the of the User's attempt (thl-session) that "
"is associated with this data",
examples=[uuid4().hex],
)
- phase: Optional[Phase] = Field(
+ phase: Phase | None = Field(
description="The phase of a thl-session in which this data was collected",
default=Phase.OFFERWALL_ENTER,
)
- product_id: Optional[UUIDStr] = Field(
+ product_id: UUIDStr | None = Field(
default=None,
description="The Brokerage Product ID (BPID)",
examples=[uuid4().hex],
)
- product_user_id: Optional[str] = Field(
+ product_user_id: str | None = Field(
default=None,
description="The Brokerage Product User ID (BPUID).",
examples=["test-user-2dbeaaf4"],
@@ -167,7 +169,7 @@ class GrlIqData(BaseModel):
description="This comes from the actual web request's headers",
examples=["72.39.217.116"],
)
- client_ip_detail: Optional[GeoIPInformation] = Field(default=None)
+ client_ip_detail: GeoIPInformation | None = Field(default=None)
created_at: AwareDatetimeISO = Field(
description="When we actually received this data. The timestamp field "
@@ -175,16 +177,16 @@ class GrlIqData(BaseModel):
"by a baddie."
)
- request_headers: Dict = Field(
+ request_headers: dict = Field(
description="The full request headers from the actual HTTP call that was made."
)
# data: Dict = Field()
# result_data: Dict = Field()
# fraud_score: int = Field()
- # is_attempt_allowed: Optional[bool] = Field(default=None)
- results: Optional[GrlIqCheckerResults] = Field(default=None)
- category_result: Optional[GrlIqForensicCategoryResult] = Field(
+ # is_attempt_allowed: bool | None = Field(default=None)
+ results: GrlIqCheckerResults | None = Field(default=None)
+ category_result: GrlIqForensicCategoryResult | None = Field(
default=None, description="Saved in the database as a jsonb"
)
@@ -212,19 +214,19 @@ class GrlIqData(BaseModel):
"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
]
)
- user_agent_str_2: Optional[str] = Field(
+ user_agent_str_2: str | None = Field(
description="This will only be set if different than user_agent_str"
)
- user_agent_hints: Optional[UserAgentHints] = Field(
+ user_agent_hints: UserAgentHints | None = Field(
description="Comes from the User-Agent Client Hints API", default=None
)
platform: Platform = Field(description="navigator.platform")
- platform_2: Optional[Platform] = Field(description="navigator.platform")
- platform_3: Optional[Platform] = Field()
+ platform_2: Platform | None = Field(description="navigator.platform")
+ platform_3: Platform | None = Field()
language: str = Field(examples=["en-US"])
language_2: str = Field(examples=["en-US"])
- language_3: Optional[str] = Field()
+ language_3: str | None = Field()
calender_locale: str = Field(examples=["en-US"])
screen_width: NonNegativeInt = Field()
@@ -241,10 +243,10 @@ class GrlIqData(BaseModel):
app_name: Literal["Netscape"] = Field(
description="Navigator.appName. Always 'Netscape'"
)
- product_sub: Optional[Literal["20030107", "20100101"]] = Field(
+ product_sub: Literal["20030107", "20100101"] | None = Field(
description="Navigator.productSub"
)
- vendor: Optional[Literal["Apple Computer, Inc.", "Google Inc.", "NAVER Corp."]] = (
+ vendor: Literal["Apple Computer, Inc.", "Google Inc.", "NAVER Corp."] | None = (
Field(description="Navigator.vendor.")
)
@@ -253,14 +255,14 @@ class GrlIqData(BaseModel):
webrtc_is_supported: PassFailError = Field()
webrtc_error: bool = Field()
webrtc_local_ip: str = Field()
- webrtc_ip: Optional[IPvAnyAddressStr] = Field(examples=[fake.ipv4_public()])
- webrtc_ip_detail: Optional[GeoIPInformation] = Field(default=None)
+ webrtc_ip: IPvAnyAddressStr | None = Field(examples=[fake.ipv4_public()])
+ webrtc_ip_detail: GeoIPInformation | None = Field(default=None)
- hardware_concurrency: Optional[int] = Field(
+ hardware_concurrency: int | None = Field(
description="Sometimes this is an empty str"
)
- hardware_concurrency_2: Optional[int] = Field()
- hardware_concurrency_3: Optional[int] = Field()
+ hardware_concurrency_2: int | None = Field()
+ hardware_concurrency_3: int | None = Field()
# Browser/session properties
navigator_java_enabled: bool = Field()
@@ -340,13 +342,13 @@ class GrlIqData(BaseModel):
"'de355917bf33e0789539450797b843f9|5' (windows, iphone, mac) or '|0' (typically android). "
)
chrome_extensions: str = Field(description="comma sep str of chrome extensions")
- audio_codecs: Optional[str] = Field(
+ audio_codecs: str | None = Field(
examples=["1,1,1,1,1,3,1,3,1,3,3,1,1,3,3,3,3,1,3,3,3,2,1,1"],
description="canPlayType: {'3': probably, '2': maybe, '1': no, 0: error}",
min_length=47,
max_length=47,
)
- video_codecs: Optional[str] = Field(
+ video_codecs: str | None = Field(
examples=["1,3,3,3,3,3,3,3,3,3,1,1,1,1,1,1,3,1,1,1,3,3,1"],
description="canPlayType: {'3': probably, '2': maybe, '1': no, 0: error}",
min_length=45,
@@ -380,13 +382,13 @@ class GrlIqData(BaseModel):
ontouchstart: bool = Field()
# todo: confirm this is 5 for an iphone
max_touch_points: int = Field()
- navigator_deviceMemory: Optional[float] = Field()
+ navigator_deviceMemory: float | None = Field()
memory_jsHeapSizeLimit: int = Field()
navigator_mediaDevices_len: int = Field()
unmasked_vendor_webgl: str = Field()
unmasked_renderer_webgl: str = Field()
keyboard_detected: bool = Field()
- keyboard_layout_size: Optional[int] = Field(description="mobile safari None?")
+ keyboard_layout_size: int | None = Field(description="mobile safari None?")
window_orientation: int = Field(description="0 or 1. idk which is which")
# Session properties
@@ -394,25 +396,23 @@ class GrlIqData(BaseModel):
# *different* each (to make sure they aren't reusing posts)
execution_time_ms: float = Field()
performance_loop_time: float = Field()
- connection_rtt: Optional[int] = Field()
- connection_downlink: Optional[float] = Field()
+ connection_rtt: int | None = Field()
+ connection_downlink: float | None = Field()
connection_type: str = Field()
connection_effectiveType: str = Field()
# fingerprint stuff
canvas_support_level: SupportLevel = Field()
- canvas_hash: Optional[Hash128] = Field(
- description="dfiq's canvas image fingerprint"
- )
- canvas_hash_2: Optional[Hash128] = Field(
+ canvas_hash: Hash128 | None = Field(description="dfiq's canvas image fingerprint")
+ canvas_hash_2: Hash128 | None = Field(
description="simpler canvas image fingerprint stolen from amiunique.org",
default=None,
)
- webgl_hash: Optional[Hash128] = Field(
+ webgl_hash: Hash128 | None = Field(
description="DFIQ's version of webgl hash. It has stuff included in the hash: anisotropy, supported "
"extensions, etc."
)
- webgl_context: Optional[
+ webgl_context: (
Literal[
"webgl2",
"webgl",
@@ -421,17 +421,18 @@ class GrlIqData(BaseModel):
"webkit-3d",
"moz-webgl",
]
- ] = Field(default=None)
- webgl_max_anisotropy: Optional[int] = Field(default=None, examples=[16])
- webgl_shading_language_version: Optional[str] = Field(
+ | None
+ ) = Field(default=None)
+ webgl_max_anisotropy: int | None = Field(default=None, examples=[16])
+ webgl_shading_language_version: str | None = Field(
default=None,
examples=["WebGL GLSL ES 3.00 (OpenGL ES GLSL ES 3.0 Chromium)"],
)
- webgl_hash_2: Optional[Hash128] = Field(
+ webgl_hash_2: Hash128 | None = Field(
description="hash128 of the canvas image, without additional stuff concatenated to it",
default=None,
)
- webgl_extensions: Optional[str] = Field(
+ webgl_extensions: str | None = Field(
description="pipe-separated list of webgl extensions",
default=None,
examples=[
@@ -439,22 +440,22 @@ class GrlIqData(BaseModel):
],
)
- audio_context_hash: Optional[Hash128] = Field()
- audio_intensity_fingerprint: Optional[float] = Field()
- audio_compressor_reduction: Optional[float] = Field()
- speech_synthesis_voice_hash: Optional[Hash128] = Field()
+ audio_context_hash: Hash128 | None = Field()
+ audio_intensity_fingerprint: float | None = Field()
+ audio_compressor_reduction: float | None = Field()
+ speech_synthesis_voice_hash: Hash128 | None = Field()
speech_synthesis_avail_voices_count: int = Field()
path_fingerprint: int = Field(
description="maybe is consistent? sum of pixel value of some path."
)
- text_2d_fingerprint: Optional[Hash128] = Field()
+ text_2d_fingerprint: Hash128 | None = Field()
canvas_fingerprint: int = Field()
# User Preferences
- color_gamut: Optional[Literal["1", "2", "3", "0"]] = Field(
+ color_gamut: Literal["1", "2", "3", "0"] | None = Field(
description="{'1': 'rec2020', '2':'p3', '3':'srgb', '0': none} # p3 typically used in macbooks n stuff"
)
- prefers_contrast: Optional[Literal["0", "1", "2", "3", "4", "5", "9"]] = Field(
+ prefers_contrast: Literal["0", "1", "2", "3", "4", "5", "9"] | None = Field(
description="{'no-preference': 0, 'high': 1, 'more': 2, 'low': 3, 'less': 4, 'forced': 5, None: 9}"
)
prefers_reduced_motion: bool = Field(description="reduce (1) vs no-preference (0)")
@@ -464,12 +465,12 @@ class GrlIqData(BaseModel):
prefers_color_scheme: bool = Field(description="{'dark': 1, '?': 0}")
# Battery Info
- battery_charging: Optional[bool] = Field(default=None)
- battery_charging_time: Optional[float] = Field(default=None)
- battery_discharging_time: Optional[float] = Field(default=None)
- battery_level: Optional[float] = Field(default=None, ge=0, le=1)
+ battery_charging: bool | None = Field(default=None)
+ battery_charging_time: float | None = Field(default=None)
+ battery_discharging_time: float | None = Field(default=None)
+ battery_level: float | None = Field(default=None, ge=0, le=1)
- supported_fonts_str: Optional[str] = Field(
+ supported_fonts_str: str | None = Field(
default=None,
description="Bit-packed string for font support. Each element is 32 bits, with each bit representing T/F for "
"font support.",
@@ -481,7 +482,7 @@ class GrlIqData(BaseModel):
)
# Time it took for the client to download the logo.jpg
- logo_download_ms: Optional[float] = Field(default=None, gt=0)
+ logo_download_ms: float | None = Field(default=None, gt=0)
# --- Not from post body ----
@@ -490,15 +491,15 @@ class GrlIqData(BaseModel):
)
# Can optionally be loaded from the grliq_forensicevents table
- events: Optional[List[Dict]] = Field(default=None)
- pointer_move_events: Optional[List[PointerMove]] = Field(default=None)
- mouse_events: Optional[List[MouseEvent]] = Field(default=None)
- keyboard_events: Optional[List[KeyboardEvent]] = Field(default=None)
+ events: list[dict] | None = Field(default=None)
+ pointer_move_events: list[PointerMove] | None = Field(default=None)
+ mouse_events: list[MouseEvent] | None = Field(default=None)
+ keyboard_events: list[KeyboardEvent] | None = Field(default=None)
- timing_data: Optional[TimingData] = Field(default=None)
+ timing_data: TimingData | None = Field(default=None)
@property
- def session_uuid(self) -> Optional[UUIDStr]:
+ def session_uuid(self) -> UUIDStr | None:
return self.mid
@cached_property
@@ -506,7 +507,7 @@ class GrlIqData(BaseModel):
return GrlUserAgent.from_ua_str(self.user_agent_str)
@cached_property
- def fingerprint_keys(self) -> List[str]:
+ def fingerprint_keys(self) -> list[str]:
fp_cols = [
"country_iso",
"canvas_hash",
@@ -557,7 +558,7 @@ class GrlIqData(BaseModel):
return hashlib.md5(s.encode()).hexdigest()
@cached_property
- def audio_codecs_named(self) -> Dict[str, bool]:
+ def audio_codecs_named(self) -> dict[str, bool]:
return dict(
zip(
AUDIO_CODEC_NAMES,
@@ -566,7 +567,7 @@ class GrlIqData(BaseModel):
)
@cached_property
- def video_codecs_named(self) -> Dict[str, bool]:
+ def video_codecs_named(self) -> dict[str, bool]:
return dict(
zip(
VIDEO_CODEC_NAMES,
@@ -584,13 +585,13 @@ class GrlIqData(BaseModel):
)[-len(SUPPORTED_FONTS) :]
@cached_property
- def supported_fonts(self) -> Set[str]:
+ def supported_fonts(self) -> set[str]:
return {
f for x, f in zip(self.supported_fonts_binary, SUPPORTED_FONTS) if x == "1"
}
@cached_property
- def audio_intensity_rounded(self) -> Optional[float]:
+ def audio_intensity_rounded(self) -> float | None:
# The audio intensity fingerprint seems to be purposely manipulated
# to add randomness, but the level of randomness if very low, past
# the 6th decimal point.
@@ -622,7 +623,7 @@ class GrlIqData(BaseModel):
mode="before",
)
@classmethod
- def str_to_int_or_null(cls, value: str) -> Optional[int]:
+ def str_to_int_or_null(cls, value: str) -> int | None:
return int(value) if value not in {None, ""} else None
@field_validator(
@@ -633,7 +634,7 @@ class GrlIqData(BaseModel):
mode="before",
)
@classmethod
- def str_to_float_or_null(cls, value: str) -> Optional[float]:
+ def str_to_float_or_null(cls, value: str) -> float | None:
return float(value) if value not in {None, ""} else None
@field_validator(
@@ -649,7 +650,7 @@ class GrlIqData(BaseModel):
mode="before",
)
@classmethod
- def str_or_null(cls, value: str) -> Optional[str]:
+ def str_or_null(cls, value: str) -> str | None:
return value or None
@field_validator(
@@ -724,7 +725,7 @@ class GrlIqData(BaseModel):
@field_validator("platform", "platform_2", "platform_3", mode="before")
@classmethod
- def platform_enum_or_other(cls, value: Optional[str]) -> Optional[Platform]:
+ def platform_enum_or_other(cls, value: str | None) -> Platform | None:
if value is None or value == "":
return None
try:
@@ -734,7 +735,7 @@ class GrlIqData(BaseModel):
@field_validator("webrtc_ip", mode="before")
@classmethod
- def preprocess_ip(cls, ip: str) -> Optional[str]:
+ def preprocess_ip(cls, ip: str) -> str | None:
# Strip square brackets if present
return re.sub(r"^\[|\]$", "", ip) if ip else None
@@ -781,7 +782,7 @@ class GrlIqData(BaseModel):
return None
- def model_dump_sql(self, **kwargs) -> Dict[str, Any]:
+ def model_dump_sql(self, **kwargs) -> dict[str, Any]:
d = dict()
d["uuid"] = self.uuid
d["session_uuid"] = self.mid
@@ -802,7 +803,7 @@ class GrlIqData(BaseModel):
return d
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> Self:
+ def from_db(cls, d: dict[str, Any]) -> Self:
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 d89f681..3fe6481 100644
--- a/generalresearch/grliq/models/forensic_result.py
+++ b/generalresearch/grliq/models/forensic_result.py
@@ -1,7 +1,6 @@
from __future__ import annotations
from enum import Enum
-from typing import List, Optional, Set
from uuid import uuid4
from pydantic import BaseModel, ConfigDict, Field, computed_field
@@ -43,13 +42,13 @@ class GrlIqForensicCategoryResult(BaseModel):
model_config = ConfigDict(extra="forbid", validate_assignment=True)
- uuid: Optional[UUIDStr] = Field(
+ uuid: UUIDStr | None = Field(
description="The uuid for the GrlIqData model these results are based on",
default=None,
examples=[uuid4().hex],
)
- updated_at: Optional[AwareDatetimeISO] = Field(default=None)
+ updated_at: AwareDatetimeISO | None = Field(default=None)
is_complete: bool = Field(
description="This is based on whether or not the GrlIqCheckerResults"
"object that this data was based on was complete at that time.",
@@ -108,7 +107,7 @@ class GrlIqForensicCategoryResult(BaseModel):
)
@staticmethod
- def model_score_fields() -> List[str]:
+ def model_score_fields() -> list[str]:
return [
"is_bot",
"is_velocity",
@@ -143,7 +142,7 @@ class GrlIqCheckerResult(BaseModel):
model_config = ConfigDict(extra="forbid", validate_assignment=True)
score: GrlIqScore = Field(default=0)
- msg: Optional[str] = Field(default=None)
+ msg: str | None = Field(default=None)
@property
def passes(self) -> bool:
@@ -170,27 +169,27 @@ class GrlIqObservations(BaseModel):
default=0, description="Count of unique timezones (by IP) (past 30 days)"
)
- paste_event_count: Optional[int] = Field(
+ paste_event_count: int | None = Field(
default=None, description="Count of paste events (user pasted in text)"
)
- visibilitychange_event_count: Optional[int] = Field(
+ visibilitychange_event_count: int | None = Field(
default=None,
description="Count of visibilitychange events (entire page isn't visible)",
)
- blur_event_count: Optional[int] = Field(
+ blur_event_count: int | None = Field(
default=None, description="Count of blur events (page lost focus)"
)
- devicemotion_event_count: Optional[int] = Field(
+ devicemotion_event_count: int | None = Field(
default=None,
description="Count of devicemotion events (device gyroscope motion)",
)
- click_event_count: Optional[int] = Field(
+ click_event_count: int | None = Field(
default=None,
description="Count of click events (any pointer type)",
)
# all clicks are marked as pointerType = 'mouse', but other pointermove events have a pointerType
# of 'touch' or 'mouse'
- pointermove_pointer_types: Optional[Set[str]] = Field(
+ pointermove_pointer_types: set[str] | None = Field(
default=None, description="pointer types"
)
@@ -203,15 +202,15 @@ class GrlIqCheckerResults(BaseModel):
model_config = ConfigDict(extra="forbid", validate_assignment=True)
- uuid: Optional[UUIDStr] = Field(
+ uuid: UUIDStr | None = Field(
description="The uuid for the GrlIqData model these results are based on",
default=None,
examples=[uuid4().hex],
)
- updated_at: Optional[AwareDatetimeISO] = Field(default=None)
+ updated_at: AwareDatetimeISO | None = Field(default=None)
- observations: Optional[GrlIqObservations] = Field(default=None)
+ observations: GrlIqObservations | None = Field(default=None)
# browser_props
check_environment: GrlIqCheckerResult = Field()
@@ -268,16 +267,16 @@ class GrlIqCheckerResults(BaseModel):
check_ip_webrtc_ip_detail: GrlIqCheckerResult = Field()
# websocket (events)
- check_page_load_events: Optional[GrlIqCheckerResult] = Field(default=None)
- check_grliq_events: Optional[GrlIqCheckerResult] = Field(default=None)
- check_pasting: Optional[GrlIqCheckerResult] = Field(default=None)
- check_pointer_movements: Optional[GrlIqCheckerResult] = Field(default=None)
- check_device_motion: Optional[GrlIqCheckerResult] = Field(default=None)
- check_pointer_type: Optional[GrlIqCheckerResult] = Field(default=None)
- check_for_bad_events: Optional[GrlIqCheckerResult] = Field(default=None)
+ check_page_load_events: GrlIqCheckerResult | None = Field(default=None)
+ check_grliq_events: GrlIqCheckerResult | None = Field(default=None)
+ check_pasting: GrlIqCheckerResult | None = Field(default=None)
+ check_pointer_movements: GrlIqCheckerResult | None = Field(default=None)
+ check_device_motion: GrlIqCheckerResult | None = Field(default=None)
+ check_pointer_type: GrlIqCheckerResult | None = Field(default=None)
+ check_for_bad_events: GrlIqCheckerResult | None = Field(default=None)
# websocket (ping)
- check_average_rtt: Optional[GrlIqCheckerResult] = Field(default=None)
+ check_average_rtt: GrlIqCheckerResult | None = Field(default=None)
# todo: we might also have a "fingerprint" in here ???
diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py
index a40e103..b0112e7 100644
--- a/generalresearch/grliq/models/forensic_summary.py
+++ b/generalresearch/grliq/models/forensic_summary.py
@@ -2,11 +2,8 @@ from __future__ import annotations
import random
from typing import (
- Dict,
List,
Literal,
- Optional,
- Tuple,
Union,
get_args,
get_origin,
@@ -49,25 +46,25 @@ class UserForensicSummary(BaseModel):
model_config = ConfigDict(extra="forbid", validate_assignment=True)
- period_start: Optional[AwareDatetimeISO] = Field(
+ period_start: AwareDatetimeISO | None = Field(
default=None,
description="Timestamp of the earliest attempt included in this summary (UTC)",
)
- period_end: Optional[AwareDatetimeISO] = Field(
+ period_end: AwareDatetimeISO | None = Field(
default=None,
description="Timestamp of the latest attempt included in this summary (UTC)",
)
# These must be nullable in case a user has 0 attempts!
- category_result_summary: Optional[GrlIqForensicCategorySummary] = Field(
+ category_result_summary: GrlIqForensicCategorySummary | = Field(
default=None
)
- checker_result_summary: Optional[GrlIqCheckerResultsSummary] = Field(default=None)
+ checker_result_summary: GrlIqCheckerResultsSummary | None = Field(default=None)
- country_timing_data_summary: Dict[CountryISO, TimingDataCountrySummary] = Field(
+ country_timing_data_summary: dict[CountryISO, TimingDataCountrySummary] = Field(
default_factory=dict
)
- ip_timing_data_summary: Dict[IPvAnyAddressStr, IPTimingDataSummary] = Field(
+ ip_timing_data_summary: dict[IPvAnyAddressStr, IPTimingDataSummary] = Field(
default_factory=dict
)
@@ -102,7 +99,7 @@ class GrlIqForensicCategorySummary(BaseModel):
platform_ip_inconsistent_avg: GrlIqAvgScore = Field(
examples=[random.randint(0, 100)]
)
- fraud_score_avg: Optional[GrlIqAvgScore] = Field(
+ fraud_score_avg: GrlIqAvgScore | None = Field(
default=None, examples=[random.randint(0, 100)]
)
@@ -171,7 +168,7 @@ class TimingDataCountrySummary(BaseModel):
rtt_q75: float = Field(gt=0, examples=[220.232])
rtt_max: float = Field(gt=0, examples=[890.006])
- expected_rtt_range: Tuple[float, float] = Field(
+ 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)],
)
@@ -191,8 +188,8 @@ 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: Optional[UserType] = Field(default=None, examples=[UserType.RESIDENTIAL])
- expected_rtt_range: Tuple[float, float] = Field(
+ user_type: UserType | None = Field(default=None, examples=[UserType.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)],
)
@@ -213,13 +210,13 @@ class CountryRTTDistribution(BaseModel):
description="Country client_ip is located in", examples=["fr"]
)
# For users marked as fraud or not
- is_fraud: Optional[bool] = 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: Optional[UserType] = Field(
+ user_type: UserType|None = Field(
default=None,
description="user_type of the client_ip as determined by MaxMind",
examples=[UserType.RESIDENTIAL],
@@ -243,7 +240,7 @@ class CountryRTTDistribution(BaseModel):
description="The 95% confidence interval calculated in log-space",
)
@property
- def expected_rtt_range(self) -> Tuple[float, float]:
+ def expected_rtt_range(self) -> tuple[float, float]:
# This is the log_mean +- 2 log_std, then converted back to non-log space.
# This is not just the mean + 2x std b/c we calculate the expected
# range in log-space (due to high skewness)
diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py
index 6cea2dc..1953f6d 100644
--- a/generalresearch/grliq/models/useragents.py
+++ b/generalresearch/grliq/models/useragents.py
@@ -1,6 +1,7 @@
+from __future__ import annotations
+
import hashlib
from enum import Enum
-from typing import Dict, List, Optional
from pydantic import BaseModel, ConfigDict, Field, field_validator
from typing_extensions import Self
@@ -105,7 +106,7 @@ mobile_families = {
class OSInfo(BaseModel):
family: OSFamily = Field()
- version_string: Optional[str] = Field()
+ version_string: str | None = Field()
@field_validator("family", mode="before")
@classmethod
@@ -118,7 +119,7 @@ class OSInfo(BaseModel):
class BrowserInfo(BaseModel):
family: BrowserFamily = Field()
- version_string: Optional[str] = Field()
+ version_string: str | None = Field(default=None)
@field_validator("family", mode="before")
@classmethod
@@ -182,7 +183,7 @@ class GrlUserAgent(BaseModel):
is_bot: bool = Field()
@property
- def ua_string_values(self) -> Dict[str, str]:
+ 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()
@@ -236,11 +237,11 @@ class UserAgentHints(BaseModel):
extra="forbid", validate_assignment=True, populate_by_name=True
)
- brands: Optional[List[Dict]] = Field(validation_alias="b", default=None)
- brands_full: Optional[List[Dict]] = Field(validation_alias="fv", default=None)
+ brands: list[dict] | None = Field(validation_alias="b", default=None)
+ brands_full: list[dict] | None = Field(validation_alias="fv", default=None)
mobile: bool = Field(validation_alias="m", default=False)
- model: Optional[str] = Field(validation_alias="md", default=None)
- platform: Optional[str] = Field(validation_alias="o", default=None)
- platform_version: Optional[str] = Field(validation_alias="ov", default=None)
- architecture: Optional[str] = Field(validation_alias="a", default=None)
- bitness: Optional[str] = Field(validation_alias="bt", default=None)
+ model: str | None = Field(validation_alias="md", default=None)
+ platform: str | None = Field(validation_alias="o", default=None)
+ platform_version: str | None = Field(validation_alias="ov", default=None)
+ architecture: str | None = Field(validation_alias="a", default=None)
+ bitness: str | None = Field(validation_alias="bt", default=None)
diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py
index c772e7f..95390a8 100644
--- a/generalresearch/grliq/utils.py
+++ b/generalresearch/grliq/utils.py
@@ -1,7 +1,8 @@
+from __future__ import annotations
+
import os
from datetime import datetime, timezone
from pathlib import Path
-from typing import Optional, Union
from uuid import UUID
# from generalresearch.config import
@@ -10,11 +11,11 @@ from generalresearch.models.custom_types import UUIDStr
def get_screenshot_fp(
created_at: datetime,
- forensic_uuid: Union[UUIDStr, UUID],
+ forensic_uuid: UUIDStr | UUID,
grliq_archive_dir: Path = "/tmp",
grliq_ss_dir_name: str = "canvas2html",
create_dir_if_not_exists: bool = True,
-) -> Optional[Path]:
+) -> Path | None:
assert created_at.tzinfo == timezone.utc
if isinstance(forensic_uuid, UUID):
diff --git a/generalresearch/grpc.py b/generalresearch/grpc.py
index bf3f0e2..040fd26 100644
--- a/generalresearch/grpc.py
+++ b/generalresearch/grpc.py
@@ -1,5 +1,6 @@
-from datetime import datetime, timedelta
-from typing import Optional
+from __future__ import annotations
+
+from datetime import datetime, timedelta, timezone
from google.protobuf.duration_pb2 import Duration
from google.protobuf.timestamp_pb2 import Timestamp
@@ -11,7 +12,7 @@ def timestamp_from_datetime(dt: datetime) -> Timestamp:
return ts
-def timestamp_from_datetime_nullable(dt: Optional[datetime]) -> Timestamp:
+def timestamp_from_datetime_nullable(dt: datetime | None) -> Timestamp:
ts = Timestamp()
if dt:
ts.FromDatetime(dt)
@@ -19,17 +20,17 @@ def timestamp_from_datetime_nullable(dt: Optional[datetime]) -> Timestamp:
def timestamp_to_datetime(ts: Timestamp) -> datetime:
- return datetime.utcfromtimestamp(ts.seconds + ts.nanos / 1e9)
+ return datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc)
-def timestamp_to_datetime_nullable(ts: Timestamp) -> Optional[datetime]:
+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.utcfromtimestamp(0)
- d = datetime.utcfromtimestamp(ts.seconds + ts.nanos / 1e9)
+ default = datetime.fromtimestamp(0, tz=timezone.utc)
+ d = datetime.fromtimestamp(ts.seconds + ts.nanos / 1e9, tz=timezone.utc)
return None if d == default else d
-def timestamp_to_json_nullable(ts: Timestamp) -> Optional[str]:
+def timestamp_to_json_nullable(ts: Timestamp) -> str | None:
# 1) grpc converts a null timestamp to '1970-01-01T00:00:00Z'. Not what we want...
# 2) grpc uses different formatting for the microseconds depending on if it's divisible by 0, 3, or 6 digits.
# I don't understand why anyone would want to do this ...
diff --git a/generalresearch/healing_ppe.py b/generalresearch/healing_ppe.py
index e8ee144..254a893 100644
--- a/generalresearch/healing_ppe.py
+++ b/generalresearch/healing_ppe.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import logging
import os
import signal
@@ -5,7 +7,6 @@ import time
from collections import defaultdict
from concurrent import futures
from concurrent.futures.process import BrokenProcessPool
-from typing import Optional
logger = logging.getLogger()
@@ -17,8 +18,8 @@ signal_int_name = defaultdict(
class HealingProcessPoolExecutor:
def __init__(
self,
- max_workers: Optional[int] = None,
- name: Optional[str] = None,
+ max_workers: int | None = None,
+ name: str | None = None,
):
if not name:
try:
diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py
index 9049e8f..a14b438 100644
--- a/generalresearch/incite/collections/__init__.py
+++ b/generalresearch/incite/collections/__init__.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import logging
import os
import subprocess
@@ -5,7 +7,7 @@ import time
from datetime import datetime
from enum import Enum
from sys import platform
-from typing import Any, Dict, List, Optional
+from typing import Any
import dask
import dask.dataframe as dd
@@ -167,10 +169,10 @@ class DFCollectionItem(CollectionItemBase):
return self.to_archive(ddf=dd.from_pandas(_df, npartitions=1), is_partial=True)
# --- ORM / Data handlers---
- def to_dict(self, *args, **kwargs) -> Dict[str, Any]:
+ def to_dict(self) -> dict[str, Any]:
return self._to_dict()
- def from_mysql(self, since: Optional[datetime] = None) -> Optional[pd.DataFrame]:
+ 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
@@ -181,9 +183,7 @@ class DFCollectionItem(CollectionItemBase):
else:
return self.from_postgres_standard(since=since)
- def from_mysql_standard(
- self, since: Optional[datetime] = None
- ) -> Optional[pd.DataFrame]:
+ def from_mysql_standard(self, since: datetime | None = None) -> pd.DataFrame | None:
assert (
self._collection.data_type != DFCollectionType.LEDGER
@@ -235,8 +235,8 @@ class DFCollectionItem(CollectionItemBase):
return df
def from_postgres_standard(
- self, since: Optional[datetime] = None
- ) -> Optional[pd.DataFrame]:
+ self, since: datetime | None = None
+ ) -> pd.DataFrame | None:
assert (
self._collection.data_type != DFCollectionType.LEDGER
), "Can't call from_postgres_standard for Ledger DFCollectionItem"
@@ -285,7 +285,7 @@ class DFCollectionItem(CollectionItemBase):
return df
- def from_postgres_ledger(self) -> Optional[pd.DataFrame]:
+ def from_postgres_ledger(self) -> pd.DataFrame | None:
assert (
self._collection.data_type == DFCollectionType.LEDGER
), "Can only call from_postgres_ledger on Ledger DFCollectionItem"
@@ -399,7 +399,7 @@ class DFCollectionItem(CollectionItemBase):
"""
assert isinstance(ddf, dd.DataFrame), "must pass dask df"
- client: Optional[Client] = self._collection._client
+ client: Client | None = self._collection._client
# client = None
if client:
@@ -419,7 +419,7 @@ class DFCollectionItem(CollectionItemBase):
def _to_archive(
self,
- ddf: dd.DataFrame,
+ ddf: dd.DataFrame | None,
is_empty: bool,
overwrite: bool = False,
) -> bool:
@@ -503,7 +503,7 @@ class DFCollectionItem(CollectionItemBase):
subprocess.call(["mv", "-T", tmp_path.as_posix(), self.path.as_posix()])
return True
- def to_archive_numbered_partial(self, ddf: Optional[dd.DataFrame] = None) -> bool:
+ 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
@@ -516,7 +516,7 @@ class DFCollectionItem(CollectionItemBase):
collection = self._collection
schema = collection._schema
- client: Optional[Client] = collection._client
+ client: Client | None = collection._client
next_numbered_path = self.next_numbered_path(self.partial_path)
partial_path = self.partial_path
@@ -572,7 +572,7 @@ class DFCollectionItem(CollectionItemBase):
assert self.should_archive(), "not ready to archive!"
- df: Optional[pd.DataFrame] = self.from_mysql()
+ df: pd.DataFrame | None = self.from_mysql()
if df is None:
self.set_empty()
@@ -589,11 +589,11 @@ class DFCollectionItem(CollectionItemBase):
class DFCollection(CollectionBase):
- data_type: Optional[DFCollectionType] = Field(default=None)
+ data_type: DFCollectionType | None = Field(default=None)
# --- Private ---
- pg_config: Optional[PostgresConfig] = Field(default=None)
- sql_helper: Optional[SqlHelper] = Field(default=None)
+ pg_config: PostgresConfig | None = Field(default=None)
+ sql_helper: SqlHelper | None = Field(default=None)
def __repr__(self):
res = self.signature() + "\n"
@@ -632,7 +632,7 @@ class DFCollection(CollectionBase):
# --- Properties ---
@property
- def items(self) -> List[DFCollectionItem]:
+ def items(self) -> list[DFCollectionItem]:
items = []
for iv in self.interval_range:
cm = DFCollectionItem(start=iv[0])
@@ -648,12 +648,12 @@ class DFCollection(CollectionBase):
def initial_load(
self,
- client: Optional[Client] = None,
+ client: Client | None = None,
sync: bool = True,
- since: Optional[datetime] = None,
- client_resources: Optional[Dict[str, Any]] = None,
- timeout: Optional[float] = None,
- ) -> List[Future]:
+ 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
@@ -692,7 +692,7 @@ class DFCollection(CollectionBase):
else:
return client.compute(fs, sync=True, priority=2, resources=client_resources)
- def fetch_force_rr_latest(self, sources) -> List[FilePath]:
+ def fetch_force_rr_latest(self, sources) -> list[FilePath]:
LOG.info(
f"{self.data_type.value}.fetch_force_rr_latest(sources={len(sources)})"
)
@@ -736,9 +736,9 @@ class DFCollection(CollectionBase):
def force_rr_latest(
self,
client: Client,
- client_resources: Optional[Dict[str, Any]] = None,
+ client_resources: dict[str, Any] | None = None,
sync: bool = True,
- ) -> List[Future]:
+ ) -> list[Future]:
# For forcing update of any partials asynchronously if desired
LOG.info(f"{self.data_type.value}.force_rr_latest({client=})")
diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py
index 22ac603..0698b82 100644
--- a/generalresearch/incite/mergers/__init__.py
+++ b/generalresearch/incite/mergers/__init__.py
@@ -4,13 +4,13 @@ import subprocess
from datetime import datetime, timezone
from enum import Enum
from sys import platform
-from typing import Optional, List, Type
+from typing import List, Optional, Type
import dask.dataframe as dd
import pandas as pd
from dask.distributed import Client
from pandera import DataFrameSchema
-from pydantic import Field, field_validator, ValidationInfo, model_validator
+from pydantic import Field, ValidationInfo, field_validator, model_validator
from typing_extensions import Self
from generalresearch.incite.base import CollectionBase, CollectionItemBase
@@ -28,9 +28,9 @@ from generalresearch.incite.schemas.mergers.foundations.user_id_product import (
UserIdProductSchema,
)
from generalresearch.incite.schemas.mergers.nginx import (
- NGINXGRSSchema,
NGINXCoreSchema,
NGINXFSBSchema,
+ NGINXGRSSchema,
)
from generalresearch.incite.schemas.mergers.pop_ledger import (
PopLedgerSchema,
diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py
index 9ce9d91..f7a45a8 100644
--- a/generalresearch/incite/mergers/foundations/__init__.py
+++ b/generalresearch/incite/mergers/foundations/__init__.py
@@ -1,5 +1,8 @@
+from __future__ import annotations
+
import logging
-from typing import Any, Collection, Dict, List
+from collections.abc import Collection
+from typing import Any
import pandas as pd
from more_itertools import chunked
@@ -28,7 +31,7 @@ def annotate_product_id(
assert len(user_ids) >= 1, "must have user_ids"
LOG.warning(f"annotate_product_id.len(user_ids): {len(user_ids)}")
- res: List[Dict[str, Any]] = []
+ res: list[dict[str, Any]] = []
with pg_config.make_connection() as conn:
for chunk in chunked(user_ids, chunksize):
try:
@@ -54,7 +57,7 @@ def annotate_product_id(
def lookup_product_and_team_id(
user_ids: Collection[int],
pg_config: PostgresConfig,
-) -> List[Dict[str, Any]]:
+) -> list[dict[str, Any]]:
user_ids = set(user_ids)
LOG.info(f"lookup_product_and_team_id: {len(user_ids)}")
@@ -64,7 +67,7 @@ def lookup_product_and_team_id(
assert len(user_ids) >= 1, "must have user_ids"
assert len(user_ids) <= 1000, "you should chunk this bro"
- res: List[Dict[str, Any]] = []
+ res: list[dict[str, Any]] = []
with pg_config.make_connection() as conn:
try:
with conn.cursor() as c:
@@ -108,7 +111,7 @@ def annotate_product_and_team_id(
assert len(user_ids) >= 1, "must have user_ids"
LOG.warning(f"annotate_product_and_team_id.len(user_ids): {len(user_ids)}")
- res: List[Dict[str, Any]] = []
+ res: list[dict[str, Any]] = []
with pg_config.make_connection() as conn:
for chunk in chunked(user_ids, chunksize):
try:
@@ -146,7 +149,7 @@ def annotate_product_user(
assert len(user_ids) >= 1, "must have user_ids"
LOG.warning(f"annotate_product_user.len(user_ids): {len(user_ids)}")
- res: List[Dict[str, Any]] = []
+ res: list[dict[str, Any]] = []
with pg_config.make_connection() as conn:
for chunk in chunked(user_ids, chunksize):
try:
diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py
index 4b87df7..32ce7ef 100644
--- a/generalresearch/incite/mergers/foundations/enriched_session.py
+++ b/generalresearch/incite/mergers/foundations/enriched_session.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import logging
from datetime import timedelta
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
+from typing import TYPE_CHECKING, Any, Literal
import dask.dataframe as dd
import pandas as pd
@@ -44,8 +46,8 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
session_coll: SessionDFCollection,
wall_coll: WallDFCollection,
pg_config: PostgresConfig,
- client: Optional[Client] = None,
- client_resources: Optional[Dict[str, Any]] = None,
+ client: Client | None = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
ir: pd.Interval = self.interval
@@ -55,7 +57,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
# Skip which already exist
if self.has_archive(include_empty=True):
- return None
+ return
# --- Session ---
LOG.warning(f"EnrichedSessionMergeItem: get session_collection")
@@ -64,12 +66,12 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
LOG.warning(f"EnrichedSessionMergeItem: no session items. set_empty.")
if self.should_archive():
self.set_empty()
- return None
+ return
if not (
session_items[-1].has_partial_archive() or session_items[-1].has_archive()
):
LOG.warning(f"EnrichedSessionMergeItem: session isn't updated!")
- return None
+ return
sddf = session_coll.ddf(
items=session_items,
@@ -94,7 +96,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
if len(wall_items) == 0:
LOG.error(f"EnrichedSessionMergeItem: no wall items")
- return None
+ return
wddf = wall_coll.ddf(
items=wall_items,
@@ -108,7 +110,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
)
if wddf is None:
- return None
+ return
attempt_cnt_ddf = (
wddf.groupby("session_id").size().rename("attempt_count").to_frame()
@@ -231,8 +233,8 @@ class EnrichedSessionMerge(MergeCollection):
self,
rr: "ReportRequest",
client: Client,
- product_ids: Optional[List[UUIDStr]] = None,
- user: Optional[User] = None,
+ product_ids: list[UUIDStr] | None = None,
+ user: User | None = None,
) -> pd.DataFrame:
"""
We don't have the concept of a Team yet so product_ids will be a list
diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
index d2a8aa1..f3ab8d8 100644
--- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
+++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import logging
-from typing import Any, Dict, Literal, Optional
+from typing import Any, Literal
import dask.dataframe as dd
import pandas as pd
@@ -40,7 +42,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
enriched_wall: EnrichedWallMerge,
pg_config: PostgresConfig,
client: Client,
- client_resources: Optional[Dict[str, Any]] = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
"""
TaskAdjustments are always partial because they could be revoked
@@ -61,7 +63,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
if len(task_adj_coll_items) == 0:
raise Exception("TaskAdjColl item collection failed")
- ddf: Optional[dd.DataFrame] = task_adj_coll.ddf(
+ ddf: dd.DataFrame | None = task_adj_coll.ddf(
items=task_adj_coll_items,
include_partial=True,
force_rr_latest=False,
diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py
index d69293b..bd77937 100644
--- a/generalresearch/incite/mergers/foundations/enriched_wall.py
+++ b/generalresearch/incite/mergers/foundations/enriched_wall.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import logging
from datetime import timedelta
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
+from typing import TYPE_CHECKING, Any, Literal
import dask.dataframe as dd
import pandas as pd
@@ -39,8 +41,8 @@ class EnrichedWallMergeItem(MergeCollectionItem):
wall_coll: WallDFCollection,
session_coll: SessionDFCollection,
pg_config: PostgresConfig,
- client: Optional[Client] = None,
- client_resources: Optional[Dict[str, Any]] = None,
+ client: Client | None = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
ir: pd.Interval = self.interval
@@ -50,7 +52,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
# Skip which already exist
if self.has_archive(include_empty=True):
- return None
+ return
# --- Wall ---
LOG.warning(f"EnrichedWallMergeItem: get wall_collection")
@@ -59,7 +61,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
LOG.warning(f"EnrichedWallMergeItem: no wall items. set_empty.")
if self.should_archive():
self.set_empty()
- return None
+ return
wdf = wall_coll.ddf(
items=wall_items,
@@ -85,7 +87,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
)
if wdf is None:
- return None
+ return
wdf = wdf.repartition(npartitions=1)
wdf = wdf.reset_index(drop=False)
@@ -106,7 +108,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
if len(session_items) == 0:
LOG.error(f"EnrichedWallMergeItem: no session items. breaking early.")
- return None
+ return
sdf = session_coll.ddf(
items=session_items,
@@ -230,8 +232,8 @@ class EnrichedWallMerge(MergeCollection):
self,
rr: "ReportRequest",
client: Client,
- product_ids: Optional[List[UUIDStr]] = None,
- user: Optional[User] = None,
+ product_ids: list[UUIDStr] | None = None,
+ user: User | None = None,
) -> pd.DataFrame:
"""We don't have the concept of a Team yet so product_ids will be a list"""
diff --git a/generalresearch/incite/mergers/foundations/user_id_product.py b/generalresearch/incite/mergers/foundations/user_id_product.py
index 73ee36a..863a741 100644
--- a/generalresearch/incite/mergers/foundations/user_id_product.py
+++ b/generalresearch/incite/mergers/foundations/user_id_product.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import logging
-from typing import Any, Dict, Literal, Optional
+from typing import Any, Literal
from distributed import Client
@@ -19,7 +21,7 @@ class UserIdProductMergeItem(MergeCollectionItem):
self,
client: Client,
user_coll: UserDFCollection,
- client_resources: Optional[Dict[str, Any]] = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
LOG.warning(f"UserIdProductMergeItem.build({self.interval})")
@@ -37,13 +39,13 @@ class UserIdProductMergeItem(MergeCollectionItem):
class UserIdProductMerge(MergeCollection):
merge_type: Literal[MergeType.USER_ID_PRODUCT] = MergeType.USER_ID_PRODUCT
collection_item_class: Literal[UserIdProductMergeItem] = UserIdProductMergeItem
- offset: Optional[str] = None
+ offset: str | None = None
def build(
self,
client: Client,
user_coll: UserDFCollection,
- client_resources: Optional[Dict[str, Any]] = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
LOG.info(f"UserIdProductMerge.build(user_coll={user_coll.signature()})")
diff --git a/generalresearch/incite/mergers/nginx_core.py b/generalresearch/incite/mergers/nginx_core.py
deleted file mode 100644
index 343da8f..0000000
--- a/generalresearch/incite/mergers/nginx_core.py
+++ /dev/null
@@ -1,146 +0,0 @@
-import json
-import logging
-from datetime import datetime, timedelta
-from typing import List, Literal
-from urllib.parse import parse_qs, urlsplit
-
-import dask.bag as db
-import pandas as pd
-from sentry_sdk import capture_exception
-
-from generalresearch.incite.mergers import (
- MergeCollection,
- MergeCollectionItem,
- MergeType,
-)
-from generalresearch.incite.schemas.mergers.nginx import NGINXCoreSchema
-from generalresearch.models.thl.definitions import ReservedQueryParameters
-
-LOG = logging.getLogger("incite")
-
-uuid4hex = r"[a-f0-9]{8}-?[a-f0-9]{4}-?4[a-f0-9]{3}-?[89ab][a-f0-9]{3}-?[a-f0-9]{12}"
-
-
-class NginxCoreMergeItem(MergeCollectionItem):
-
- def build(self) -> None:
- ir: pd.Interval = self.interval
- is_partial = not self.should_archive()
- coll: MergeCollection = self._collection
-
- start, end = ir.left.to_pydatetime(), ir.right.to_pydatetime()
- __name__ = coll.merge_type.value
-
- reserved_kwargs = set([e.value for e in ReservedQueryParameters]) | set(
- ["AC5AD0DDBC0C", "66482fb"]
- )
- LOG.info(f"{__name__}: {self._collection._client} {reserved_kwargs=}")
-
- # --- READ ---
- _start = start.replace(hour=0)
- _end = end.replace(hour=0)
- days: List[str] = [
- (start + timedelta(days=i)).strftime("%Y-%m-%d")
- for i in range((end - start).days + 1)
- ]
- LOG.info(f"{__name__}: READ start")
- lines = db.read_text(
- urlpath=[f"/tmp/thl-core-logs/access.log-{day}-*.gz" for day in days],
- compression="gzip",
- include_path=False,
- )
-
- # --- PROCESS ---
- LOG.info(f"{__name__}: PROCESS start")
-
- def process_core_entry(x: dict) -> dict:
- request: str = x["request"].split(" ")[1] # GET full_url_path HTTP/1.1.1
- referer: str = x["referer"]
-
- request_split = urlsplit(request)
- request_query_dict = parse_qs(request_split.query)
-
- # -- couldn't get to work well with .astype. I know the 0 or 0.0 isn't good but too frustrating for now
- try:
- upstream_status = int(x["upstream_status"])
- except (Exception,):
- upstream_status = 0
-
- try:
- status = int(x["status"])
- except (Exception,):
- status = 0
-
- try:
- request_time = float(x["request_time"])
- except (Exception,):
- request_time = 0.0
-
- try:
- upstream_response_time = float(x["upstream_response_time"])
- except (Exception,):
- upstream_response_time = 0.0
-
- return {
- "time": datetime.fromtimestamp(float(x["time"])),
- "method": x.get("method", None),
- "user_agent": x.get("user_agent", None),
- "upstream_route": x.get("upstream_route", None),
- "host": x.get("host", None),
- "upstream_status": upstream_status,
- "status": status,
- "request_time": request_time,
- "upstream_response_time": upstream_response_time,
- "upstream_cache_hit": x.get("upstream_cache_hit") == "True",
- # GRL custom
- "request_path": request_split.path,
- "referer": referer,
- "session_id": request_query_dict.get("AC5AD0DDBC0C", [None])[0],
- "request_id": request_query_dict.get("66482fb", [None])[0],
- "nudge_id": request_query_dict.get("5e0e0323", [None])[0],
- "request_custom_query_params": ",".join(
- [
- qk
- for qk in request_query_dict.keys()
- if qk not in reserved_kwargs
- ]
- ),
- }
-
- LOG.info(f"{__name__}: PROCESS - records maps")
- records = lines.map(json.loads).map(process_core_entry)
- LOG.info(f"{__name__}: PROCESS - .to_dataframe()")
- ddf = records.to_dataframe()
-
- # --- for "partition_on" ---
- ddf = ddf[
- ddf["time"].between(ir.left.to_datetime64(), ir.right.to_datetime64())
- ]
- ddf = ddf.repartition(npartitions=1)
-
- # --- SAVE ---
- LOG.info(f"{__name__}: SAVE start")
- self.to_archive(ddf=ddf, is_partial=is_partial)
- LOG.info(f"{__name__}: SAVE end")
-
- return None
-
-
-class NginxCoreMerge(MergeCollection):
- merge_type: Literal[MergeType.NGINX_CORE] = MergeType.NGINX_CORE
- _schema = NGINXCoreSchema
- collection_item_class = NginxCoreMergeItem
-
- def build(self) -> None:
- LOG.info(f"NginxCoreMerge.build()")
-
- for item in reversed(self.items):
- item: NginxCoreMergeItem
-
- try:
- item.build()
- except (Exception,) as e:
- capture_exception(error=e)
- pass
-
- return None
diff --git a/generalresearch/incite/mergers/nginx_fsb.py b/generalresearch/incite/mergers/nginx_fsb.py
deleted file mode 100644
index 9cb71d3..0000000
--- a/generalresearch/incite/mergers/nginx_fsb.py
+++ /dev/null
@@ -1,150 +0,0 @@
-import json
-import logging
-import re
-from datetime import datetime, timedelta
-from typing import List, Literal
-from urllib.parse import parse_qs, urlsplit
-
-import dask.bag as db
-import pandas as pd
-from sentry_sdk import capture_exception
-
-from generalresearch.incite.mergers import (
- MergeCollection,
- MergeCollectionItem,
- MergeType,
-)
-from generalresearch.incite.schemas.mergers.nginx import NGINXFSBSchema
-from generalresearch.models.thl.definitions import ReservedQueryParameters
-
-LOG = logging.getLogger("incite")
-
-uuid4hex = r"[a-f0-9]{8}-?[a-f0-9]{4}-?4[a-f0-9]{3}-?[89ab][a-f0-9]{3}-?[a-f0-9]{12}"
-
-
-class NginxFSBMergeItem(MergeCollectionItem):
-
- def build(self) -> None:
- ir: pd.Interval = self.interval
- is_partial = not self.should_archive()
- coll: MergeCollection = self._collection
-
- start, end = ir.left.to_pydatetime(), ir.right.to_pydatetime()
- __name__ = coll.merge_type.value
-
- reserved_kwargs = set([e.value for e in ReservedQueryParameters])
- LOG.info(f"{__name__}: {coll._client} {reserved_kwargs=}")
-
- # --- READ ---
- _start = start.replace(hour=0)
- _end = end.replace(hour=0)
- days: List[str] = [
- (start + timedelta(days=i)).strftime("%Y-%m-%d")
- for i in range((end - start).days + 1)
- ]
- LOG.info(f"{__name__}: READ start: {days=}")
- lines = db.read_text(
- urlpath=[f"/tmp/fsb-logs/access.log-{day}-*.gz" for day in days],
- compression="gzip",
- include_path=False,
- )
-
- # --- PROCESS ---
- LOG.info(f"{__name__}: PROCESS start")
-
- def process_fsb_entry(x: dict) -> dict:
- request: str = x["request"].split(" ")[1] # GET full_url_path HTTP/1.1.1
- url_split = urlsplit(request)
- query_dict = parse_qs(url_split.query)
- product_ids = re.findall(uuid4hex, request)
- product_id = (
- product_ids[0] if len(product_ids) else "-"
- ) # Cannot (categorize) convert non-finite values
- is_offerwall = "/offerwall/" in url_split.path
- offerwall = "-" # Cannot (categorize) convert non-finite values
- if is_offerwall:
- offerwall = url_split.path.split("/offerwall/")[1][:-1] or "-"
- is_report = "/report/" in url_split.path
-
- # -- couldn't get to work well with .astype. I know the 0 or 0.0 isn't good but too frustrating for now
- try:
- status = int(x["status"])
- except (Exception,):
- status = 0
-
- try:
- upstream_status = int(x["upstream_status"])
- except (Exception,):
- upstream_status = 0
-
- try:
- request_time = float(x["request_time"])
- except (Exception,):
- request_time = 0.0
-
- try:
- upstream_response_time = float(x["upstream_response_time"])
- except (Exception,):
- upstream_response_time = 0.0
-
- return {
- "time": datetime.fromtimestamp(float(x["time"])),
- "method": x.get("method", None),
- "user_agent": x.get("user_agent", None),
- "upstream_route": x.get("upstream_route", None),
- "host": x.get("host", None),
- "status": status,
- "upstream_status": upstream_status,
- "request_time": request_time,
- "upstream_response_time": upstream_response_time,
- "upstream_cache_hit": x.get("upstream_cache_hit") == "True",
- # GRL custom
- "product_id": product_id,
- "product_user_id": query_dict.get("bpuid", [None])[0],
- "n_bins": query_dict.get("n_bins", [None])[0],
- "is_offerwall": is_offerwall,
- "offerwall": offerwall,
- "is_report": is_report,
- "custom_query_params": ",".join(
- [qk for qk in query_dict.keys() if qk not in reserved_kwargs]
- ),
- }
-
- LOG.info(f"{__name__}: PROCESS - records maps")
- records = lines.map(json.loads).map(process_fsb_entry)
- LOG.info(f"{__name__}: PROCESS - .to_dataframe()")
- ddf = records.to_dataframe()
-
- # -- for "partition_on"
- LOG.info(f"{__name__}: PROCESS - cleanup")
- ddf = ddf[
- ddf["time"].between(ir.left.to_datetime64(), ir.right.to_datetime64())
- ]
- ddf = ddf.repartition(npartitions=1)
-
- # --- SAVE ---
- LOG.info(f"{__name__}: SAVE start")
- self.to_archive(ddf=ddf, is_partial=is_partial)
- LOG.info(f"{__name__}: SAVE finish")
-
- return None
-
-
-class NginxFSBMerge(MergeCollection):
- merge_type: Literal[MergeType.NGINX_FSB] = MergeType.NGINX_FSB
- _schema = NGINXFSBSchema
- collection_item_class = NginxFSBMergeItem
-
- def build(self) -> None:
- LOG.info(f"NginxFSBMerge.build()")
-
- for item in reversed(self.items):
- item: NginxFSBMergeItem
-
- try:
- item.build()
- except (Exception,) as e:
- capture_exception(error=e)
- pass
-
- return None
diff --git a/generalresearch/incite/mergers/nginx_grs.py b/generalresearch/incite/mergers/nginx_grs.py
deleted file mode 100644
index fb22070..0000000
--- a/generalresearch/incite/mergers/nginx_grs.py
+++ /dev/null
@@ -1,141 +0,0 @@
-import json
-import logging
-from datetime import datetime, timedelta
-from typing import List, Literal
-from urllib.parse import parse_qs, urlsplit
-
-import dask.bag as db
-import dask.dataframe as dd
-from sentry_sdk import capture_exception
-
-from generalresearch.incite.mergers import (
- MergeCollection,
- MergeCollectionItem,
- MergeType,
-)
-from generalresearch.incite.schemas.mergers.nginx import NGINXGRSSchema
-
-LOG = logging.getLogger("incite")
-
-
-uuid4hex = r"[a-f0-9]{8}-?[a-f0-9]{4}-?4[a-f0-9]{3}-?[89ab][a-f0-9]{3}-?[a-f0-9]{12}"
-
-
-class NginxGRSMergeItem(MergeCollectionItem):
-
- def build(self) -> None:
- ir = self.interval
- coll: MergeCollection = self._collection
- is_partial = not self.should_archive()
-
- start, end = ir.left.to_pydatetime(), ir.right.to_pydatetime()
- __name__ = self._collection.merge_type.value
-
- reserved_kwargs = set(["39057c8b", "c184efc0", "0bb50182"])
- LOG.info(f"{__name__}: {coll._client} {reserved_kwargs=}")
-
- # --- READ ---
- _start = start.replace(hour=0)
- _end = end.replace(hour=0)
- days: List[str] = [
- (start + timedelta(days=i)).strftime("%Y-%m-%d")
- for i in range((end - start).days + 1)
- ]
- LOG.info(f"{__name__}: READ start: {days}")
- lines = db.read_text(
- urlpath=[f"/tmp/grs-logs/access.log-{day}-*.gz" for day in days],
- compression="gzip",
- include_path=False,
- )
-
- # --- PROCESS ---
- LOG.info(f"{MergeType.NGINX_GRS.value}: PROCESS start")
-
- def process_grs_entry(x: dict) -> dict:
- request: str = x["request"].split(" ")[1] # GET full_url_path HTTP/1.1.1
-
- referer_split = urlsplit(x["referer"])
- referer_query_dict = parse_qs(referer_split.query)
-
- # -- couldn't get to work well with .astype. I know the 0 or 0.0 isn't good but too frustrating for now
- try:
- upstream_status = int(x["upstream_status"])
- except (Exception,):
- upstream_status = 0
-
- try:
- status = int(x["status"])
- except (Exception,):
- status = 0
-
- try:
- request_time = float(x["request_time"])
- except (Exception,):
- request_time = 0.00
-
- try:
- upstream_response_time = float(x["upstream_response_time"])
- except (Exception,):
- upstream_response_time = 0.0
- return {
- "time": datetime.fromtimestamp(float(x["time"])),
- "method": x.get("method", None),
- "user_agent": x.get("user_agent", None),
- "upstream_route": x.get("upstream_route", None),
- "host": x.get("host", None),
- "status": status,
- "upstream_status": upstream_status,
- "request_time": request_time,
- "upstream_response_time": upstream_response_time,
- "upstream_cache_hit": x.get("upstream_cache_hit") == "True",
- # GRL custom
- "product_id": referer_query_dict.get("39057c8b", [None])[0],
- "product_user_id": referer_query_dict.get("c184efc0", [None])[0],
- "wall_uuid": referer_query_dict.get("0bb50182", [None])[0],
- "custom_query_params": ",".join(
- [
- qk
- for qk in referer_query_dict.keys()
- if qk not in reserved_kwargs
- ]
- ),
- }
-
- LOG.info(f"{__name__}: PROCESS - records maps")
- records = lines.map(json.loads).map(process_grs_entry)
- LOG.info(f"{__name__}: PROCESS - .to_dataframe()")
- ddf: dd.DataFrame = records.to_dataframe()
-
- # -- for "partition_on"
- LOG.info(f"{__name__}: PROCESS - cleanup")
- ddf = ddf[
- ddf["time"].between(ir.left.to_datetime64(), ir.right.to_datetime64())
- ]
- ddf = ddf.repartition(npartitions=1)
-
- # --- SAVE ---
- LOG.info(f"{__name__}: SAVE start")
- self.to_archive(ddf=ddf, is_partial=is_partial)
- LOG.info(f"{__name__}: SAVE finish")
-
- return None
-
-
-class NginxGRSMerge(MergeCollection):
- merge_type: Literal[MergeType.NGINX_GRS] = MergeType.NGINX_GRS
- _schema = NGINXGRSSchema
- collection_item_class = NginxGRSMergeItem
-
- def build(self) -> None:
- LOG.info(f"NginxGRSMerge.build()")
-
- for item in reversed(self.items):
- item: NginxGRSMergeItem
-
- try:
- item.build()
- except (Exception,) as e:
- capture_exception(e)
- pass
-
- return None
diff --git a/generalresearch/incite/mergers/pop_ledger.py b/generalresearch/incite/mergers/pop_ledger.py
index 0475df2..b32503c 100644
--- a/generalresearch/incite/mergers/pop_ledger.py
+++ b/generalresearch/incite/mergers/pop_ledger.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import logging
-from typing import Any, Dict, Literal, Optional
+from typing import Any, Literal
import dask.dataframe as dd
import pandas as pd
@@ -23,8 +25,8 @@ class PopLedgerMergeItem(MergeCollectionItem):
def build(
self,
ledger_coll: LedgerDFCollection,
- client: Optional[Client] = None,
- client_resources: Optional[Dict[str, Any]] = None,
+ client: Client | None = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
ir: pd.Interval = self.interval
@@ -44,14 +46,14 @@ class PopLedgerMergeItem(MergeCollectionItem):
)
if ddf is None:
- return None
+ return
ddf = ddf[ddf["created"].between(start, end)]
df: pd.DataFrame = client.compute(ddf, resources=client_resources, sync=True)
if df.empty:
# self.set_empty()
- return None
+ return
df["direction_name"] = df["direction"].apply(lambda x: Direction(x).name)
diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py
index 4c2defb..c060aae 100644
--- a/generalresearch/incite/mergers/ym_survey_wall.py
+++ b/generalresearch/incite/mergers/ym_survey_wall.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import logging
from datetime import timedelta
-from typing import Any, Dict, Literal, Optional
+from typing import Any, Literal
import dask.dataframe as dd
import pandas as pd
@@ -30,8 +32,8 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
self,
wall_coll: WallDFCollection,
enriched_session: EnrichedSessionMerge,
- client: Optional[Client] = None,
- client_resources: Optional[Dict[str, Any]] = None,
+ client: Client | None = None,
+ client_resources: dict[str, Any] | None = None,
) -> None:
LOG.info(f"YMSurveyWallMerge.build({self.start=}, {self.finish=})")
ir: pd.Interval = self.interval
@@ -114,7 +116,7 @@ class YMSurveyWallMerge(MergeCollection):
collection_item_class: Literal[YMSurveyWallMergeCollectionItem] = (
YMSurveyWallMergeCollectionItem
)
- start: Optional[AwareDatetimeISO] = None
+ start: AwareDatetimeISO | None = None
offset: str = "10D"
_schema = YMSurveyWallSchema
diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py
index c7b871c..2f5995f 100644
--- a/generalresearch/incite/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/mergers/ym_wall_summary.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
from datetime import datetime, time, timedelta
-from typing import List, Literal, Optional, Type
+from typing import Literal, Type
import dask.dataframe as dd
import pandas as pd
@@ -81,16 +83,16 @@ class YMWallSummaryMerge(MergeCollection):
merge_type: Literal[MergeType.YM_WALL_SUMMARY] = MergeType.YM_WALL_SUMMARY
_schema = YMWallSummarySchema
collection_item_class: Type[YMWallSummaryMergeItem] = YMWallSummaryMergeItem
- items: List[YMWallSummaryMergeItem] = Field(default_factory=list)
+ items: list[YMWallSummaryMergeItem] = Field(default_factory=list)
@field_validator("offset")
- def check_offset_ym_wall_summary(cls, v: Optional[str]):
+ def check_offset_ym_wall_summary(cls, v: str | None):
# the offset MUST be on a whole day, no hourly
assert v.endswith("D"), "offset must be in days"
return v
@field_validator("start")
- def check_start_ym_wall_summary(cls, v: Optional[datetime]):
+ def check_start_ym_wall_summary(cls, v: datetime | None):
# the start MUST be start on midnight exactly
assert v.time() == time(0, 0, 0, 0), "start must no have a time component"
return v
diff --git a/generalresearch/incite/schemas/mergers/nginx.py b/generalresearch/incite/schemas/mergers/nginx.py
deleted file mode 100644
index 3d16738..0000000
--- a/generalresearch/incite/schemas/mergers/nginx.py
+++ /dev/null
@@ -1,140 +0,0 @@
-# MergeType.NGINX_GRS: NGINXGRSSchema,
-# MergeType.NGINX_FSB: NGINXFSBSchema,
-# MergeType.NGINX_CORE: NGINXCoreSchema,
-
-from datetime import timedelta
-
-import pandas as pd
-from pandera import Check, Column, DataFrameSchema, Index
-
-from generalresearch.incite.schemas import ARCHIVE_AFTER, PARTITION_ON
-
-NGINXBaseSchema = DataFrameSchema(
- columns={
- "time": Column(dtype=pd.DatetimeTZDtype(tz="UTC"), nullable=False),
- "method": Column(
- dtype=str, checks=[Check.str_length(max_value=8)], nullable=True
- ),
- "user_agent": Column(
- dtype=str, checks=[Check.str_length(max_value=3_000)], nullable=True
- ),
- "upstream_route": Column(
- dtype=str, checks=[Check.str_length(max_value=255)], nullable=True
- ),
- "host": Column(
- dtype=str, checks=[Check.str_length(max_value=255)], nullable=True
- ),
- "status": Column(
- dtype="Int32",
- checks=[Check.between(min_value=0, max_value=600)],
- nullable=False,
- ),
- "upstream_status": Column(
- dtype="Int32",
- checks=[Check.between(min_value=0, max_value=600)],
- nullable=False,
- ),
- "request_time": Column(
- dtype=float,
- checks=[Check.greater_than_or_equal_to(min_value=0)],
- nullable=False,
- ),
- "upstream_response_time": Column(
- dtype=float,
- checks=[Check.greater_than_or_equal_to(min_value=0)],
- nullable=False,
- ),
- "upstream_cache_hit": Column(dtype=bool, nullable=False),
- }
-)
-
-NGINXGRSSchema = DataFrameSchema(
- index=Index(dtype=int, checks=Check.greater_than_or_equal_to(0)),
- columns=NGINXBaseSchema.columns
- | {
- # --- GRL Custom
- "product_id": Column(
- dtype=str, checks=Check.str_length(min_value=1, max_value=32), nullable=True
- ),
- "product_user_id": Column(
- dtype=str,
- checks=Check.str_length(min_value=1, max_value=128),
- nullable=True,
- ),
- "wall_uuid": Column(
- dtype=str,
- # It's modified by some people and so this breaks..
- # checks=[Check.str_length(min_value=32, max_value=32)],
- nullable=True,
- ),
- "custom_query_params": Column(
- dtype=str, checks=[Check.str_length(max_value=3_000)], nullable=True
- ),
- },
- checks=[],
- coerce=True,
- metadata={PARTITION_ON: ["product_id"], ARCHIVE_AFTER: timedelta(minutes=1)},
-)
-
-NGINXCoreSchema = DataFrameSchema(
- index=Index(dtype=int, checks=Check.greater_than_or_equal_to(0)),
- columns=NGINXBaseSchema.columns
- | {
- # --- GRL Custom
- "request_path": Column(
- dtype=str,
- checks=Check.str_length(min_value=1, max_value=3_000),
- nullable=False,
- ),
- "referer": Column(
- dtype=str,
- checks=Check.str_length(min_value=1, max_value=128),
- nullable=True,
- ),
- "session_id": Column(
- dtype=str, checks=Check.str_length(max_value=3_000), nullable=True
- ),
- "request_id": Column(
- dtype=str, checks=Check.str_length(max_value=3_000), nullable=True
- ),
- "nudge_id": Column(
- dtype=str, checks=Check.str_length(max_value=3_000), nullable=True
- ),
- "request_custom_query_params": Column(
- dtype=str, checks=[Check.str_length(max_value=3_000)], nullable=True
- ),
- },
- checks=[],
- coerce=True,
- metadata={PARTITION_ON: None, ARCHIVE_AFTER: timedelta(minutes=1)},
-)
-
-NGINXFSBSchema = DataFrameSchema(
- index=Index(dtype=int, checks=Check.greater_than_or_equal_to(0)),
- columns=NGINXBaseSchema.columns
- | {
- # --- GRL Custom
- "product_id": Column(
- dtype=str, checks=Check.str_length(min_value=1, max_value=32), nullable=True
- ),
- "product_user_id": Column(
- dtype=str,
- checks=Check.str_length(min_value=1, max_value=128),
- nullable=True,
- ),
- "n_bins": Column(
- dtype="Int32",
- checks=Check.greater_than_or_equal_to(min_value=0),
- nullable=True,
- ),
- "is_offerwall": Column(dtype=bool, nullable=False),
- "offerwall": Column(dtype=bool, nullable=False),
- "is_report": Column(dtype=bool, nullable=False),
- "custom_query_params": Column(
- dtype=str, checks=[Check.str_length(max_value=3_000)], nullable=True
- ),
- },
- checks=[],
- coerce=True,
- metadata={PARTITION_ON: ["product_id"], ARCHIVE_AFTER: timedelta(minutes=1)},
-)
diff --git a/generalresearch/incite/schemas/mergers/ym_wall_summary.py b/generalresearch/incite/schemas/mergers/ym_wall_summary.py
index 5e97b47..37a517f 100644
--- a/generalresearch/incite/schemas/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/schemas/mergers/ym_wall_summary.py
@@ -1,5 +1,6 @@
+from __future__ import annotations
+
from datetime import timedelta
-from typing import Set
from pandera import Check, Column, DataFrameSchema, Index
@@ -7,7 +8,7 @@ from generalresearch.incite.schemas import ARCHIVE_AFTER
from generalresearch.locales import Localelator
from generalresearch.models import Source
-COUNTRY_ISOS: Set[str] = Localelator().get_all_countries()
+COUNTRY_ISOS: set[str] = Localelator().get_all_countries()
kosovo = "xk"
COUNTRY_ISOS.add(kosovo)
diff --git a/generalresearch/managers/base.py b/generalresearch/managers/base.py
index bb9ca75..b935413 100644
--- a/generalresearch/managers/base.py
+++ b/generalresearch/managers/base.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
+from collections.abc import Collection
from enum import Enum
-from typing import Collection, Optional
from generalresearch.pg_helper import PostgresConfig
from generalresearch.redis_helper import RedisConfig
@@ -21,7 +23,7 @@ class SqlManager(Manager):
def __init__(
self,
sql_helper: SqlHelper,
- permissions: Optional[Collection[Permission]] = None,
+ permissions: Collection[Permission] | None = None,
**kwargs,
):
super().__init__(**kwargs)
@@ -36,7 +38,7 @@ class PostgresManager(Manager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
**kwargs,
):
super().__init__(**kwargs)
@@ -50,7 +52,7 @@ class RedisManager(Manager):
def __init__(
self,
redis_config: RedisConfig,
- cache_prefix: Optional[str] = None,
+ cache_prefix: str | None = None,
**kwargs,
):
super().__init__(**kwargs)
@@ -64,8 +66,8 @@ class SqlManagerWithRedis(SqlManager, RedisManager):
self,
sql_helper: SqlHelper,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
):
super().__init__(
sql_helper=sql_helper,
@@ -80,8 +82,8 @@ class PostgresManagerWithRedis(PostgresManager, RedisManager):
self,
pg_config: PostgresConfig,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
):
super().__init__(
pg_config=pg_config,
diff --git a/generalresearch/managers/cint/profiling.py b/generalresearch/managers/cint/profiling.py
index 5e6e46a..d549e94 100644
--- a/generalresearch/managers/cint/profiling.py
+++ b/generalresearch/managers/cint/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from typing import Collection
from generalresearch.models.cint.question import CintQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[CintQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[CintQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py
index e8298fd..f80542e 100644
--- a/generalresearch/managers/cint/survey.py
+++ b/generalresearch/managers/cint/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional, Set
import pymysql
from pymysql import IntegrityError
@@ -31,13 +31,13 @@ class CintSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- exclude_fields: Optional[Set[str]] = None,
- ) -> List[CintSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ exclude_fields: set[str] | None = None,
+ ) -> list[CintSurvey]:
"""
Accepts lots of optional filters.
@@ -106,7 +106,7 @@ class CintSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[CintSurvey]) -> bool:
+ def update(self, surveys: list[CintSurvey]) -> bool:
now = datetime.now(tz=timezone.utc)
for survey in surveys:
survey.last_updated = now
@@ -117,7 +117,7 @@ class CintSurveyManager(SurveyManager):
self.sql_helper.bulk_update("cint_survey", survey_fields, survey_data)
return True
- def create_or_update(self, surveys: List[CintSurvey]):
+ def create_or_update(self, surveys: list[CintSurvey]):
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py
index 4cf3a3e..fe70732 100644
--- a/generalresearch/managers/criteria.py
+++ b/generalresearch/managers/criteria.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
from abc import ABC
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, Dict, Set
from more_itertools import chunked
@@ -30,7 +32,7 @@ class CriteriaManager(SqlManager, ABC):
"""
...
- def filter(self, hashes: Collection[str]) -> Dict[str, MarketplaceCondition]:
+ def filter(self, hashes: Collection[str]) -> dict[str, MarketplaceCondition]:
"""
Filter for criterion from the db
"""
@@ -44,7 +46,7 @@ class CriteriaManager(SqlManager, ABC):
)
return {x["hash"]: self.CONDITION_MODEL.from_mysql(x) for x in res}
- def filter_exists(self, hashes: Set[str]) -> Set[str]:
+ def filter_exists(self, hashes: set[str]) -> set[str]:
"""Returns hashes that exist in the db"""
res = self.sql_helper.execute_sql_query(
query=f"""
diff --git a/generalresearch/managers/dynata/profiling.py b/generalresearch/managers/dynata/profiling.py
index 39e6591..25afbe2 100644
--- a/generalresearch/managers/dynata/profiling.py
+++ b/generalresearch/managers/dynata/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from typing import Collection
from generalresearch.models.dynata.question import DynataQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[DynataQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[DynataQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py
index d20c1c8..372a57d 100644
--- a/generalresearch/managers/dynata/survey.py
+++ b/generalresearch/managers/dynata/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional
import pymysql
from pymysql import IntegrityError
@@ -50,12 +50,12 @@ class DynataSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- ) -> List[DynataSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ ) -> list[DynataSurvey]:
"""
Accepts lots of optional filters.
@@ -122,7 +122,7 @@ class DynataSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[DynataSurvey]) -> bool:
+ def update(self, surveys: list[DynataSurvey]) -> bool:
now = datetime.now(tz=timezone.utc)
update_fields = self.SURVEY_FIELDS + ["last_updated"]
@@ -131,7 +131,7 @@ class DynataSurveyManager(SurveyManager):
self.sql_helper.bulk_update("dynata_survey", update_fields, survey_data)
return True
- def create_or_update(self, surveys: List[DynataSurvey]):
+ def create_or_update(self, surveys: list[DynataSurvey]):
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py
index c0c0d0f..0be2bb9 100644
--- a/generalresearch/managers/events.py
+++ b/generalresearch/managers/events.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import logging
import math
import socket
@@ -5,7 +7,7 @@ import threading
import time
from datetime import datetime, timedelta, timezone
from decimal import Decimal
-from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
+from typing import TYPE_CHECKING
from redis.client import PubSub, Redis
@@ -119,7 +121,7 @@ class UserStatsManager(RedisManager):
# The key is the user's ID so this can be called multiple times
# without side effects.
if user.created is None:
- return None
+ return
now = round(time.time())
sec_24hr = round(timedelta(hours=24).total_seconds())
@@ -127,7 +129,7 @@ class UserStatsManager(RedisManager):
expires_at = minute + sec_24hr
ttl = expires_at - now
if ttl <= 0:
- return None
+ return
pipe = self.redis_client.pipeline()
name = "signups_last_24h"
@@ -137,7 +139,6 @@ class UserStatsManager(RedisManager):
pipe.hset(name, user.product_user_id, now)
pipe.hexpire(name, ttl, user.product_user_id)
pipe.execute()
- return None
def mark_user_active(self, user: User) -> None:
now = datetime.now(tz=timezone.utc).isoformat()
@@ -190,7 +191,6 @@ class UserStatsManager(RedisManager):
pipe.hset(name, key, now)
pipe.hexpire(name, timedelta(hours=1), key)
pipe.execute()
- return None
def unmark_user_inprogress(self, user: User):
# Call when a user exits a Session
@@ -205,7 +205,6 @@ class UserStatsManager(RedisManager):
name = "in_progress_users"
pipe.hdel(name, user.uuid)
pipe.execute()
- return None
def clear_global_user_stats(self) -> None:
# For testing
@@ -214,7 +213,6 @@ class UserStatsManager(RedisManager):
r.delete("active_users_last_24h")
r.delete("signups_last_24h")
r.delete("in_progress_users")
- return None
class TaskStatsManager(RedisManager):
@@ -252,7 +250,7 @@ class TaskStatsManager(RedisManager):
"TaskStatsManager:latest", res.model_dump_json(), ex=timedelta(hours=24)
)
- def get_latest_task_stats(self) -> Optional[TaskStatsSnapshot]:
+ def get_latest_task_stats(self) -> TaskStatsSnapshot | None:
res = self.redis_client.get("TaskStatsManager:latest")
if res is not None:
return TaskStatsSnapshot.model_validate_json(res)
@@ -286,7 +284,6 @@ class TaskStatsManager(RedisManager):
pipe.hexpire(name_24h_source, ttl_24hr, key)
pipe.execute()
- return None
def _set_live_task_stats(
self, source: Source, live_task_count: int, live_tasks_max_payout: Decimal
@@ -306,12 +303,12 @@ class TaskStatsManager(RedisManager):
pipe.execute()
- def get_active_sources(self) -> List[Source]:
+ def get_active_sources(self) -> list[Source]:
return [Source(x) for x in self.redis_client.hkeys("live_task_count")]
def get_task_stats_raw(
self,
- ) -> Dict[str, Union[AggregateBySource, MaxGaugeBySource]]:
+ ) -> dict[str, AggregateBySource | MaxGaugeBySource | None]:
sources = self.get_active_sources()
pipe = self.redis_client.pipeline(transaction=False)
@@ -375,8 +372,6 @@ class TaskStatsManager(RedisManager):
)
self.redis_client.delete(*keys)
- return None
-
class SessionStatsManager(RedisManager):
"""
@@ -413,7 +408,6 @@ class SessionStatsManager(RedisManager):
self.session_on_complete(session=session, user=user)
else:
self.session_on_fail(session=session, user=user)
- return None
def session_on_fail(self, session: Session, user: User):
r = self.redis_client
@@ -564,7 +558,7 @@ class SessionStatsManager(RedisManager):
self.calculate_avg_stats(res)
return res
- def calculate_avg_stats(self, res: Dict[str, Optional[float | int]]):
+ def calculate_avg_stats(self, res: dict[str, float | int | None]):
res["session_avg_payout_last_24h"] = None
res["session_avg_user_payout_last_24h"] = None
res["session_complete_avg_loi_last_24h"] = None
@@ -634,7 +628,7 @@ class EventManager(StatsManager):
def get_last_stats_key(self, product_id: UUIDStr):
return f"{self.cache_prefix}:last_stats:{product_id}"
- def get_active_subscribers(self) -> Set[UUIDStr]:
+ 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}
return product_ids
@@ -661,7 +655,8 @@ class EventManager(StatsManager):
res = self.redis_client.set(lock_key, 1, ex=120, nx=True)
if not res:
logging.debug("failed to acquire stats_worker_task lock")
- return None
+ return
+
logging.info("Acquired stats_worker_task lock")
for product_id in self.get_active_subscribers():
@@ -683,7 +678,7 @@ class EventManager(StatsManager):
self.redis_client.delete(lock_key)
- return None
+ return
def make_influx_point(self, channel: str, numsub: int):
return {
@@ -729,7 +724,7 @@ class EventManager(StatsManager):
)
)
self.publish_event(msg, product_id=user.product_id)
- return None
+ return
def handle_task_finish(self, wall: Wall, session: Session, user: User):
self.mark_user_active(user=user)
@@ -811,8 +806,8 @@ class EventSubscriber(RedisManager):
def __init__(self, *args, product_id: UUIDStr, **kwargs):
super().__init__(*args, **kwargs)
self.product_id = product_id
- self.pubsub_client: Optional[Redis] = None
- self.pubsub: Optional[PubSub] = None
+ self.pubsub_client: Redis | None = None
+ self.pubsub: PubSub | None = None
self._subscribe()
def _subscribe(self):
@@ -823,7 +818,7 @@ class EventSubscriber(RedisManager):
p.subscribe(self.get_channel_name())
self.pubsub_client = r
self.pubsub = p
- return None
+ return
def get_channel_name(self):
return f"{self.cache_prefix}:event-channel:{self.product_id}"
@@ -834,7 +829,7 @@ class EventSubscriber(RedisManager):
def get_last_stats_key(self):
return f"{self.cache_prefix}:last_stats:{self.product_id}"
- def get_last_stats_msg(self) -> Optional[StatsMessage]:
+ def get_last_stats_msg(self) -> StatsMessage | None:
raw = self.redis_client.get(self.get_last_stats_key())
if raw is not None:
return StatsMessage.model_validate_json(raw)
@@ -847,7 +842,7 @@ class EventSubscriber(RedisManager):
raw.reverse()
return [ServerToClientMessageAdapter.validate_json(x) for x in raw]
- def poll_message(self) -> Optional[ServerToClientMessage]:
+ def poll_message(self) -> ServerToClientMessage | None:
res = self.pubsub.get_message(ignore_subscribe_messages=True)
if res is None:
return None
diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py
index 7b8e526..21c8793 100644
--- a/generalresearch/managers/gr/authentication.py
+++ b/generalresearch/managers/gr/authentication.py
@@ -1,8 +1,10 @@
+from __future__ import annotations
+
import binascii
import logging
import os
from datetime import datetime, timezone
-from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
+from typing import TYPE_CHECKING, Any
from uuid import uuid4
from psycopg import sql
@@ -23,7 +25,7 @@ class GRUserManager(PostgresManagerWithRedis):
def create_dummy(
self,
- sub: Optional[str] = None,
+ sub: str | None = None,
is_superuser: bool = False,
) -> "GRUser":
sub = sub or f"{uuid4().hex}-{uuid4().hex}"
@@ -53,14 +55,12 @@ class GRUserManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO gr_user
(sub, is_superuser, date_joined)
VALUES (%(sub)s, %(is_superuser)s, %(date_joined)s)
RETURNING id
- """
- )
+ """)
c.execute(query=query, params=data)
gr_user_id: int = c.fetchone()["id"]
conn.commit()
@@ -68,7 +68,7 @@ class GRUserManager(PostgresManagerWithRedis):
instance.id = gr_user_id
return instance
- def get_by_id(self, gr_user_id: int) -> Optional["GRUser"]:
+ def get_by_id(self, gr_user_id: int) -> "GRUser" | None:
from generalresearch.models.gr.authentication import GRUser
with self.pg_config.make_connection() as conn:
@@ -96,7 +96,7 @@ class GRUserManager(PostgresManagerWithRedis):
assert isinstance(gr_user, GRUser), "GRUser not serialized correctly"
return gr_user
- def get_by_sub(self, sub: str, raises=True) -> Optional["GRUser"]:
+ def get_by_sub(self, sub: str, raises=True) -> "GRUser" | None:
from generalresearch.models.gr.authentication import GRUser
with self.pg_config.make_connection() as conn:
@@ -131,22 +131,20 @@ class GRUserManager(PostgresManagerWithRedis):
def get_by_sub_or_create(self, sub: str) -> "GRUser":
return self.get_by_sub(sub=sub, raises=False) or self.create(sub=sub)
- def get_all(self) -> List["GRUser"]:
+ def get_all(self) -> list["GRUser"]:
from generalresearch.models.gr.authentication import GRUser
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- c.execute(
- query="""
+ c.execute(query="""
SELECT u.*
FROM gr_user AS u
- """
- )
+ """)
res = c.fetchall()
return [GRUser.from_postgresql(i) for i in res]
- def get_by_team(self, team_id: PositiveInt) -> List["GRUser"]:
+ def get_by_team(self, team_id: PositiveInt) -> list["GRUser"]:
from generalresearch.models.gr.authentication import GRUser
with self.pg_config.make_connection() as conn:
@@ -172,7 +170,7 @@ class GRUserManager(PostgresManagerWithRedis):
def list_product_uuids(
self, user: "GRUser", thl_pg_config: PostgresConfig
- ) -> Optional[List[UUIDStr]]:
+ ) -> list[UUIDStr] | None:
if user.business_uuids is None:
LOG.warning("prefetch not run")
return None
@@ -193,10 +191,10 @@ class GRTokenManager(PostgresManager):
def get_by_key(
self,
api_key: str,
- jwks: Optional[Dict[str, Any]] = None,
- audience: Optional[str] = None,
- issuer: Optional[Union[AnyHttpUrl, str]] = None,
- gr_redis_config: Optional[RedisConfig] = None,
+ jwks: dict[str, Any] | None = None,
+ audience: str | None = None,
+ issuer: AnyHttpUrl | str | None = None,
+ gr_redis_config: RedisConfig | None = None,
) -> "GRToken":
"""Return the GRToken for this API Token.
@@ -244,14 +242,12 @@ class GRTokenManager(PostgresManager):
# API Key
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- query = sql.SQL(
- """
+ query = sql.SQL("""
SELECT grk.*
FROM gr_token AS grk
WHERE grk.key = %s
LIMIT 1
- """
- )
+ """)
c.execute(query=query, params=(api_key,))
res = c.fetchall()
@@ -284,35 +280,31 @@ class GRTokenManager(PostgresManager):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO gr_token (key, user_id, created)
VALUES (%(key)s, %(user_id)s, %(created)s)
- """
- ),
+ """),
params=data,
)
conn.commit()
- return None
+ return
- def get_by_user_id(self, user_id: PositiveInt) -> Optional["GRToken"]:
+ 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
from generalresearch.models.gr.authentication import GRToken
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- query = sql.SQL(
- """
+ query = sql.SQL("""
SELECT grt.*
FROM gr_token AS grt
LEFT JOIN gr_user AS u
ON u.id = grt.user_id
WHERE u.id = %s
LIMIT 1;
- """
- )
+ """)
c.execute(query=query, params=(user_id,))
diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py
index e9ec580..63ba474 100644
--- a/generalresearch/managers/gr/business.py
+++ b/generalresearch/managers/gr/business.py
@@ -1,4 +1,6 @@
-from typing import TYPE_CHECKING, List, Optional
+from __future__ import annotations
+
+from typing import TYPE_CHECKING, List
from uuid import UUID, uuid4
from psycopg import sql
@@ -27,12 +29,12 @@ class BusinessBankAccountManager(PostgresManager):
def create_dummy(
self,
business_id: PositiveInt,
- uuid: Optional[UUIDStr] = None,
- transfer_method: Optional["TransferMethod"] = None,
- account_number: Optional[str] = None,
- routing_number: Optional[str] = None,
- iban: Optional[str] = None,
- swift: Optional[str] = None,
+ 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
@@ -51,10 +53,10 @@ class BusinessBankAccountManager(PostgresManager):
business_id: PositiveInt,
uuid: UUIDStr,
transfer_method: "TransferMethod",
- account_number: Optional[str] = None,
- routing_number: Optional[str] = None,
- iban: Optional[str] = None,
- swift: Optional[str] = None,
+ account_number: str | None = None,
+ routing_number: str | None = None,
+ iban: str | None = None,
+ swift: str | None = None,
) -> "BusinessBankAccount":
from generalresearch.models.gr.business import BusinessBankAccount
@@ -75,8 +77,7 @@ class BusinessBankAccountManager(PostgresManager):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_bankaccount
(uuid, transfer_method, account_number,
routing_number, iban, swift, business_id)
@@ -84,8 +85,7 @@ class BusinessBankAccountManager(PostgresManager):
(%(uuid)s, %(transfer_method)s, %(account_number)s,
%(routing_number)s, %(iban)s, %(swift)s, %(business_id)s)
RETURNING id
- """
- ),
+ """),
params=data,
)
ba_id = c.fetchone()["id"] # type: ignore
@@ -100,13 +100,11 @@ class BusinessBankAccountManager(PostgresManager):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT ba.*
FROM common_bankaccount AS ba
WHERE ba.business_id = %s
- """
- ),
+ """),
params=(business_id,),
)
res = c.fetchall()
@@ -119,14 +117,14 @@ class BusinessAddressManager(PostgresManager):
def create_dummy(
self,
business_id: PositiveInt,
- uuid: Optional[UUIDStr] = None,
- line_1: Optional[str] = None,
- line_2: Optional[str] = None,
- city: Optional[str] = None,
- state: Optional[str] = None,
- postal_code: Optional[str] = None,
- phone_number: Optional[PhoneNumber] = None,
- country: Optional[str] = None,
+ uuid: UUIDStr | None = None,
+ line_1: str | None = None,
+ line_2: str | None = None,
+ city: str | None = None,
+ state: str | None = None,
+ postal_code: str | None = None,
+ phone_number: PhoneNumber | None = None,
+ country: str | None = None,
):
uuid = uuid or uuid4().hex
line_1 = line_1 or "abc"
@@ -153,13 +151,13 @@ class BusinessAddressManager(PostgresManager):
self,
business_id: PositiveInt,
uuid: UUIDStr,
- line_1: Optional[str] = None,
- line_2: Optional[str] = None,
- city: Optional[str] = None,
- state: Optional[str] = None,
- postal_code: Optional[str] = None,
- phone_number: Optional[PhoneNumber] = None,
- country: Optional[str] = None,
+ line_1: str | None = None,
+ line_2: str | None = None,
+ city: str | None = None,
+ state: str | None = None,
+ postal_code: str | None = None,
+ phone_number: PhoneNumber | None = None,
+ country: str | None = None,
) -> "BusinessAddress":
from generalresearch.models.gr.business import BusinessAddress
@@ -181,8 +179,7 @@ class BusinessAddressManager(PostgresManager):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_businessaddress
(uuid, line_1, line_2, city, country, state,
postal_code, phone_number, business_id)
@@ -191,8 +188,7 @@ class BusinessAddressManager(PostgresManager):
%(country)s, %(state)s, %(postal_code)s,
%(phone_number)s, %(business_id)s)
RETURNING id
- """
- ),
+ """),
params=data,
)
ba_id = c.fetchone()["id"] # type: ignore
@@ -220,10 +216,10 @@ class BusinessManager(PostgresManagerWithRedis):
def get_or_create(
self,
uuid: UUIDStr,
- name: Optional[str] = None,
- team: Optional["Team"] = None,
- kind: Optional["BusinessType"] = None,
- tax_number: Optional[str] = None,
+ name: str | None = None,
+ team: "Team" | None = None,
+ kind: "BusinessType" | None = None,
+ tax_number: str | None = None,
) -> "Business":
"""
Warning: this ** does not ** update the name, team, kind, tax_number
@@ -243,11 +239,11 @@ class BusinessManager(PostgresManagerWithRedis):
def create_dummy(
self,
- uuid: Optional[UUIDStr] = None,
- name: Optional[str] = None,
- team: Optional["Team"] = None,
- kind: Optional["BusinessType"] = None,
- tax_number: Optional[str] = None,
+ 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
@@ -262,10 +258,10 @@ class BusinessManager(PostgresManagerWithRedis):
def create(
self,
name: str,
- kind: Optional["BusinessType"] = None,
- uuid: Optional[UUIDStr] = None,
- team: Optional["Team"] = None,
- tax_number: Optional[str] = None,
+ kind: "BusinessType" | None = None,
+ uuid: UUIDStr | None = None,
+ team: "Team" | None = None,
+ tax_number: str | None = None,
) -> "Business":
"""
Behavior: does this raise on duplicate?
@@ -289,13 +285,11 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_business (uuid, kind, name, tax_number)
VALUES (%(uuid)s, %(kind)s, %(name)s, %(tax_number)s)
RETURNING id
- """
- ),
+ """),
params=data,
)
business_id = c.fetchone()["id"] # type: ignore
@@ -310,7 +304,7 @@ class BusinessManager(PostgresManagerWithRedis):
return business
- def get_all(self) -> List["Business"]:
+ def get_all(self) -> list["Business"]:
"""WARNING: This should be access by the /god/ page only, and only
used by GRUser.is_staff as it doesn't provide any authentication
on it's own. This is used because the .get_by_team_id() and
@@ -324,14 +318,10 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- c.execute(
- query=sql.SQL(
- """
+ c.execute(query=sql.SQL("""
SELECT b.id, b.uuid, b.kind, b.name, b.tax_number
FROM common_business AS b
- """
- )
- )
+ """))
res = c.fetchall()
response = []
@@ -348,21 +338,19 @@ class BusinessManager(PostgresManagerWithRedis):
def get_by_team(
self,
team_id: PositiveInt,
- ) -> List["Business"]:
+ ) -> list["Business"]:
# conn: psycopg.Connection = GR_POSTGRES_C.make_connection()
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT b.id, b.uuid, b.kind, b.name, b.tax_number
FROM common_business AS b
INNER JOIN common_team_businesses as tb
ON tb.business_id = b.id
WHERE tb.team_id = %s
- """
- ),
+ """),
params=(team_id,),
)
@@ -381,14 +369,13 @@ class BusinessManager(PostgresManagerWithRedis):
def get_by_user_id(
self,
user_id: PositiveInt,
- ) -> List["Business"]:
+ ) -> list["Business"]:
from generalresearch.models.gr.business import Business
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT b.id, b.uuid, b.kind, b.name, b.tax_number
FROM common_business AS b
INNER JOIN common_team_businesses AS tb
@@ -396,8 +383,7 @@ class BusinessManager(PostgresManagerWithRedis):
INNER JOIN common_membership AS m
ON m.team_id = tb.team_id
WHERE m.user_id = %s
- """
- ),
+ """),
params=(user_id,),
)
@@ -411,7 +397,7 @@ class BusinessManager(PostgresManagerWithRedis):
return response
- def get_ids_by_user_id(self, user_id: PositiveInt) -> List[PositiveInt]:
+ def get_ids_by_user_id(self, user_id: PositiveInt) -> list[PositiveInt]:
"""
:return: Every Business UUIDStr that this GRUser has permission to view
"""
@@ -419,8 +405,7 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT b.id
FROM common_business AS b
INNER JOIN common_team_businesses AS tb
@@ -428,8 +413,7 @@ class BusinessManager(PostgresManagerWithRedis):
INNER JOIN common_membership AS cm
ON tb.team_id = cm.team_id
WHERE cm.user_id = %s
- """
- ),
+ """),
params=(user_id,),
)
@@ -437,7 +421,7 @@ class BusinessManager(PostgresManagerWithRedis):
return [i["id"] for i in res]
- def get_uuids_by_user_id(self, user_id: PositiveInt) -> List[UUIDStr]:
+ def get_uuids_by_user_id(self, user_id: PositiveInt) -> list[UUIDStr]:
"""
:return: Every Business UUIDStr that this GRUser has permission to view
"""
@@ -445,8 +429,7 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT b.uuid
FROM common_business AS b
INNER JOIN common_team_businesses AS tb
@@ -454,8 +437,7 @@ class BusinessManager(PostgresManagerWithRedis):
INNER JOIN common_membership AS cm
ON tb.team_id = cm.team_id
WHERE cm.user_id = %s
- """
- ),
+ """),
params=(user_id,),
)
@@ -466,7 +448,7 @@ class BusinessManager(PostgresManagerWithRedis):
def get_by_uuid(
self,
business_uuid: UUIDStr,
- ) -> Optional["Business"]:
+ ) -> "Business" | None:
from generalresearch.models.gr.business import Business
assert UUID(hex=business_uuid).hex == business_uuid
@@ -474,14 +456,12 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT id, uuid, kind, name, tax_number
FROM common_business
WHERE uuid = %s
LIMIT 1;
- """
- ),
+ """),
params=(business_uuid,),
)
@@ -496,7 +476,7 @@ class BusinessManager(PostgresManagerWithRedis):
# data["contact"] = BusinessContact.model_validate(data)
return Business.model_validate(data)
- def get_by_id(self, business_id: PositiveInt) -> Optional["Business"]:
+ def get_by_id(self, business_id: PositiveInt) -> "Business" | None:
from generalresearch.models.gr.business import Business
assert isinstance(business_id, int)
@@ -504,14 +484,12 @@ class BusinessManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT id, uuid, kind, name, tax_number
FROM common_business
WHERE id = %s
LIMIT 1;
- """
- ),
+ """),
params=(business_id,),
)
diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py
index c6709d0..d3c2561 100644
--- a/generalresearch/managers/gr/team.py
+++ b/generalresearch/managers/gr/team.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
from datetime import datetime, timezone
-from typing import TYPE_CHECKING, List, Optional
+from typing import TYPE_CHECKING
from uuid import uuid4
from psycopg import sql
@@ -16,8 +18,6 @@ if TYPE_CHECKING:
from generalresearch.models.gr.authentication import GRUser
from generalresearch.models.gr.business import Business
from generalresearch.models.gr.team import (
- Membership,
- MembershipPrivilege,
Team,
)
@@ -66,15 +66,13 @@ class MembershipManager(PostgresManager):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_membership
(uuid, privilege, owner, team_id, user_id, created)
VALUES (%(uuid)s, %(privilege)s, %(owner)s, %(team_id)s,
%(user_id)s, %(created)s)
RETURNING id
- """
- ),
+ """),
params=data,
)
membership_id: int = c.fetchone()["id"] # type: ignore
@@ -85,19 +83,17 @@ class MembershipManager(PostgresManager):
def exists(
self, gr_user_id: PositiveInt, team_id: PositiveInt
- ) -> Optional[Membership]:
+ ) -> Membership | None:
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT id, uuid, privilege, owner, created,
user_id, team_id
FROM common_membership
WHERE team_id = %s AND user_id = %s
LIMIT 1
- """
- ),
+ """),
params=(team_id, gr_user_id),
)
res = c.fetchone()
@@ -107,38 +103,34 @@ class MembershipManager(PostgresManager):
return Membership.model_validate(res)
- def get_by_team_id(self, team_id: PositiveInt) -> List[Membership]:
+ def get_by_team_id(self, team_id: PositiveInt) -> list[Membership]:
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT id, uuid, privilege, owner, created,
user_id, team_id
FROM common_membership
WHERE team_id = %s
LIMIT 250
- """
- ),
+ """),
params=(team_id,),
)
res = c.fetchall()
return [Membership.model_validate(i) for i in res]
- def get_by_gr_user_id(self, gr_user_id: PositiveInt) -> List[Membership]:
+ def get_by_gr_user_id(self, gr_user_id: PositiveInt) -> list[Membership]:
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT id, uuid, privilege, owner, created,
user_id, team_id
FROM common_membership
WHERE user_id = %s
LIMIT 250
- """
- ),
+ """),
params=(gr_user_id,),
)
res = c.fetchall()
@@ -149,7 +141,7 @@ class MembershipManager(PostgresManager):
class TeamManager(PostgresManagerWithRedis):
def get_or_create(
- self, uuid: Optional[UUIDStr] = None, name: Optional[str] = None
+ self, uuid: UUIDStr | None = None, name: str | None = None
) -> "Team":
team = self.get_by_uuid(team_uuid=uuid)
@@ -159,25 +151,21 @@ class TeamManager(PostgresManagerWithRedis):
return self.create(uuid=uuid, name=name or "< Unknown >")
- def get_all(self) -> List["Team"]:
+ def get_all(self) -> list["Team"]:
from generalresearch.models.gr.team import Team
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
- c.execute(
- query=sql.SQL(
- """
+ c.execute(query=sql.SQL("""
SELECT t.id, t.uuid, t.name
FROM common_team AS t
- """
- )
- )
+ """))
res = c.fetchall()
return [Team.model_validate(i) for i in res]
def create_dummy(
- self, uuid: Optional[UUIDStr] = None, name: Optional[str] = None
+ self, uuid: UUIDStr | None = None, name: str | None = None
) -> "Team":
uuid = uuid or uuid4().hex
name = name or f"name-{uuid4().hex[:12]}"
@@ -187,7 +175,7 @@ class TeamManager(PostgresManagerWithRedis):
def create(
self,
name: str,
- uuid: Optional[UUIDStr] = None,
+ uuid: UUIDStr | None = None,
) -> "Team":
from generalresearch.models.gr.team import Team
@@ -196,13 +184,11 @@ class TeamManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_team (uuid, name)
VALUES (%s, %s)
RETURNING id
- """
- ),
+ """),
params=[team.uuid, team.name],
)
team_id = c.fetchone()["id"] # type: ignore
@@ -227,13 +213,11 @@ class TeamManager(PostgresManagerWithRedis):
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
INSERT INTO common_team_businesses
(team_id, business_id)
VALUES (%s, %s)
- """
- ),
+ """),
params=(
team.id,
business.id,
@@ -241,9 +225,7 @@ class TeamManager(PostgresManagerWithRedis):
)
conn.commit()
- return None
-
- def get_by_uuid(self, team_uuid: UUIDStr) -> Optional["Team"]:
+ def get_by_uuid(self, team_uuid: UUIDStr) -> "Team" | None:
from generalresearch.models.gr.team import Team
with self.pg_config.make_connection() as conn:
@@ -265,20 +247,18 @@ class TeamManager(PostgresManagerWithRedis):
return Team.model_validate(res)
- def get_by_id(self, team_id: PositiveInt) -> Optional["Team"]:
+ def get_by_id(self, team_id: PositiveInt) -> "Team" | None:
from generalresearch.models.gr.team import Team
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT t.id, t.uuid, t.name
FROM common_team AS t
WHERE t.id = %s
LIMIT 1;
- """
- ),
+ """),
params=(team_id,),
)
@@ -289,21 +269,19 @@ class TeamManager(PostgresManagerWithRedis):
return Team.model_validate(res)
- def get_by_user(self, gr_user: "GRUser") -> List["Team"]:
+ def get_by_user(self, gr_user: "GRUser") -> list["Team"]:
from generalresearch.models.gr.team import Team
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(
- query=sql.SQL(
- """
+ query=sql.SQL("""
SELECT team.*
FROM common_team AS team
INNER JOIN common_membership AS mem
ON mem.team_id = team.id
WHERE mem.user_id = %s
- """
- ),
+ """),
params=(gr_user.id,),
)
diff --git a/generalresearch/managers/innovate/profiling.py b/generalresearch/managers/innovate/profiling.py
index 81cb29e..bfa2685 100644
--- a/generalresearch/managers/innovate/profiling.py
+++ b/generalresearch/managers/innovate/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.innovate.question import InnovateQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_keys: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[InnovateQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_keys: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[InnovateQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py
index bd86e34..7db2f49 100644
--- a/generalresearch/managers/innovate/survey.py
+++ b/generalresearch/managers/innovate/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional, Set
import pymysql
from pymysql import IntegrityError
@@ -63,13 +63,13 @@ class InnovateSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- exclude_fields: Optional[Set[str]] = None,
- ) -> List[InnovateSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ exclude_fields: set[str] | None = None,
+ ) -> list[InnovateSurvey]:
"""
Accepts lots of optional filters.
:param country_iso: filters on country_iso field
@@ -141,7 +141,7 @@ class InnovateSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[InnovateSurvey]) -> bool:
+ def update(self, surveys: list[InnovateSurvey]) -> bool:
now = datetime.now(tz=timezone.utc)
update_fields = self.SURVEY_FIELDS + ["updated"]
@@ -155,7 +155,7 @@ class InnovateSurveyManager(SurveyManager):
return True
- def create_or_update(self, surveys: List[InnovateSurvey]) -> None:
+ def create_or_update(self, surveys: list[InnovateSurvey]) -> None:
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
@@ -181,5 +181,3 @@ class InnovateSurveyManager(SurveyManager):
else:
raise e
self.update([surveys[sn] for sn in existing_sns])
-
- return None
diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py
index 18dec92..52312e5 100644
--- a/generalresearch/managers/leaderboard/manager.py
+++ b/generalresearch/managers/leaderboard/manager.py
@@ -1,7 +1,9 @@
+from __future__ import annotations
+
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from functools import cached_property
-from typing import TYPE_CHECKING, List, Optional
+from typing import TYPE_CHECKING
import pandas as pd
from pandas import Period
@@ -28,7 +30,7 @@ class LeaderboardManager:
freq: LeaderboardFrequency,
product_id: str,
country_iso: str,
- within_time: Optional[NaiveDatetime | AwareDatetime] = None,
+ within_time: NaiveDatetime | AwareDatetime | None = None,
):
"""
:param within_time: Any local datetime falling within the desired leaderboard period.
@@ -87,8 +89,8 @@ class LeaderboardManager:
def get_leaderboard_rows(
self,
- limit: Optional[int] = None,
- ) -> List[LeaderboardRow]:
+ limit: int | None = None,
+ ) -> list[LeaderboardRow]:
limit = limit if limit else 0
res = self.redis_client.zrange(
self.key, start=0, end=limit - 1, withscores=True, desc=True
@@ -104,8 +106,8 @@ class LeaderboardManager:
]
def get_personal_leaderboard_rows(
- self, bp_user_id: str, limit: Optional[int] = 5
- ) -> List[LeaderboardRow]:
+ self, bp_user_id: str, limit: int | None = 5
+ ) -> list[LeaderboardRow]:
# We can't just grab this user's rank and nearby rows b/c redis does
# not handle ties the same way we do (in redis, each value is a
# unique rank, we use lowest rank for all ties). So we have to just
@@ -129,8 +131,8 @@ class LeaderboardManager:
def get_leaderboard(
self,
- limit: Optional[int] = None,
- bp_user_id: Optional[str] = None,
+ limit: int | None = None,
+ bp_user_id: str | None = None,
) -> Leaderboard:
if bp_user_id:
@@ -168,8 +170,6 @@ class LeaderboardManager:
self.redis_client.zincrby(self.key, amount=1, value=product_user_id)
self.redis_client.expire(self.key, time=self.expiration)
- return None
-
def hit_sum_payouts(self, product_user_id: str, user_payout: Decimal) -> None:
assert (
self.board_code == LeaderboardCode.SUM_PAYOUTS
@@ -179,8 +179,6 @@ class LeaderboardManager:
)
self.redis_client.expire(self.key, time=self.expiration)
- return None
-
def hit_largest_payout(self, product_user_id: str, user_payout: Decimal) -> None:
assert (
self.board_code == LeaderboardCode.LARGEST_PAYOUT
@@ -191,8 +189,6 @@ class LeaderboardManager:
)
self.redis_client.expire(self.key, time=self.expiration)
- return None
-
def hit(self, session: "Session") -> None:
user = session.user
match self.board_code:
diff --git a/generalresearch/managers/lucid/profiling.py b/generalresearch/managers/lucid/profiling.py
index ac2556d..fdd2d52 100644
--- a/generalresearch/managers/lucid/profiling.py
+++ b/generalresearch/managers/lucid/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from pydantic import ValidationError
@@ -10,11 +12,11 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- pks: Optional[Collection[Tuple[str | int, str, str]]] = None,
-) -> List[LucidQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ pks: Collection[tuple[str | int, str, str]] | None = None,
+) -> list[LucidQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/marketplace/user_pid.py b/generalresearch/managers/marketplace/user_pid.py
index cadaea9..15d8a19 100644
--- a/generalresearch/managers/marketplace/user_pid.py
+++ b/generalresearch/managers/marketplace/user_pid.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
from abc import ABC
-from typing import Collection, Dict, List, Optional
+from collections.abc import Collection
from uuid import UUID
from generalresearch.managers.base import SqlManager
@@ -12,14 +14,14 @@ class UserPidManager(SqlManager, ABC):
For getting user pids across marketplaces
"""
- SOURCE: Optional[Source] = None
+ SOURCE: Source | None = None
TABLE_NAME = None
def filter(
self,
- user_ids: Optional[Collection[int]] = None,
- pids: Optional[Collection[str]] = None,
- ) -> List[Dict[str, str]]:
+ user_ids: Collection[int] | None = None,
+ pids: Collection[str] | None = None,
+ ) -> list[dict[str, str]]:
"""
Filter by user_id or user_pid
"""
@@ -65,11 +67,11 @@ class UserPidMultiManager:
For looking up marketplace user_pids by user_id across multiple marketplaces
"""
- def __init__(self, sql_helper: SqlHelper, managers: List[UserPidManager]):
+ def __init__(self, sql_helper: SqlHelper, managers: list[UserPidManager]):
self.sql_helper = sql_helper
self.managers = managers
- def filter(self, user_ids: Optional[Collection[int]] = None):
+ def filter(self, user_ids: Collection[int] | None = None):
# You can only query across all marketplaces by user_id.
# If you are looking by user_pid, it is assumed
# you know which marketplace you are looking in.
@@ -77,14 +79,11 @@ class UserPidMultiManager:
assert isinstance(user_ids, (list, set)), "must pass a collection of user_ids"
params = [set(user_ids)] * len(self.managers)
- queries = [
- f"""
+ queries = [f"""
SELECT user_id, pid, '{m.SOURCE.value}' as source
FROM {m.mysql_db_table}
WHERE user_id IN %s
- """
- for m in self.managers
- ]
+ """ for m in self.managers]
query = "\nUNION ".join(queries)
res = self.sql_helper.execute_sql_query(query=query, params=params)
for x in res:
diff --git a/generalresearch/managers/morning/profiling.py b/generalresearch/managers/morning/profiling.py
index 494e406..01f99f3 100644
--- a/generalresearch/managers/morning/profiling.py
+++ b/generalresearch/managers/morning/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.morning.question import MorningQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,14 +9,14 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- source: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[MorningQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ source: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[MorningQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py
index 3478e92..2d86f0f 100644
--- a/generalresearch/managers/morning/survey.py
+++ b/generalresearch/managers/morning/survey.py
@@ -2,8 +2,8 @@ from __future__ import annotations
import json
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional
import pymysql
from pymysql import IntegrityError
@@ -67,12 +67,12 @@ class MorningSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- ) -> List[MorningBid]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ ) -> list[MorningBid]:
"""
Accepts lots of optional filters.
:param country_iso: filters on country_iso field
@@ -178,13 +178,13 @@ class MorningSurveyManager(SurveyManager):
return True
- def update(self, surveys: List[MorningBid]) -> None:
+ def update(self, surveys: list[MorningBid]) -> None:
now = datetime.now(tz=timezone.utc)
for survey in surveys:
self.update_one(survey, now=now)
- def update_one(self, bid: MorningBid, now: Optional[datetime] = None) -> bool:
+ def update_one(self, bid: MorningBid, now: datetime | None = None) -> bool:
if now is None:
now = datetime.now(tz=timezone.utc)
d = bid.to_mysql()
@@ -235,7 +235,7 @@ class MorningSurveyManager(SurveyManager):
conn.commit()
return bool(c.rowcount >= 1)
- def create_or_update(self, surveys: List[MorningBid]):
+ def create_or_update(self, surveys: list[MorningBid]):
surveys = {s.id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py
index 65c63e5..1efe875 100644
--- a/generalresearch/managers/network/label.py
+++ b/generalresearch/managers/network/label.py
@@ -1,8 +1,10 @@
-from datetime import datetime, timezone, timedelta
-from typing import Collection, Optional, List
+from __future__ import annotations
+
+from collections.abc import Collection
+from datetime import datetime, timedelta, timezone
from psycopg import sql
-from pydantic import TypeAdapter, IPvAnyNetwork
+from pydantic import IPvAnyNetwork, TypeAdapter
from generalresearch.managers.base import PostgresManager
from generalresearch.models.custom_types import (
@@ -15,8 +17,7 @@ from generalresearch.models.network.label import IPLabel, IPLabelKind, IPLabelSo
class IPLabelManager(PostgresManager):
def create(self, ip_label: IPLabel) -> IPLabel:
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO network_iplabel (
ip, labeled_at, created_at,
label_kind, source, confidence,
@@ -25,8 +26,7 @@ class IPLabelManager(PostgresManager):
%(ip)s, %(labeled_at)s, %(created_at)s,
%(label_kind)s, %(source)s, %(confidence)s,
%(provider)s, %(metadata)s
- ) RETURNING id;"""
- )
+ ) RETURNING id;""")
params = ip_label.model_dump_postgres()
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
@@ -36,14 +36,14 @@ class IPLabelManager(PostgresManager):
def make_filter_str(
self,
- ips: Optional[Collection[IPvAnyNetworkStr]] = None,
- ip_in_network: Optional[IPvAnyAddressStr] = None,
- label_kind: Optional[IPLabelKind] = None,
- source: Optional[IPLabelSource] = None,
- labeled_at: Optional[AwareDatetimeISO] = None,
- labeled_after: Optional[AwareDatetimeISO] = None,
- labeled_before: Optional[AwareDatetimeISO] = None,
- provider: Optional[str] = None,
+ 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 = {}
@@ -85,15 +85,15 @@ class IPLabelManager(PostgresManager):
def filter(
self,
- ips: Optional[Collection[IPvAnyNetworkStr]] = None,
- ip_in_network: Optional[IPvAnyAddressStr] = None,
- label_kind: Optional[IPLabelKind] = None,
- source: Optional[IPLabelSource] = None,
- labeled_at: Optional[AwareDatetimeISO] = None,
- labeled_after: Optional[AwareDatetimeISO] = None,
- labeled_before: Optional[AwareDatetimeISO] = None,
- provider: Optional[str] = None,
- ) -> List[IPLabel]:
+ 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,
diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py
index 9e4d773..19c5caf 100644
--- a/generalresearch/managers/network/mtr.py
+++ b/generalresearch/managers/network/mtr.py
@@ -1,4 +1,4 @@
-from typing import Optional
+from __future__ import annotations
from psycopg import Cursor, sql
@@ -8,12 +8,11 @@ from generalresearch.models.network.tool_run import MTRRun
class MTRRunManager(PostgresManager):
- def _create(self, run: MTRRun, c: Optional[Cursor] = None) -> None:
+ 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(
- """
+ query = sql.SQL("""
INSERT INTO network_mtr (
run_id, source_ip, facility_id,
protocol, port, parsed,
@@ -24,20 +23,17 @@ class MTRRunManager(PostgresManager):
%(protocol)s, %(port)s, %(parsed)s,
%(started_at)s, %(ip)s, %(scan_group_id)s
);
- """
- )
+ """)
params = run.model_dump_postgres()
- query_hops = sql.SQL(
- """
+ 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]
diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py
index f26fd44..a8470c8 100644
--- a/generalresearch/managers/network/nmap.py
+++ b/generalresearch/managers/network/nmap.py
@@ -1,4 +1,4 @@
-from typing import Optional
+from __future__ import annotations
from psycopg import Cursor, sql
@@ -8,13 +8,12 @@ from generalresearch.models.network.tool_run import NmapRun
class NmapRunManager(PostgresManager):
- def _create(self, run: NmapRun, c: Optional[Cursor] = None) -> None:
+ 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(
- """
+ query = sql.SQL("""
INSERT INTO network_portscan (
run_id, xml_version, host_state,
host_state_reason, latency_ms, distance,
@@ -29,12 +28,10 @@ class NmapRunManager(PostgresManager):
%(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(
- """
+ query_ports = sql.SQL("""
INSERT INTO network_portscanport (
port_scan_id, protocol, port,
state, reason, reason_ttl,
@@ -44,8 +41,7 @@ class NmapRunManager(PostgresManager):
%(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]
@@ -59,5 +55,3 @@ class NmapRunManager(PostgresManager):
c.execute(query, params)
if nmap_run.ports:
c.executemany(query_ports, params_ports)
-
- return None
diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py
index 41e4138..0b41a9a 100644
--- a/generalresearch/managers/network/rdns.py
+++ b/generalresearch/managers/network/rdns.py
@@ -1,4 +1,4 @@
-from typing import Optional
+from __future__ import annotations
from psycopg import Cursor
@@ -8,7 +8,7 @@ from generalresearch.models.network.tool_run import RDNSRun
class RDNSRunManager(PostgresManager):
- def _create(self, run: RDNSRun, c: Optional[Cursor] = None) -> None:
+ 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
"""
diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py
index 17f4935..026b3d3 100644
--- a/generalresearch/managers/network/tool_run.py
+++ b/generalresearch/managers/network/tool_run.py
@@ -1,19 +1,20 @@
-from typing import Collection, List, Dict
+from __future__ import annotations
-from psycopg import Cursor, sql
+from collections.abc import Collection
-from generalresearch.managers.base import PostgresManager, Permission
+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.managers.network.mtr import MTRRunManager
from generalresearch.models.network.rdns.result import RDNSResult
from generalresearch.models.network.tool_run import (
+ MTRRun,
NmapRun,
RDNSRun,
- MTRRun,
- ToolRun,
ToolName,
+ ToolRun,
)
from generalresearch.pg_helper import PostgresConfig
@@ -22,7 +23,7 @@ class ToolRunManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
):
super().__init__(pg_config=pg_config, permissions=permissions)
self.nmap_manager = NmapRunManager(self.pg_config)
@@ -30,8 +31,7 @@ class ToolRunManager(PostgresManager):
self.mtr_manager = MTRRunManager(self.pg_config)
def _create_tool_run(self, run: NmapRun | RDNSRun | MTRRun, c: Cursor):
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO network_toolrun (
ip, scan_group_id, tool_class,
tool_name, tool_version, started_at,
@@ -44,8 +44,7 @@ class ToolRunManager(PostgresManager):
%(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"]
@@ -62,7 +61,7 @@ class ToolRunManager(PostgresManager):
else:
raise ValueError("unrecognized run type")
- def get_latest_runs_by_tool(self, ip: str) -> Dict[ToolName, ToolRun]:
+ def get_latest_runs_by_tool(self, ip: str) -> dict[ToolName, ToolRun]:
query = """
SELECT DISTINCT ON (tool_name) *
FROM network_toolrun
diff --git a/generalresearch/managers/pollfish/profiling.py b/generalresearch/managers/pollfish/profiling.py
index 4431784..daf529b 100644
--- a/generalresearch/managers/pollfish/profiling.py
+++ b/generalresearch/managers/pollfish/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.pollfish.question import PollfishQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[PollfishQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[PollfishQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/precision/profiling.py b/generalresearch/managers/precision/profiling.py
index c4b24d8..449fd25 100644
--- a/generalresearch/managers/precision/profiling.py
+++ b/generalresearch/managers/precision/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.precision.question import PrecisionQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[PrecisionQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[PrecisionQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py
index 112c938..6fb30f2 100644
--- a/generalresearch/managers/precision/survey.py
+++ b/generalresearch/managers/precision/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional
import pymysql
from pymysql import IntegrityError
@@ -49,12 +49,12 @@ class PrecisionSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- ) -> List[PrecisionSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ ) -> list[PrecisionSurvey]:
"""
Accepts lots of optional filters.
:param country_iso: filters on country_iso field
@@ -145,7 +145,7 @@ class PrecisionSurveyManager(SurveyManager):
return True
- def update(self, surveys: List[PrecisionSurvey]) -> bool:
+ def update(self, surveys: list[PrecisionSurvey]) -> bool:
for survey in surveys:
self.update_one(survey)
return True
@@ -218,7 +218,7 @@ class PrecisionSurveyManager(SurveyManager):
return True
- def create_or_update(self, surveys: List[PrecisionSurvey]):
+ def create_or_update(self, surveys: list[PrecisionSurvey]):
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
diff --git a/generalresearch/managers/prodege/profiling.py b/generalresearch/managers/prodege/profiling.py
index bf4e3cf..54a7b57 100644
--- a/generalresearch/managers/prodege/profiling.py
+++ b/generalresearch/managers/prodege/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.prodege.question import ProdegeQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[ProdegeQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[ProdegeQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py
index 75fa34e..f555290 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 typing import Collection, List, Optional
import pymysql
@@ -43,12 +43,12 @@ class ProdegeSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- ) -> List[ProdegeSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ ) -> list[ProdegeSurvey]:
"""
Accepts lots of optional filters.
@@ -113,7 +113,7 @@ class ProdegeSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[ProdegeSurvey]) -> None:
+ def update(self, surveys: list[ProdegeSurvey]) -> None:
now = datetime.now(tz=timezone.utc)
# Do to stupidity with bid/actual loi/ir values (see ProdegeSurvey.to_mysql), we now
diff --git a/generalresearch/managers/repdata/profiling.py b/generalresearch/managers/repdata/profiling.py
index 6a63c38..4b97abd 100644
--- a/generalresearch/managers/repdata/profiling.py
+++ b/generalresearch/managers/repdata/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.repdata.question import RepDataQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[RepDataQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[RepDataQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py
index eecdc04..2e2224f 100644
--- a/generalresearch/managers/repdata/survey.py
+++ b/generalresearch/managers/repdata/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import json
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional
import pymysql
@@ -58,12 +58,12 @@ class RepDataSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- ) -> List[RepDataSurveyHashed]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ ) -> list[RepDataSurveyHashed]:
"""
Accepts lots of optional filters.
:param country_iso: filters on country_iso field
@@ -159,7 +159,7 @@ class RepDataSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[RepDataSurveyHashed]) -> bool:
+ def update(self, surveys: list[RepDataSurveyHashed]) -> bool:
now = datetime.now(tz=timezone.utc)
update_fields = self.SURVEY_FIELDS + ["last_updated"]
diff --git a/generalresearch/managers/sago/profiling.py b/generalresearch/managers/sago/profiling.py
index 6d00ec4..4f5b2f5 100644
--- a/generalresearch/managers/sago/profiling.py
+++ b/generalresearch/managers/sago/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.sago.question import SagoQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[SagoQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[SagoQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py
index 12b37e0..325639f 100644
--- a/generalresearch/managers/sago/survey.py
+++ b/generalresearch/managers/sago/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional, Set
import pymysql
from pymysql import IntegrityError
@@ -47,13 +47,13 @@ class SagoSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- exclude_fields: Optional[Set[str]] = None,
- ) -> List[SagoSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ exclude_fields: set[str] | None = None,
+ ) -> list[SagoSurvey]:
"""
Accepts lots of optional filters.
@@ -121,7 +121,7 @@ class SagoSurveyManager(SurveyManager):
)
return True
- def update(self, surveys: List[SagoSurvey]) -> bool:
+ def update(self, surveys: list[SagoSurvey]) -> bool:
now = datetime.now(tz=timezone.utc)
update_fields = self.SURVEY_FIELDS + ["updated"]
@@ -155,7 +155,7 @@ class SagoSurveyManager(SurveyManager):
raise ValueError("this should never happen")
return True
- def create_or_update(self, surveys: List[SagoSurvey]):
+ def create_or_update(self, surveys: list[SagoSurvey]) -> None:
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
@@ -182,5 +182,3 @@ class SagoSurveyManager(SurveyManager):
raise e
self.update([surveys[sn] for sn in existing_sns])
-
- return None
diff --git a/generalresearch/managers/spectrum/profiling.py b/generalresearch/managers/spectrum/profiling.py
index 8a9b9a9..8a0904a 100644
--- a/generalresearch/managers/spectrum/profiling.py
+++ b/generalresearch/managers/spectrum/profiling.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
import json
-from typing import Collection, List, Optional, Tuple
+from collections.abc import Collection
from generalresearch.models.spectrum.question import SpectrumQuestion
from generalresearch.sql_helper import SqlHelper
@@ -7,13 +9,13 @@ from generalresearch.sql_helper import SqlHelper
def get_profiling_library(
sql_helper: SqlHelper,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- question_ids: Optional[Collection[str]] = None,
- max_options: Optional[int] = None,
- is_live: Optional[bool] = None,
- pks: Optional[Collection[Tuple[str, str, str]]] = None,
-) -> List[SpectrumQuestion]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ question_ids: Collection[str] | None = None,
+ max_options: int | None = None,
+ is_live: bool | None = None,
+ pks: Collection[tuple[str, str, str]] | None = None,
+) -> list[SpectrumQuestion]:
"""
Accepts lots of optional filters.
diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py
index 190450d..3ff2db8 100644
--- a/generalresearch/managers/spectrum/survey.py
+++ b/generalresearch/managers/spectrum/survey.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, List, Optional
import pymysql
from pymysql import IntegrityError
@@ -56,13 +56,13 @@ class SpectrumSurveyManager(SurveyManager):
def get_survey_library(
self,
- country_iso: Optional[str] = None,
- language_iso: Optional[str] = None,
- survey_ids: Optional[Collection[str]] = None,
- is_live: Optional[bool] = None,
- updated_since: Optional[datetime] = None,
- fields: List[str] = None,
- ) -> List[SpectrumSurvey]:
+ country_iso: str | None = None,
+ language_iso: str | None = None,
+ survey_ids: Collection[str] | None = None,
+ is_live: bool | None = None,
+ updated_since: datetime | None = None,
+ fields: list[str] | None = None,
+ ) -> list[SpectrumSurvey]:
"""
Accepts lots of optional filters.
:param country_iso: filters on country_iso field
@@ -133,7 +133,7 @@ class SpectrumSurveyManager(SurveyManager):
return True
- def update(self, surveys: List[SpectrumSurvey]) -> bool:
+ def update(self, surveys: list[SpectrumSurvey]) -> bool:
now = datetime.now(tz=timezone.utc)
# Due to stupidity with bid/actual loi/ir values (last block nonsense),
@@ -144,7 +144,7 @@ class SpectrumSurveyManager(SurveyManager):
return True
- def update_one(self, survey: SpectrumSurvey, now=None) -> bool:
+ def update_one(self, survey: SpectrumSurvey, now: datetime | None = None) -> bool:
if now is None:
now = datetime.now(tz=timezone.utc)
@@ -188,7 +188,7 @@ class SpectrumSurveyManager(SurveyManager):
return c.rowcount == 1
- def create_or_update(self, surveys: List[SpectrumSurvey]) -> None:
+ def create_or_update(self, surveys: list[SpectrumSurvey]) -> None:
surveys = {s.survey_id: s for s in surveys}
sns = set(surveys.keys())
existing_sns = {
@@ -215,5 +215,3 @@ class SpectrumSurveyManager(SurveyManager):
raise e
self.update([surveys[sn] for sn in existing_sns])
-
- return None
diff --git a/generalresearch/managers/survey.py b/generalresearch/managers/survey.py
index 3e3f4ee..964344f 100644
--- a/generalresearch/managers/survey.py
+++ b/generalresearch/managers/survey.py
@@ -1,5 +1,6 @@
+from __future__ import annotations
+
from abc import ABC
-from typing import List
from generalresearch.managers.base import SqlManager
from generalresearch.models.thl.survey import MarketplaceTask
@@ -13,7 +14,7 @@ class SurveyManager(SqlManager, ABC):
"""
...
- def update(self, surveys: List[MarketplaceTask]) -> bool:
+ 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.
diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py
index b264945..04452cd 100644
--- a/generalresearch/managers/thl/buyer.py
+++ b/generalresearch/managers/thl/buyer.py
@@ -1,5 +1,7 @@
+from __future__ import annotations
+
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Collection, Dict, List, Optional
from generalresearch.managers.base import Permission, PostgresManager
from generalresearch.models import Source
@@ -12,12 +14,12 @@ class BuyerManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
):
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] = dict()
+ self.source_code_pk: dict[str, int] = dict()
self.populate_caches()
def populate_caches(self):
@@ -36,13 +38,13 @@ class BuyerManager(PostgresManager):
def get(self, source: Source, code: str) -> Buyer:
return self.source_code_buyer[f"{source.value}:{code}"]
- def get_if_exists(self, source: Source, code: str) -> Optional[Buyer]:
+ def get_if_exists(self, source: Source, code: str) -> Buyer | None:
try:
return self.get(source=source, code=code)
except KeyError:
return None
- def bulk_get_or_create(self, source: Source, codes: Collection[str]) -> List[Buyer]:
+ def bulk_get_or_create(self, source: Source, codes: Collection[str]) -> list[Buyer]:
now = datetime.now(tz=timezone.utc)
buyers = []
params_seq = []
@@ -111,5 +113,3 @@ class BuyerManager(PostgresManager):
else:
buyer.id = pk
conn.commit()
-
- return None
diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py
index accbb41..7365878 100644
--- a/generalresearch/managers/thl/cashout_method.py
+++ b/generalresearch/managers/thl/cashout_method.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
+from collections.abc import Collection
from copy import copy
from datetime import datetime, timezone
-from typing import Any, Collection, Dict, List, Optional
+from typing import Any
from uuid import UUID, uuid4
from pydantic import NonNegativeInt
@@ -41,9 +44,7 @@ class CashoutMethodManager(PostgresManager):
self.pg_config.execute_write(query, values)
- return None
-
- def delete_cashout_method(self, cm_id: str):
+ def delete_cashout_method(self, cm_id: str) -> None:
db_res = self.pg_config.execute_sql_query(
query="""
SELECT id::uuid, user_id
@@ -152,11 +153,11 @@ class CashoutMethodManager(PostgresManager):
@staticmethod
def make_filter_str(
- uuid: Optional[str] = None,
- user: Optional[User] = None,
- ext_id: Optional[str] = None,
- payout_types: Optional[Collection[PayoutType]] = None,
- is_live: Optional[bool] = True,
+ uuid: str | None = None,
+ user: User | None = None,
+ ext_id: str | None = None,
+ payout_types: Collection[PayoutType] | None = None,
+ is_live: bool | None = True,
):
filters = []
params = dict()
@@ -183,11 +184,11 @@ class CashoutMethodManager(PostgresManager):
def filter_count(
self,
- uuid: Optional[str] = None,
- user: Optional[User] = None,
- ext_id: Optional[str] = None,
- payout_types: Optional[Collection[PayoutType]] = None,
- is_live: Optional[bool] = True,
+ uuid: str | None = None,
+ user: User | None = None,
+ ext_id: str | None = None,
+ payout_types: Collection[PayoutType] | None = None,
+ is_live: bool | None = True,
) -> NonNegativeInt:
filter_str, params = self.make_filter_str(
uuid=uuid,
@@ -208,12 +209,12 @@ class CashoutMethodManager(PostgresManager):
def filter(
self,
- uuid: Optional[str] = None,
- user: Optional[User] = None,
- ext_id: Optional[str] = None,
- payout_types: Optional[Collection[PayoutType]] = None,
- is_live: Optional[bool] = True,
- ) -> List[CashoutMethod]:
+ uuid: str | None = None,
+ user: User | None = None,
+ ext_id: str | None = None,
+ payout_types: Collection[PayoutType] | None = None,
+ is_live: bool | None = True,
+ ) -> list[CashoutMethod]:
filter_str, params = self.make_filter_str(
uuid=uuid,
user=user,
@@ -231,7 +232,7 @@ class CashoutMethodManager(PostgresManager):
)
return [self.format_from_db(x, user=user) for x in res]
- def get_cashout_methods(self, user: User) -> List[CashoutMethod]:
+ def get_cashout_methods(self, user: User) -> list[CashoutMethod]:
"""
The provider column is PayoutType. Some are only user-scoped,
and some are global.
@@ -279,7 +280,7 @@ class CashoutMethodManager(PostgresManager):
return cms
@staticmethod
- def format_from_db(x: Dict[str, Any], user: Optional[User] = None) -> CashoutMethod:
+ def format_from_db(x: dict[str, Any], user: User | None = None) -> CashoutMethod:
x["id"] = UUID(x["id"]).hex
# The data column here is inconsistent. Pulling keys from the mysql 'data' col
diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py
index 15c309b..05ceb8f 100644
--- a/generalresearch/managers/thl/category.py
+++ b/generalresearch/managers/thl/category.py
@@ -1,4 +1,6 @@
-from typing import Collection, Dict
+from __future__ import annotations
+
+from collections.abc import Collection
from generalresearch.managers.base import Permission, PostgresManager
from generalresearch.models.custom_types import UUIDStr
@@ -13,11 +15,11 @@ class CategoryManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ 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] = dict()
+ self.category_label_map: dict[str, Category] = dict()
self.populate_caches()
def populate_caches(self):
diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py
index 696a497..517f677 100644
--- a/generalresearch/managers/thl/contest_manager.py
+++ b/generalresearch/managers/thl/contest_manager.py
@@ -1,5 +1,8 @@
+from __future__ import annotations
+
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Any, Collection, Dict, List, Literal, Optional, Tuple, cast
+from typing import Any, Literal, cast
from uuid import UUID
import redis
@@ -163,7 +166,7 @@ class ContestBaseManager(PostgresManager):
d = res[0]
return model_cls[d["contest_type"]].model_validate_mysql(d)
- def get_if_exists(self, contest_uuid: UUIDStr) -> Optional[Contest]:
+ def get_if_exists(self, contest_uuid: UUIDStr) -> Contest | None:
try:
return self.get(contest_uuid=contest_uuid)
@@ -174,15 +177,15 @@ class ContestBaseManager(PostgresManager):
@staticmethod
def make_filter_str(
- product_id: Optional[str] = None,
- status: Optional[ContestStatus] = None,
- contest_type: Optional[ContestType] = None,
- starts_at_before: Optional[datetime | bool] = None,
- name: Optional[str] = None,
- name_contains: Optional[str] = None,
- uuids: Optional[Collection[str]] = None,
- has_participants: Optional[bool] = None,
- ) -> Tuple[str, Dict[str, Any]]:
+ product_id: str | None = None,
+ status: ContestStatus | None = None,
+ contest_type: ContestType | None = None,
+ starts_at_before: datetime | bool | None = None,
+ name: str | None = None,
+ name_contains: str | None = None,
+ uuids: Collection[str] | None = None,
+ has_participants: bool | None = None,
+ ) -> tuple[str, dict[str, Any]]:
filters = []
params = dict()
@@ -225,18 +228,18 @@ class ContestBaseManager(PostgresManager):
def get_many(
self,
- product_id: Optional[str] = None,
- status: Optional[ContestStatus] = None,
- contest_type: Optional[ContestType] = None,
- starts_at_before: Optional[datetime | bool] = None,
- name: Optional[str] = None,
- name_contains: Optional[str] = None,
- uuids: Optional[Collection[str]] = None,
- has_participants: Optional[bool] = None,
- page: Optional[int] = None,
- size: Optional[int] = None,
+ product_id: str | None = None,
+ status: ContestStatus | None = None,
+ contest_type: ContestType | None = None,
+ starts_at_before: datetime | bool | None = None,
+ name: str | None = None,
+ name_contains: str | None = None,
+ uuids: Collection[str] | None = None,
+ has_participants: bool | None = None,
+ page: int | None = None,
+ size: int | None = None,
include_winners: bool = True,
- ) -> List[Contest]:
+ ) -> list[Contest]:
filter_str, params = self.make_filter_str(
product_id=product_id,
@@ -310,35 +313,35 @@ class ContestBaseManager(PostgresManager):
def get_many_by_user_eligible_raffle(
self, user: User, country_iso: str
- ) -> List[RaffleUserView]:
+ ) -> list[RaffleUserView]:
# Seems like this is a known pycharm bug. Doing it this way to be explicit.
# https://youtrack.jetbrains.com/issue/PY-42473/Type-inference-broken-for-Literal-with-Enum
cs = self.get_many_by_user_eligible(
user=user, country_iso=country_iso, contest_type=ContestType.RAFFLE
)
- return cast(List[RaffleUserView], cs)
+ return cast(list[RaffleUserView], cs)
def get_many_by_user_eligible_milestone(
self,
user: User,
country_iso: str,
- entry_trigger: Optional[ContestEntryTrigger] = None,
- ) -> List[MilestoneUserView]:
+ entry_trigger: ContestEntryTrigger | None = None,
+ ) -> list[MilestoneUserView]:
cs = self.get_many_by_user_eligible(
user=user,
country_iso=country_iso,
contest_type=ContestType.MILESTONE,
entry_trigger=entry_trigger,
)
- return cast(List[MilestoneUserView], cs)
+ return cast(list[MilestoneUserView], cs)
def get_many_by_user_eligible(
self,
user: User,
country_iso: str,
- contest_type: Optional[ContestType] = None,
- entry_trigger: Optional[ContestEntryTrigger] = None,
- ) -> List[ContestUserView]:
+ contest_type: ContestType | None = None,
+ entry_trigger: ContestEntryTrigger | None = None,
+ ) -> list[ContestUserView]:
# Get by product_id, and status OPEN. Then we have to filter in python.
# (could also add country filter into mysql)
assert user.user_id, "invalid user"
@@ -387,9 +390,9 @@ class ContestBaseManager(PostgresManager):
def get_many_by_user_entered(
self,
user: User,
- limit: Optional[PositiveInt] = 100,
+ limit: PositiveInt | None = 100,
order_by: Literal["recent_enter", "ending_soon"] = "recent_enter",
- ) -> List[ContestUserView]:
+ ) -> list[ContestUserView]:
"""
This sets the user_contest_info field as well, which calculates the
user's entry count and win percentages.
@@ -445,8 +448,8 @@ class ContestBaseManager(PostgresManager):
def get_many_by_user_won(
self,
user: User,
- limit: Optional[PositiveInt] = 100,
- ) -> List[ContestUserView]:
+ limit: PositiveInt | None = 100,
+ ) -> list[ContestUserView]:
"""
This sets the user_contest_info field as well, which calculates the
user's entry count and win percentages.
@@ -493,7 +496,7 @@ class ContestBaseManager(PostgresManager):
return res
@staticmethod
- def parse_user_from_row(d: Dict):
+ def parse_user_from_row(d: dict):
return User(
uuid=UUID(d["user_uuid"]).hex,
user_id=d["user_id"],
@@ -501,7 +504,7 @@ class ContestBaseManager(PostgresManager):
product_id=UUID(d["product_id"]).hex,
)
- def get_winnings_by_user(self, user: User) -> List[ContestWinner]:
+ def get_winnings_by_user(self, user: User) -> list[ContestWinner]:
assert user.user_id, "invalid user"
sql_res = self.pg_config.execute_sql_query(
query=f"""
@@ -530,7 +533,7 @@ class ContestBaseManager(PostgresManager):
return res
- def get_entries_by_contest_id(self, contest_id: PositiveInt) -> List[ContestEntry]:
+ def get_entries_by_contest_id(self, contest_id: PositiveInt) -> list[ContestEntry]:
res = self.pg_config.execute_sql_query(
query=f"""
@@ -598,7 +601,6 @@ class ContestBaseManager(PostgresManager):
assert c.rowcount == 1, "Contest changed during write"
conn.commit()
ledger_manager.create_tx_contest_close(contest=contest)
- return None
def cancel_contest(self, contest: Contest) -> int:
assert contest.status == ContestStatus.CANCELLED, "status must be cancelled"
@@ -895,7 +897,6 @@ class MilestoneContestManager(ContestBaseManager):
)
assert c.rowcount == 1, "Contest changed during write"
conn.commit()
- return None
def award_milestone_contest(
self,
@@ -940,7 +941,6 @@ class MilestoneContestManager(ContestBaseManager):
conn.commit()
contest.win_count += win_count
ledger_manager.create_tx_milestone_winner(contest=contest, winners=winners)
- return None
class LeaderboardContestManager(ContestBaseManager):
@@ -1013,7 +1013,7 @@ class ContestManager(
ledger_manager: ThlLedgerManager,
redis_client: Redis,
user_manager: UserManager,
- ) -> Dict[str, NonNegativeInt]:
+ ) -> dict[str, NonNegativeInt]:
# This is an administrative function that we'll run on a schedule,
# that will check for any open contests, for any BP, that should be
# closed, and then do it!
diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py
index f86573b..d88594d 100644
--- a/generalresearch/managers/thl/ipinfo.py
+++ b/generalresearch/managers/thl/ipinfo.py
@@ -1,7 +1,9 @@
+from __future__ import annotations
+
import ipaddress
+from collections.abc import Collection
from decimal import Decimal
from random import randint
-from typing import Collection, Dict, List, Optional
import faker
import pymysql
@@ -33,19 +35,19 @@ class IPGeonameManager(PostgresManager):
def create_dummy(
self,
- geoname_id: Optional[PositiveInt] = None,
- continent_code: Optional[str] = None,
- continent_name: Optional[str] = None,
- country_iso: Optional[str] = None,
- country_name: Optional[str] = None,
- subdivision_1_iso: Optional[str] = None,
- subdivision_1_name: Optional[str] = None,
- subdivision_2_iso: Optional[str] = None,
- subdivision_2_name: Optional[str] = None,
- city_name: Optional[str] = None,
- metro_code: Optional[int] = None,
- time_zone: Optional[str] = None,
- is_in_european_union: Optional[bool] = None,
+ geoname_id: PositiveInt | None = None,
+ continent_code: str | None = None,
+ continent_name: str | None = None,
+ country_iso: str | None = None,
+ country_name: str | None = None,
+ subdivision_1_iso: str | None = None,
+ subdivision_1_name: str | None = None,
+ subdivision_2_iso: str | None = None,
+ subdivision_2_name: str | None = None,
+ city_name: str | None = None,
+ metro_code: int | None = None,
+ time_zone: str | None = None,
+ is_in_european_union: bool | None = None,
) -> IPGeoname:
return self.create(
geoname_id=geoname_id or randint(1, 999_999_999),
@@ -121,16 +123,16 @@ class IPGeonameManager(PostgresManager):
geoname_id: PositiveInt,
continent_code: str,
continent_name: str,
- country_iso: Optional[str],
- country_name: Optional[str] = None,
- subdivision_1_iso: Optional[str] = None,
- subdivision_1_name: Optional[str] = None,
- subdivision_2_iso: Optional[str] = None,
- subdivision_2_name: Optional[str] = None,
- city_name: Optional[str] = None,
- metro_code: Optional[int] = None,
- time_zone: Optional[str] = None,
- is_in_european_union: Optional[bool] = None,
+ country_iso: str | None,
+ country_name: str | None = None,
+ subdivision_1_iso: str | None = None,
+ subdivision_1_name: str | None = None,
+ subdivision_2_iso: str | None = None,
+ subdivision_2_name: str | None = None,
+ city_name: str | None = None,
+ metro_code: int | None = None,
+ time_zone: str | None = None,
+ is_in_european_union: bool | None = None,
) -> IPGeoname:
instance = IPGeoname.model_validate(
@@ -181,8 +183,8 @@ class IPGeonameManager(PostgresManager):
def fetch_geoname_ids(
self,
- filter_ids: List[PositiveInt],
- ) -> List[IPGeoname]:
+ filter_ids: list[PositiveInt],
+ ) -> list[IPGeoname]:
if len(filter_ids) == 0:
return []
@@ -203,8 +205,8 @@ class IPGeonameManager(PostgresManager):
def fetch_geoname_ids_(
self,
c: Cursor,
- filter_ids: List[PositiveInt],
- ) -> List[IPGeoname]:
+ filter_ids: list[PositiveInt],
+ ) -> list[IPGeoname]:
assert len(filter_ids) <= 500, "chunk me"
@@ -230,30 +232,30 @@ class IPInformationManager(PostgresManager):
def create_dummy(
self,
- ip: Optional[IPvAnyAddressStr] = None,
- geoname_id: Optional[PositiveInt] = None,
- country_iso: Optional[str] = None,
- registered_country_iso: Optional[str] = None,
- is_anonymous: Optional[bool] = None,
- is_anonymous_vpn: Optional[bool] = None,
- is_hosting_provider: Optional[bool] = None,
- is_public_proxy: Optional[bool] = None,
- is_tor_exit_node: Optional[bool] = None,
- is_residential_proxy: Optional[bool] = None,
- autonomous_system_number: Optional[PositiveInt] = None,
- autonomous_system_organization: Optional[str] = None,
- domain: Optional[str] = None,
- isp: Optional[str] = None,
- mobile_country_code: Optional[str] = None,
- mobile_network_code: Optional[str] = None,
- network: Optional[str] = None,
- organization: Optional[str] = None,
- static_ip_score: Optional[float] = None,
- user_type: Optional[UserType] = None,
- postal_code: Optional[str] = None,
- latitude: Optional[Decimal] = None,
- longitude: Optional[Decimal] = None,
- accuracy_radius: Optional[int] = None,
+ ip: IPvAnyAddressStr | None = None,
+ geoname_id: PositiveInt | None = None,
+ country_iso: str | None = None,
+ registered_country_iso: str | None = None,
+ is_anonymous: bool | None = None,
+ is_anonymous_vpn: bool | None = None,
+ is_hosting_provider: bool | None = None,
+ is_public_proxy: bool | None = None,
+ is_tor_exit_node: bool | None = None,
+ is_residential_proxy: bool | None = None,
+ autonomous_system_number: PositiveInt | None = None,
+ autonomous_system_organization: str | None = None,
+ domain: str | None = None,
+ isp: str | None = None,
+ mobile_country_code: str | None = None,
+ mobile_network_code: str | None = None,
+ network: str | None = None,
+ organization: str | None = None,
+ static_ip_score: float | None = None,
+ user_type: UserType | None = None,
+ postal_code: str | None = None,
+ latitude: Decimal | None = None,
+ longitude: Decimal | None = None,
+ accuracy_radius: int | None = None,
) -> "IPInformation":
return self.create(
ip=ip or fake.ipv4_public(),
@@ -313,29 +315,29 @@ class IPInformationManager(PostgresManager):
def create(
self,
ip: IPvAnyAddressStr,
- geoname_id: Optional[PositiveInt] = None,
- country_iso: Optional[str] = None,
- registered_country_iso: Optional[str] = None,
- is_anonymous: Optional[bool] = None,
- is_anonymous_vpn: Optional[bool] = None,
- is_hosting_provider: Optional[bool] = None,
- is_public_proxy: Optional[bool] = None,
- is_tor_exit_node: Optional[bool] = None,
- is_residential_proxy: Optional[bool] = None,
- autonomous_system_number: Optional[PositiveInt] = None,
- autonomous_system_organization: Optional[str] = None,
- domain: Optional[str] = None,
- isp: Optional[str] = None,
- mobile_country_code: Optional[str] = None,
- mobile_network_code: Optional[str] = None,
- network: Optional[str] = None,
- organization: Optional[str] = None,
- static_ip_score: Optional[float] = None,
- user_type: Optional[UserType] = None,
- postal_code: Optional[str] = None,
- latitude: Optional[Decimal] = None,
- longitude: Optional[Decimal] = None,
- accuracy_radius: Optional[int] = None,
+ geoname_id: PositiveInt | None = None,
+ country_iso: str | None = None,
+ registered_country_iso: str | None = None,
+ is_anonymous: bool | None = None,
+ is_anonymous_vpn: bool | None = None,
+ is_hosting_provider: bool | None = None,
+ is_public_proxy: bool | None = None,
+ is_tor_exit_node: bool | None = None,
+ is_residential_proxy: bool | None = None,
+ autonomous_system_number: PositiveInt | None = None,
+ autonomous_system_organization: str | None = None,
+ domain: str | None = None,
+ isp: str | None = None,
+ mobile_country_code: str | None = None,
+ mobile_network_code: str | None = None,
+ network: str | None = None,
+ organization: str | None = None,
+ static_ip_score: float | None = None,
+ user_type: UserType | None = None,
+ postal_code: str | None = None,
+ latitude: Decimal | None = None,
+ longitude: Decimal | None = None,
+ accuracy_radius: int | None = None,
) -> "IPInformation":
instance = IPInformation.model_validate(
@@ -423,7 +425,7 @@ class IPInformationManager(PostgresManager):
"""
self.pg_config.execute_write(query, params=data)
- def get_ip_info(self, ip: IPvAnyAddressStr) -> Optional["IPInformation"]:
+ def get_ip_info(self, ip: IPvAnyAddressStr) -> "IPInformation" | None:
res = self.fetch_ip_information(filter_ips=[ip])
if len(res) != 1:
return None
@@ -432,8 +434,8 @@ class IPInformationManager(PostgresManager):
def fetch_ip_information(
self,
- filter_ips: List[IPvAnyAddressStr],
- ) -> List["IPInformation"]:
+ filter_ips: list[IPvAnyAddressStr],
+ ) -> list["IPInformation"]:
if len(filter_ips) == 0:
return []
@@ -453,8 +455,8 @@ class IPInformationManager(PostgresManager):
def fetch_ip_information_(
self,
c: Cursor,
- filter_ips: List[IPvAnyAddressStr],
- ) -> List["IPInformation"]:
+ filter_ips: list[IPvAnyAddressStr],
+ ) -> list["IPInformation"]:
"""
IPs are converted to normalized form (/64 network exploded) for DB lookup,
and are then matched back to the original queried form for return.
@@ -521,7 +523,7 @@ class IPInformationManager(PostgresManager):
class GeoIpInfoManager(PostgresManagerWithRedis):
- def get(self, ip_address: IPvAnyAddressStr) -> Optional[GeoIPInformation]:
+ def get(self, ip_address: IPvAnyAddressStr) -> GeoIPInformation | None:
res = self.get_cache(ip_address)
if res:
return res
@@ -532,7 +534,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
def get_multi(
self, ip_addresses: Collection[IPvAnyAddressStr]
- ) -> Dict[IPvAnyAddressStr, Optional[GeoIPInformation]]:
+ ) -> dict[IPvAnyAddressStr, GeoIPInformation | None]:
if not ip_addresses:
return {}
# To deploy this, we still have (for the next 28 days) users who's
@@ -547,7 +549,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
return res
def set_cache_multi(
- self, ipinfo_map: Dict[IPvAnyAddressStr, GeoIPInformation]
+ self, ipinfo_map: dict[IPvAnyAddressStr, GeoIPInformation]
) -> None:
"""Set multiple GeoIPInformation objects in Redis in one call."""
if not ipinfo_map:
@@ -576,7 +578,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
def get_cache_multi(
self, ip_addresses: Collection[IPvAnyAddressStr]
- ) -> Dict[IPvAnyAddressStr, Optional[GeoIPInformation]]:
+ ) -> dict[IPvAnyAddressStr, GeoIPInformation | None]:
"""Get multiple GeoIPInformation objects from Redis in one call.
Returns a dict mapping IP address -> GeoIPInformation (or None if not in cache).
@@ -627,7 +629,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
self.get_cache_key(ip_address=ipinfo.ip), data, ex=3 * 24 * 3600
)
- def get_cache(self, ip_address: IPvAnyAddressStr) -> Optional[GeoIPInformation]:
+ def get_cache(self, ip_address: IPvAnyAddressStr) -> GeoIPInformation | None:
normalized_ip, lookup_prefix = normalize_ip(ip_address)
res: str = self.get_cache_raw(normalized_ip)
if not res:
@@ -716,7 +718,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
def get_mysql_multi(
self,
ips: Collection[IPvAnyAddressStr],
- ) -> Dict[IPvAnyAddressStr, Optional[GeoIPInformation]]:
+ ) -> dict[IPvAnyAddressStr, GeoIPInformation | None]:
if len(ips) == 0:
return {}
@@ -736,8 +738,8 @@ class GeoIpInfoManager(PostgresManagerWithRedis):
def get_mysql_multi_chunk(
self,
c: Cursor,
- ips: List[IPvAnyAddressStr],
- ) -> Dict[IPvAnyAddressStr, Optional[GeoIPInformation]]:
+ ips: list[IPvAnyAddressStr],
+ ) -> dict[IPvAnyAddressStr, GeoIPInformation | None]:
assert len(ips) <= 500, "chunk me"
diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py
index 399d28f..fc56550 100644
--- a/generalresearch/managers/thl/ledger_manager/conditions.py
+++ b/generalresearch/managers/thl/ledger_manager/conditions.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import logging
from datetime import datetime, timedelta, timezone
-from typing import TYPE_CHECKING, Callable, Optional, Tuple
+from typing import TYPE_CHECKING, Callable
from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF
from generalresearch.currency import USDCent
@@ -70,12 +72,12 @@ def generate_condition_bp_payout(
payoutevent_uuid: UUIDStr,
skip_one_per_day_check: bool = False,
skip_wallet_balance_check: bool = False,
-) -> Callable[..., Tuple[bool, str]]:
+) -> Callable[..., tuple[bool, str]]:
created = datetime.now(tz=timezone.utc)
def _condition(
lm: "ThlLedgerManager",
- ) -> Tuple[bool, str]:
+ ) -> tuple[bool, str]:
bp_wallet_account = lm.get_account_or_create_bp_wallet(product=product)
tag = f"{lm.currency.value}:bp_payout:{payoutevent_uuid}"
txs_ids = lm.get_tx_ids_by_tag(tag=tag)
@@ -109,7 +111,7 @@ def generate_condition_bp_payout(
def generate_condition_user_payout_request(
- user: User, payoutevent_uuid: UUIDStr, min_balance: Optional[int] = None
+ user: User, payoutevent_uuid: UUIDStr, min_balance: int | None = None
) -> Callable[..., bool]:
"""This returns a function that checks if `user` has at least
`min_balance` in their wallet and that a payout request hasn't already
@@ -150,14 +152,14 @@ def generate_condition_user_payout_request(
def generate_condition_enter_contest(
user: User, tag: str, min_balance: USDCent
-) -> Callable[..., Tuple[bool, str]]:
+) -> Callable[..., tuple[bool, str]]:
"""This returns a function that checks if `user` has at least
`min_balance` in their wallet and that a tx doesn't already exist
with this tag
"""
assert isinstance(min_balance, USDCent), "balance must be USDCent"
- def _condition(lm: "ThlLedgerManager") -> Tuple[bool, str]:
+ def _condition(lm: "ThlLedgerManager") -> tuple[bool, str]:
txs_ids = lm.get_tx_ids_by_tag(tag)
if len(txs_ids) != 0:
logger.info(f"{tag} failed condition check duplicate transaction")
diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py
index d463b55..864f1dd 100644
--- a/generalresearch/managers/thl/ledger_manager/ledger.py
+++ b/generalresearch/managers/thl/ledger_manager/ledger.py
@@ -1,7 +1,10 @@
+from __future__ import annotations
+
import logging
from collections import defaultdict
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
-from typing import Any, Callable, Collection, Dict, List, Optional, Set, Tuple, Union
+from typing import Any, Callable
from uuid import UUID
import redis
@@ -69,9 +72,9 @@ class LedgerManagerBasePostgres(PostgresManager, RedisManager):
self,
pg_config: PostgresConfig,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
- currency: Optional[LedgerCurrency] = LedgerCurrency.USD,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
+ currency: LedgerCurrency | None = LedgerCurrency.USD,
testing: bool = False,
):
if permissions is not None and (
@@ -92,11 +95,11 @@ class LedgerManagerBasePostgres(PostgresManager, RedisManager):
def make_filter_str(
self,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
- account_uuid: Optional[str] = None,
- metadata_key: Optional[str] = None,
- metadata_value: Optional[str] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
+ account_uuid: str | None = None,
+ metadata_key: str | None = None,
+ metadata_value: str | None = None,
):
filters = []
params = {}
@@ -129,11 +132,11 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def create_tx(
self,
- entries: List[LedgerEntry],
- metadata: Optional[Dict[str, str]] = None,
- ext_description: Optional[str] = None,
- tag: Optional[str] = None,
- created: Optional[AwareDatetime] = None,
+ entries: list[LedgerEntry],
+ metadata: dict[str, str] | None = None,
+ ext_description: str | None = None,
+ tag: str | None = None,
+ created: AwareDatetime | None = None,
) -> LedgerTransaction:
"""
:returns a LedgerTransaction ID. This is because we can't fully populate
@@ -206,9 +209,9 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def create_tx_protected(
self,
lock_key: str,
- condition: Callable[..., Union[bool, Tuple[bool, str]]],
+ condition: Callable[..., bool | tuple[bool, str]],
create_tx_func: Callable,
- flag_key: Optional[str] = None,
+ flag_key: str | None = None,
skip_flag_check: bool = False,
) -> LedgerTransaction:
"""
@@ -351,11 +354,11 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
raise ValueError(f"Too many txs with this tag: {tag}")
return {x["id"] for x in res}
- def get_tx_by_tag(self, tag: str) -> List[LedgerTransaction]:
+ def get_tx_by_tag(self, tag: str) -> list[LedgerTransaction]:
tx_ids = self.get_tx_ids_by_tag(tag=tag)
return self.get_tx_by_ids(transaction_ids=tx_ids)
- def get_tx_ids_by_tags(self, tags: List[str]) -> set[PositiveInt]:
+ def get_tx_ids_by_tags(self, tags: list[str]) -> set[PositiveInt]:
res = self.pg_config.execute_sql_query(
query=f"""
SELECT lt.id, lt.tag, lt.created, lt.ext_description
@@ -367,7 +370,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
return {x["id"] for x in res}
- def get_txs_by_tags(self, tags: List[str]) -> List[LedgerTransaction]:
+ def get_txs_by_tags(self, tags: list[str]) -> list[LedgerTransaction]:
tx_ids = self.get_tx_ids_by_tags(tags=tags)
return self.get_tx_by_ids(transaction_ids=tx_ids)
@@ -384,7 +387,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def get_tx_by_ids(
self,
transaction_ids: Collection[PositiveInt],
- ) -> List[LedgerTransaction]:
+ ) -> list[LedgerTransaction]:
args = {"transaction_ids": list(transaction_ids)}
@@ -407,8 +410,8 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
@staticmethod
def process_get_tx_mysql_rows_json(
- rows: Collection[Dict[str, Any]],
- ) -> List[LedgerTransaction]:
+ rows: Collection[dict[str, Any]],
+ ) -> list[LedgerTransaction]:
"""Columns: transaction_id, created, ext_description, tag,
key_value_pairs, entries_json
- key_value_pairs: &-delimited key=value pairs
@@ -456,8 +459,8 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def get_tx_filtered_by_account_summary(
self,
account_uuid: UUIDStr,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
) -> UserLedgerTransactionTypesSummary:
filter_str, params = self.make_filter_str(
time_start=time_start,
@@ -494,8 +497,8 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def get_tx_filtered_by_account_count(
self,
account_uuid: UUIDStr,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
) -> NonNegativeInt:
filter_str, params = self.make_filter_str(
time_start=time_start,
@@ -519,10 +522,10 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def get_tx_filtered_by_account(
self,
account_uuid: UUIDStr,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
- order_by: Optional[str] = "created,tag",
- ) -> List[LedgerTransaction]:
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
+ order_by: str | None = "created,tag",
+ ) -> list[LedgerTransaction]:
txs, _ = self.get_tx_filtered_by_account_paginated(
account_uuid=account_uuid,
time_start=time_start,
@@ -535,7 +538,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
self,
account_uuid: str,
oldest_created: datetime,
- exclude_txs_before: Optional[datetime] = None,
+ exclude_txs_before: datetime | None = None,
) -> NonNegativeInt:
"""
In a paginated list of txs, if I want to calculate
@@ -568,9 +571,9 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def include_running_balance(
self,
- txs: List[UserLedgerTransactionType],
+ txs: list[UserLedgerTransactionType],
account_uuid: str,
- exclude_txs_before: Optional[AwareDatetime] = None,
+ exclude_txs_before: AwareDatetime | None = None,
):
"""
exclude_txs_before is NOT for filtering. It is a "hack" to exclude
@@ -598,12 +601,12 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
def get_tx_filtered_by_account_paginated(
self,
account_uuid: UUIDStr,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
- page: Optional[int] = None,
- size: Optional[int] = None,
- order_by: Optional[str] = "created,tag",
- ) -> Tuple[List[LedgerTransaction], int]:
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
+ page: int | None = None,
+ size: int | None = None,
+ order_by: str | None = "created,tag",
+ ) -> tuple[list[LedgerTransaction], int]:
"""
If time_start and/or time_end are passed, the txs are filtered to
include only that range.
@@ -692,9 +695,9 @@ class LedgerTransactionManager(LedgerManagerBasePostgres):
self,
metadata_key: str,
metadata_value: str,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
- ) -> List[LedgerTransaction]:
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
+ ) -> list[LedgerTransaction]:
# Renamed from "get_tx_filtered" which is not a good name
filter_str, params = self.make_filter_str(
@@ -738,8 +741,8 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
"""
def get_tx_metadata_by_txs(
- self, transactions: List[LedgerTransaction]
- ) -> Dict[PositiveInt, Dict[str, Any]]:
+ self, transactions: list[LedgerTransaction]
+ ) -> dict[PositiveInt, dict[str, Any]]:
"""
Each transaction can have 1 metadata dictionary. However, each
metadata dictionary can have multiple key/value pairs that
@@ -767,12 +770,12 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
def get_tx_metadata_ids_by_tx(
self, transaction: LedgerTransaction
- ) -> Set[PositiveInt]:
+ ) -> set[PositiveInt]:
return self.get_tx_metadata_ids_by_txs(transactions=[transaction])
def get_tx_metadata_ids_by_txs(
- self, transactions: List[LedgerTransaction]
- ) -> Set[PositiveInt]:
+ self, transactions: list[LedgerTransaction]
+ ) -> set[PositiveInt]:
"""
This explicitly returns the tx_metadata database ids. Potentially,
useful for counting total key/value pairs, and/or deleting records
@@ -794,12 +797,12 @@ class LedgerMetadataManager(LedgerManagerBasePostgres):
class LedgerEntryManager(LedgerManagerBasePostgres):
- def get_tx_entries_by_tx(self, transaction: LedgerTransaction) -> List[LedgerEntry]:
+ def get_tx_entries_by_tx(self, transaction: LedgerTransaction) -> list[LedgerEntry]:
return self.get_tx_entries_by_txs(transactions=[transaction])
def get_tx_entries_by_txs(
- self, transactions: List[LedgerTransaction]
- ) -> List[LedgerEntry]:
+ self, transactions: list[LedgerTransaction]
+ ) -> list[LedgerEntry]:
tx_ids = set([tx.id for tx in transactions])
tx_entries = self.pg_config.execute_sql_query(
query="""
@@ -852,15 +855,15 @@ class LedgerAccountManager(LedgerManagerBasePostgres):
def get_account(
self, qualified_name: str, raise_on_error: bool = True
- ) -> Optional[LedgerAccount]:
+ ) -> LedgerAccount | None:
res = self.get_account_many(
qualified_names=[qualified_name], raise_on_error=raise_on_error
)
return res[0] if len(res) == 1 else None
def get_account_many_(
- self, qualified_names: List[str], raise_on_error: bool = True
- ) -> List[Dict[str, Any]]:
+ self, qualified_names: list[str], raise_on_error: bool = True
+ ) -> list[dict[str, Any]]:
assert len(qualified_names) <= 500, "chunk me"
# qualified_name has a unique index so there can only be 0 or 1 match.
@@ -882,8 +885,8 @@ class LedgerAccountManager(LedgerManagerBasePostgres):
return list(res)
def get_account_many(
- self, qualified_names: List[str], raise_on_error: bool = True
- ) -> List[LedgerAccount]:
+ self, qualified_names: list[str], raise_on_error: bool = True
+ ) -> list[LedgerAccount]:
res = flatten(
[
self.get_account_many_(chunk, raise_on_error)
@@ -893,22 +896,22 @@ class LedgerAccountManager(LedgerManagerBasePostgres):
return [LedgerAccount.model_validate(i) for i in res]
def get_account_or_create(self, account: LedgerAccount) -> LedgerAccount:
- res: Optional[LedgerAccount] = self.get_account(
+ res: LedgerAccount | None = self.get_account(
qualified_name=account.qualified_name, raise_on_error=False
)
return res or self.create_account(account=account)
- def get_accounts(self, qualified_names: List[str]) -> List[LedgerAccount]:
+ def get_accounts(self, qualified_names: list[str]) -> list[LedgerAccount]:
return self.get_account_many(qualified_names, raise_on_error=True)
- def get_accounts_if_exists(self, qualified_names: List[str]) -> List[LedgerAccount]:
+ def get_accounts_if_exists(self, qualified_names: list[str]) -> list[LedgerAccount]:
"""Rather than returning None, this may return an empty list, or
a list that has less LedgerAccount instances than the number of
qualified_names that was passed in.
"""
return self.get_account_many(qualified_names, raise_on_error=False)
- def get_account_if_exists(self, qualified_name: str) -> Optional[LedgerAccount]:
+ def get_account_if_exists(self, qualified_name: str) -> LedgerAccount | None:
return self.get_account(qualified_name, raise_on_error=False)
def get_account_balance(self, account: LedgerAccount) -> int:
@@ -939,8 +942,8 @@ class LedgerAccountManager(LedgerManagerBasePostgres):
def get_account_balance_timerange(
self,
account: LedgerAccount,
- time_start: Optional[AwareDatetime] = None,
- time_end: Optional[AwareDatetime] = None,
+ time_start: AwareDatetime | None = None,
+ time_end: AwareDatetime | None = None,
) -> int:
"""
This returns an int and not a USDCent because an Account's balance
@@ -974,8 +977,8 @@ class LedgerAccountManager(LedgerManagerBasePostgres):
account: LedgerAccount,
metadata_key: str,
metadata_value: str,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
) -> int:
"""I want the balance for this account filtered by transactions with
a certain tag.
@@ -1042,8 +1045,7 @@ class LedgerManager(
"""This is for testing only, as it'll take forever to run this if
the ledger_manager is huge
"""
- res = self.pg_config.execute_sql_query(
- f"""
+ res = self.pg_config.execute_sql_query(f"""
SELECT
SUM(CASE WHEN normal_balance = -1 THEN total ELSE 0 END) AS credit_total,
SUM(CASE WHEN normal_balance = 1 THEN total ELSE 0 END) AS debit_total
@@ -1056,17 +1058,16 @@ class LedgerManager(
ON ledger_entry.account_id = tl.uuid
GROUP BY account_id, normal_balance
) x
- """
- )[0]
+ """)[0]
return res["credit_total"] == res["debit_total"]
def get_account_debit_credit_by_metadata(
self,
account: LedgerAccount,
metadata_key: str,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
- ) -> Dict[str, Dict[str, int]]:
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
+ ) -> dict[str, dict[str, int]]:
"""Show me the sum of debit and credit scoped to this account, grouped
by all values of metadata_key
"""
@@ -1105,9 +1106,9 @@ class LedgerManager(
def get_balances_timerange(
self,
- time_start: Optional[AwareDatetime] = None,
- time_end: Optional[AwareDatetime] = None,
- ) -> Dict[str, Any]:
+ time_start: AwareDatetime | None = None,
+ time_end: AwareDatetime | None = None,
+ ) -> dict[str, Any]:
filter_str, params = self.make_filter_str(
time_end=time_end,
diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py
index 8004320..dcdfcf3 100644
--- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py
+++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py
@@ -1,7 +1,10 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
from decimal import Decimal
-from typing import TYPE_CHECKING, Callable, Collection, List, Optional
+from typing import TYPE_CHECKING, Callable
from uuid import UUID
import numpy as np
@@ -237,8 +240,8 @@ class ThlLedgerManager(LedgerManager):
def get_tx_bp_payouts(
self,
account_uuids: Collection[UUIDStr],
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
):
if time_start is None:
time_start = datetime(year=2017, month=1, day=1, tzinfo=timezone.utc)
@@ -270,7 +273,7 @@ class ThlLedgerManager(LedgerManager):
self,
wall: Wall,
user: User,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
force: bool = False,
) -> PositiveInt:
"""
@@ -297,7 +300,7 @@ class ThlLedgerManager(LedgerManager):
)
def create_tx_task_complete_(
- self, wall: Wall, user: User, created: Optional[datetime] = None
+ self, wall: Wall, user: User, created: datetime | None = None
) -> LedgerTransaction:
revenue_account = self.get_account_task_complete_revenue()
@@ -338,7 +341,7 @@ class ThlLedgerManager(LedgerManager):
return t
def create_tx_bp_payment(
- self, session: Session, created: Optional[datetime] = None, force: bool = False
+ self, session: Session, created: datetime | None = None, force: bool = False
) -> LedgerTransaction:
"""
Create a transaction when we decide to report a session as complete
@@ -366,7 +369,7 @@ class ThlLedgerManager(LedgerManager):
)
def create_tx_bp_payment_(
- self, session: Session, created: Optional[datetime] = None
+ self, session: Session, created: datetime | None = None
) -> LedgerTransaction:
user = session.user
assert user.product, "user.prefetch_product()"
@@ -469,8 +472,8 @@ class ThlLedgerManager(LedgerManager):
return t
def create_tx_task_adjustment(
- self, wall: Wall, user: User, created: Optional[datetime] = None
- ) -> Optional[LedgerTransaction]:
+ self, wall: Wall, user: User, created: datetime | None = None
+ ) -> LedgerTransaction | None:
"""
How is this different then create_tx_bp_adjustment
@@ -554,8 +557,8 @@ class ThlLedgerManager(LedgerManager):
return t
def create_tx_bp_adjustment(
- self, session: Session, created: Optional[datetime] = None
- ) -> Optional[LedgerTransaction]:
+ self, session: Session, created: datetime | None = None
+ ) -> LedgerTransaction | None:
"""
How is this different then create_tx_task_adjustment
"""
@@ -600,7 +603,7 @@ class ThlLedgerManager(LedgerManager):
if user.product.user_wallet_enabled:
# If the user wallet is enabled, the user_payout "comes out" of
# the payout
- payout_after_adj: Optional[Decimal] = (
+ payout_after_adj: Decimal | None = (
session.get_user_payout_after_adjustment()
)
if payout_after_adj is None:
@@ -869,7 +872,7 @@ class ThlLedgerManager(LedgerManager):
amount: USDCent,
created: AwareDatetime,
direction: Direction = Direction.DEBIT,
- description: Optional[str] = None,
+ description: str | None = None,
skip_flag_check: bool = False,
) -> LedgerTransaction:
"""https://en.wikipedia.org/wiki/Plug_(accounting)
@@ -935,7 +938,7 @@ class ThlLedgerManager(LedgerManager):
amount: USDCent,
created: AwareDatetime,
direction: Direction,
- description: Optional[str] = None,
+ description: str | None = None,
) -> LedgerTransaction:
assert isinstance(amount, int)
@@ -995,9 +998,9 @@ class ThlLedgerManager(LedgerManager):
self,
user: User,
payout_event: UserPayoutEvent,
- created: Optional[datetime] = None,
- skip_flag_check: Optional[bool] = False,
- skip_wallet_balance_check: Optional[bool] = False,
+ created: datetime | None = None,
+ skip_flag_check: bool | None = False,
+ skip_wallet_balance_check: bool | None = False,
) -> LedgerTransaction:
"""
The funds move from the user's wallet into the BP's "pending"
@@ -1045,7 +1048,7 @@ class ThlLedgerManager(LedgerManager):
created=created,
)
- min_balance: Optional[int] = int(amount)
+ min_balance: int | None = int(amount)
if payout_event.payout_type == PayoutType.AMT_HIT:
# We allow the user's balance to reach up to -$1.00.
min_balance = -100 + amount
@@ -1074,9 +1077,9 @@ class ThlLedgerManager(LedgerManager):
self,
user: User,
payout_event: UserPayoutEvent,
- created: Optional[datetime] = None,
- fee_amount: Optional[Decimal] = None,
- skip_flag_check: Optional[bool] = False,
+ created: datetime | None = None,
+ fee_amount: Decimal | None = None,
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
"""
Once the cashout request is approved and completed, the funds
@@ -1172,8 +1175,8 @@ class ThlLedgerManager(LedgerManager):
self,
user: User,
payout_event: UserPayoutEvent,
- created: Optional[datetime] = None,
- skip_flag_check: Optional[bool] = False,
+ created: datetime | None = None,
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
assert (
user.product.user_wallet_enabled
@@ -1214,7 +1217,7 @@ class ThlLedgerManager(LedgerManager):
user: User,
payout_event: UserPayoutEvent,
description: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
# This is the same for all user payout requests, regardless of the
# payout_type (paypal, amt, tango)
@@ -1263,7 +1266,7 @@ class ThlLedgerManager(LedgerManager):
fee_payer_account: LedgerAccount,
fee_amount: Decimal,
description: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
"""
Creates the LedgerTransaction for a completed user payout request.
@@ -1346,7 +1349,7 @@ class ThlLedgerManager(LedgerManager):
user: User,
payout_event: UserPayoutEvent,
description: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
assert user.product
@@ -1391,9 +1394,9 @@ class ThlLedgerManager(LedgerManager):
amount: Decimal,
ref_uuid: UUIDStr,
description: str,
- source_account: Optional[LedgerAccount] = None,
- created: Optional[datetime] = None,
- skip_flag_check: Optional[bool] = False,
+ source_account: LedgerAccount | None = None,
+ created: datetime | None = None,
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
"""
Pay a user into their wallet balance. There is no fee here. There
@@ -1433,8 +1436,8 @@ class ThlLedgerManager(LedgerManager):
amount: Decimal,
ref_uuid: UUIDStr,
description: str,
- source_account: Optional[LedgerAccount] = None,
- created: Optional[datetime] = None,
+ source_account: LedgerAccount | None = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
metadata = {
@@ -1477,7 +1480,7 @@ class ThlLedgerManager(LedgerManager):
self,
contest_uuid: UUIDStr,
contest_entry: ContestEntry,
- skip_flag_check: Optional[bool] = False,
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
"""
User is requesting to enter a Raffle Contest. We'll DEBIT
@@ -1525,7 +1528,7 @@ class ThlLedgerManager(LedgerManager):
amount: USDCent,
contest_uuid: UUIDStr,
tag: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
description = f"Enter contest {amount.to_usd_str()} {contest_uuid}"
metadata = {
@@ -1562,7 +1565,7 @@ class ThlLedgerManager(LedgerManager):
def create_tx_contest_close(
self,
contest: Contest,
- skip_flag_check: Optional[bool] = False,
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
"""
Contest is over. For each winner, we make a transaction.
@@ -1695,10 +1698,10 @@ class ThlLedgerManager(LedgerManager):
def create_tx_contest_close_(
self,
- entries: List[LedgerEntry],
+ entries: list[LedgerEntry],
contest_uuid: UUIDStr,
tag: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
description = f"Close contest {contest_uuid}"
metadata = {
@@ -1716,8 +1719,8 @@ class ThlLedgerManager(LedgerManager):
def create_tx_milestone_winner(
self,
contest: MilestoneContest,
- winners: List["ContestWinner"],
- skip_flag_check: Optional[bool] = False,
+ winners: list["ContestWinner"],
+ skip_flag_check: bool | None = False,
) -> LedgerTransaction:
"""
A user has reached a milestone. Pay out any cash or physical prizes,
@@ -1796,11 +1799,11 @@ class ThlLedgerManager(LedgerManager):
def create_tx_milestone_winner_(
self,
- entries: List[LedgerEntry],
+ entries: list[LedgerEntry],
contest_uuid: UUIDStr,
user_uuid: UUIDStr,
tag: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> LedgerTransaction:
description = f"Milestone award {contest_uuid}"
metadata = {
@@ -1817,7 +1820,7 @@ class ThlLedgerManager(LedgerManager):
)
def get_user_wallet_balance(
- self, user: User, since_days_ago: Optional[int] = None
+ self, user: User, since_days_ago: int | None = None
) -> int:
"""
Calculates all payments to user's wallet minus all payouts from
@@ -1932,11 +1935,11 @@ class ThlLedgerManager(LedgerManager):
def get_user_txs(
self,
user: User,
- time_start: Optional[datetime] = None,
- time_end: Optional[datetime] = None,
+ time_start: datetime | None = None,
+ time_end: datetime | None = None,
page: int = 1,
size: int = 50,
- order_by: Optional[str] = "created,tag",
+ order_by: str | None = "created,tag",
) -> UserLedgerTransactions:
user.prefetch_product(self.pg_config)
user_account = self.get_account_or_create_user_wallet(user)
diff --git a/generalresearch/managers/thl/maxmind/__init__.py b/generalresearch/managers/thl/maxmind/__init__.py
deleted file mode 100644
index 3bf0d07..0000000
--- a/generalresearch/managers/thl/maxmind/__init__.py
+++ /dev/null
@@ -1,162 +0,0 @@
-from typing import Collection, Optional
-
-import geoip2.models
-
-from generalresearch.managers.base import (
- Permission,
- PostgresManagerWithRedis,
-)
-from generalresearch.managers.thl.ipinfo import (
- GeoIpInfoManager,
- IPGeonameManager,
- IPInformationManager,
-)
-from generalresearch.managers.thl.maxmind.basic import MaxmindBasicManager
-from generalresearch.managers.thl.maxmind.insights import (
- get_insights_ip_information,
- should_call_insights,
-)
-from generalresearch.models.custom_types import IPvAnyAddressStr
-from generalresearch.models.thl.ipinfo import (
- GeoIPInformation,
- IPGeoname,
- IPInformation,
- normalize_ip,
-)
-from generalresearch.pg_helper import PostgresConfig
-from generalresearch.redis_helper import RedisConfig
-
-
-class MaxmindManager(PostgresManagerWithRedis):
- def __init__(
- self,
- maxmind_account_id: str,
- maxmind_license_key: str,
- pg_config: PostgresConfig,
- redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- ):
- self.ipinfo_manager = IPInformationManager(pg_config=pg_config)
- self.ipgeo_manager = IPGeonameManager(pg_config=pg_config)
- self.geoipinfo_manager = GeoIpInfoManager(
- pg_config=pg_config, redis_config=redis_config
- )
-
- self.basic_maxmind_manager = MaxmindBasicManager(
- data_dir="/tmp/",
- maxmind_account_id=maxmind_account_id,
- maxmind_license_key=maxmind_license_key,
- )
-
- self.maxmind_account_id = maxmind_account_id
- self.maxmind_license_key = maxmind_license_key
-
- super().__init__(
- pg_config=pg_config,
- redis_config=redis_config,
- permissions=permissions,
- )
-
- def store_basic_ip_information(self, res: geoip2.models.Country) -> None:
- geoname_id = res.country.geoname_id
- assert geoname_id, "Must have a Geoname ID to store"
-
- res_geo = self.ipgeo_manager.fetch_geoname_ids(filter_ids=[geoname_id])
- if len(res_geo) == 0:
- self.ipgeo_manager.create_basic(
- geoname_id=geoname_id,
- is_in_european_union=res.country.is_in_european_union,
- country_iso=res.country.iso_code,
- country_name=res.country.name,
- continent_name=res.continent.name,
- continent_code=res.continent.code,
- )
-
- self.ipinfo_manager.create_basic(
- ip=res.traits.ip_address,
- country_iso=res.country.iso_code,
- registered_country_iso=res.registered_country.iso_code,
- geoname_id=geoname_id,
- )
-
- def store_insights_ip_information(self, res: geoip2.models.Insights) -> None:
- ipinfo = IPInformation.from_insights(res)
- geoname_id = ipinfo.geoname_id
- res_geo = self.ipgeo_manager.fetch_geoname_ids([geoname_id])
- if len(res_geo) == 0:
- ipgeo = IPGeoname.from_insights(res)
- self.ipgeo_manager.create_or_update(ipgeo=ipgeo)
- self.ipinfo_manager.create_or_update(ipinfo=ipinfo)
-
- return None
-
- def get_or_create_ip_information(
- self,
- ip_address: IPvAnyAddressStr,
- force_insights: bool = False,
- ) -> Optional[GeoIPInformation]:
- """
- This is the 'top-level' IP handling call.
-
- - Check to see if we already 'know about' this IP. If so, return
- it. Otherwise:
- - Lookup basic or detailed info. Cache the result. maxmind lookup
- happens synchronously. If `pool`, the db operation happens async
- and we don't necessarily return the insights info.
- """
- res = self.geoipinfo_manager.get(ip_address)
- if res and (
- (force_insights is True and res.basic is False) or (force_insights is False)
- ):
- return res
- return self.run_ip_information(ip_address, force_insights=force_insights)
-
- def run_ip_information(
- self,
- ip_address: IPvAnyAddressStr,
- force_insights: bool = False,
- ) -> Optional[GeoIPInformation]:
- """
- Assumes this IP is "unknown" to us (not in the ipinformation table).
- Quick lookup IP using geoip2.Database. If its "good", lookup detailed
- info. Run db update.
- """
- # Quick lookup IP using geoip2.database
- basic_res = self.basic_maxmind_manager.get_basic_ip_information(ip_address)
- if basic_res is None:
- # IP is not 'valid'. We do nothing because if we see it again, it'll just hit the
- # geoip2.database (and redis and mysql_rr) which is ok... so no biggie.
- return None
-
- if force_insights or should_call_insights(res=basic_res):
- # IP is valid and country is good. Look up insights.
- return self.get_and_store_insights(ip_address)
-
- else:
- # IP is valid, but from a spammy country.
- self.store_basic_ip_information(res=basic_res)
- return self.geoipinfo_manager.get(ip_address)
-
- def get_and_store_insights(
- self,
- ip_address: IPvAnyAddressStr,
- ) -> GeoIPInformation:
-
- rc = self.redis_client
- normalized_ip, lookup_prefix = normalize_ip(ip_address)
- # Protect the actual calling of this with a lock
- with rc.lock(f"insights-lock:{normalized_ip}", timeout=2, blocking_timeout=1):
- # Check again we don't have it (or it is only the basic that is cached)
- res = self.geoipinfo_manager.get_cache(ip_address=ip_address)
- if res is not None and res.basic is False:
- return res
-
- res_mm = get_insights_ip_information(
- ip_address=normalized_ip,
- maxmind_account_id=self.maxmind_account_id,
- maxmind_license_key=self.maxmind_license_key,
- )
- self.store_insights_ip_information(res_mm)
- res = self.geoipinfo_manager.recreate_cache(ip_address)
-
- return res
diff --git a/generalresearch/managers/thl/maxmind/basic.py b/generalresearch/managers/thl/maxmind/basic.py
deleted file mode 100644
index 72479df..0000000
--- a/generalresearch/managers/thl/maxmind/basic.py
+++ /dev/null
@@ -1,133 +0,0 @@
-import logging
-import os
-import subprocess
-from datetime import timedelta
-from pathlib import Path
-from threading import RLock
-from typing import Optional, Union
-from uuid import uuid4
-
-import geoip2.database
-import geoip2.models
-import requests
-from cachetools import TTLCache, cached
-from geoip2.errors import AddressNotFoundError
-
-from generalresearch.managers.base import Manager
-from generalresearch.models.custom_types import (
- CountryISOLike,
- IPvAnyAddressStr,
-)
-
-logger = logging.getLogger()
-
-
-class MaxmindBasicManager(Manager):
-
- def __init__(
- self,
- data_dir: Union[str, Path],
- maxmind_account_id: str,
- maxmind_license_key: str,
- ):
-
- self.data_dir = data_dir
- self.maxmind_account_id = maxmind_account_id
- self.maxmind_license_key = maxmind_license_key
-
- self.run_update_geoip_db()
- super().__init__()
-
- @cached(
- cache=TTLCache(maxsize=1, ttl=timedelta(hours=1).total_seconds()),
- lock=RLock(),
- )
- def get_geoip_db(self):
- db_path = os.path.join(self.data_dir, "GeoIP2-Country.mmdb")
- return geoip2.database.Reader(fileish=db_path)
-
- def get_basic_ip_information(
- self, ip_address: IPvAnyAddressStr
- ) -> Optional[geoip2.models.Country]:
- try:
- return self.get_geoip_db().country(ip_address)
- except (ValueError, AddressNotFoundError):
- return None
-
- def get_country_iso_from_ip_geoip2db(
- self, ip: IPvAnyAddressStr
- ) -> Optional[CountryISOLike]:
- res = self.get_basic_ip_information(ip_address=ip)
- if res:
- return res.country.iso_code.lower()
-
- def run_update_geoip_db(self) -> None:
- # runs update_geoip_db with slack panic if fails
- db_path = os.path.join(self.data_dir, "GeoIP2-Country.mmdb")
- if os.path.exists(db_path):
- logger.info("GeoIP2-Country.mmdb already exists!")
- else:
- logger.info("Updating GeoIP2-Country.mmdb")
- try:
- self.update_geoip_db()
- except Exception as e:
- # TODO: Alert
- pass
-
- def update_geoip_db(self) -> None:
- """
- Download, checksum, extract from archive, confirm it works, then replace file on disk.
- # note: allowed 2,000 downloads per day, so I'm not bothering to implement
- # last modified or whatever checks.
- # https://support.maxmind.com/geoip-faq/databases-and-database-updates/is-there-a-limit-to-how-often-i-can
- -download-a-database-from-my-maxmind-account/
-
- """
- db_url = (
- f"https://download.maxmind.com/app/geoip_download?edition_id=GeoIP2-Country&"
- f"license_key={self.maxmind_license_key}&suffix=tar.gz"
- )
- sha256_url = (
- f"https://download.maxmind.com/app/geoip_download?edition_id=GeoIP2-Country&"
- f"license_key={self.maxmind_license_key}&suffix=tar.gz.sha256"
- )
- u = uuid4().hex
- cwd = f"/tmp/{u}/"
- os.makedirs(name=cwd, exist_ok=True)
-
- res = requests.get(db_url)
- # db_file_name looks like "GeoIP2-Country_20210806.tar.gz"
- db_file_name = res.headers.get("Content-Disposition").split("filename=")[1]
- tmp_db_file = cwd + db_file_name
- with open(tmp_db_file, "wb") as f:
- f.write(res.content)
- res = requests.get(sha256_url)
- tmp_sha256_file = cwd + "db.sha256"
- with open(tmp_sha256_file, "wb") as f:
- f.write(res.content)
- subprocess.check_call(args=["sha256sum", "-c", tmp_sha256_file], cwd=cwd)
- # Extract
- db_name = db_file_name.replace(".tar.gz", "")
- subprocess.check_call(
- args=[
- "tar",
- "-xf",
- tmp_db_file,
- "--strip-components",
- "1",
- f"{db_name}/GeoIP2-Country.mmdb",
- ],
- cwd=cwd,
- )
-
- # Confirm it works
- g = geoip2.database.Reader(fileish=cwd + "GeoIP2-Country.mmdb")
- g.country("111.111.111.111").country.iso_code.lower()
-
- # update file on disk
- prod_db = os.path.join(self.data_dir, "GeoIP2-Country.mmdb")
- subprocess.check_call(["mv", cwd + "GeoIP2-Country.mmdb", prod_db])
-
- # clean up
- assert cwd.startswith("/tmp/")
- subprocess.check_call(["rm", "-r", cwd])
diff --git a/generalresearch/managers/thl/maxmind/insights.py b/generalresearch/managers/thl/maxmind/insights.py
deleted file mode 100644
index 13dc3e8..0000000
--- a/generalresearch/managers/thl/maxmind/insights.py
+++ /dev/null
@@ -1,50 +0,0 @@
-import logging
-from typing import Optional
-
-import geoip2.models
-import geoip2.webservice
-from geoip2.errors import (
- AddressNotFoundError,
- AuthenticationError,
- InvalidRequestError,
- OutOfQueriesError,
-)
-
-logger = logging.getLogger()
-
-
-def get_insights_ip_information(
- ip_address: str,
- maxmind_account_id: str,
- maxmind_license_key: str,
-) -> Optional[geoip2.models.Insights]:
-
- # (2) We want more information, proceed further ($0.002)
- client = geoip2.webservice.Client(
- account_id=maxmind_account_id, license_key=maxmind_license_key, timeout=1
- )
- logger.info(f"get_insights_ip_information: {ip_address}")
- try:
- res = client.insights(ip_address)
-
- except (AuthenticationError, OutOfQueriesError) as e:
- # TODO: Alert
- return None
-
- except (AddressNotFoundError, InvalidRequestError):
- return None
- else:
- return res
-
-
-def should_call_insights(res: geoip2.models.Country) -> bool:
- """
- Call insights immediately if the IP is either:
- - in the continent of North America, Europe, or Oceania
- - in the country of Japan, Singapore, Israel, Hong Kong, Taiwan, South Korea
- """
- if res.continent.code.upper() in {"NA", "EU", "OC"}:
- return True
- if res.country.iso_code.upper() in {"JP", "SG", "IL", "HK", "TW", "KR"}:
- return True
- return False
diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py
index 0d00edb..43f7dc1 100644
--- a/generalresearch/managers/thl/payout.py
+++ b/generalresearch/managers/thl/payout.py
@@ -1,9 +1,12 @@
+from __future__ import annotations
+
from collections import defaultdict
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
from random import choice as rand_choice
from random import randint
from time import sleep
-from typing import Any, Collection, Dict, List, Optional, Union
+from typing import Any
from uuid import UUID, uuid4
import numpy as np
@@ -55,13 +58,11 @@ class PayoutEventManager(PostgresManagerWithRedis):
access
"""
- res = self.pg_config.execute_sql_query(
- query=f"""
+ 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}
@@ -91,10 +92,10 @@ class PayoutEventManager(PostgresManagerWithRedis):
def update(
self,
- payout_event: Union[UserPayoutEvent, BrokerageProductPayoutEvent],
+ payout_event: UserPayoutEvent | BrokerageProductPayoutEvent | None,
status: PayoutStatus,
- ext_ref_id: Optional[str] = None,
- order_data: Optional[Dict[str, Any]] = None,
+ ext_ref_id: str | None = None,
+ order_data: dict[str, Any] | None = None,
) -> None:
# These 3 things are the only modifiable attributes
ext_ref_id = ext_ref_id if ext_ref_id is not None else payout_event.ext_ref_id
@@ -102,15 +103,13 @@ class PayoutEventManager(PostgresManagerWithRedis):
payout_event.update(status=status, ext_ref_id=ext_ref_id, order_data=order_data)
d = payout_event.model_dump_mysql()
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE event_payout SET
status = %(status)s,
ext_ref_id = %(ext_ref_id)s,
order_data = %(order_data)s
WHERE uuid = %(uuid)s;
- """
- )
+ """)
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=d)
@@ -119,8 +118,6 @@ class PayoutEventManager(PostgresManagerWithRedis):
), "Nothing was updated! Are you sure this payout_event exists?"
conn.commit()
- return None
-
class UserPayoutEventManager(PayoutEventManager):
@@ -163,7 +160,7 @@ class UserPayoutEventManager(PayoutEventManager):
pe = self.get_by_uuid(pe_uuid=pe_uuid)
transaction_info = dict()
- order: Dict[str, Any] = pe.order_data
+ order: dict[str, Any] = pe.order_data
if pe.payout_type == PayoutType.TANGO and pe.status == PayoutStatus.COMPLETE:
reward = order["reward"]
if "credentialList" in reward:
@@ -190,17 +187,17 @@ class UserPayoutEventManager(PayoutEventManager):
def filter_by(
self,
- reference_uuid: Optional[str] = None,
- debit_account_uuids: Optional[Collection[UUIDStr]] = None,
- amount: Optional[int] = None,
- created: Optional[datetime] = None,
- created_after: Optional[datetime] = None,
- product_ids: Collection[str] = None,
- bp_user_ids: Optional[Collection[str]] = None,
- cashout_method_uuids: Collection[UUIDStr] = None,
- cashout_types: Optional[Collection[PayoutType]] = None,
- statuses: Optional[Collection[PayoutStatus]] = None,
- ) -> List[UserPayoutEvent]:
+ reference_uuid: 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,
+ cashout_method_uuids: Collection[UUIDStr] | None = None,
+ cashout_types: Collection[PayoutType] | None = None,
+ statuses: Collection[PayoutStatus] | None = None,
+ ) -> list[UserPayoutEvent]:
"""Try to retrieve payout events by the product_id/user_uuid, amount,
and optionally timestamp.
@@ -290,16 +287,16 @@ class UserPayoutEventManager(PayoutEventManager):
payout_type: PayoutType,
amount: PositiveInt,
# --- Optional: Default / Default Factory ---
- uuid: Optional[UUIDStr] = None,
- status: Optional[PayoutStatus] = None,
- created: Optional[AwareDatetimeISO] = None,
- request_data: Optional[Dict[str, Any]] = None,
+ uuid: UUIDStr | None = None,
+ status: PayoutStatus | None = None,
+ created: AwareDatetimeISO | None = None,
+ request_data: dict[str, Any] | None = None,
# --- Optional: None ---
- account_reference_type: Optional[str] = None,
- account_reference_uuid: Optional[UUIDStr] = None,
- description: Optional[str] = None,
- ext_ref_id: Optional[str] = None,
- order_data: Optional[Union[Dict[str, Any], CashMailOrderData]] = None,
+ account_reference_type: str | None = None,
+ account_reference_uuid: UUIDStr | None = None,
+ description: str | None = None,
+ ext_ref_id: str | None = None,
+ order_data: dict[str, Any] | CashMailOrderData | None = None,
) -> UserPayoutEvent:
payout_event = UserPayoutEvent(
@@ -344,19 +341,19 @@ class UserPayoutEventManager(PayoutEventManager):
def create_dummy(
self,
- uuid: Optional[UUIDStr] = None,
- debit_account_uuid: Optional[UUIDStr] = None,
- account_reference_type: Optional[str] = None,
- account_reference_uuid: Optional[UUIDStr] = None,
- cashout_method_uuid: Optional[UUIDStr] = None,
- description: Optional[str] = None,
- created: Optional[AwareDatetimeISO] = None,
- amount: Optional[PositiveInt] = None,
- status: Optional[PayoutStatus] = None,
- ext_ref_id: Optional[str] = None,
- payout_type: Optional[PayoutType] = None,
- request_data: Optional[Dict[str, Any]] = None,
- order_data: Optional[Union[Dict[str, Any], CashMailOrderData]] = None,
+ uuid: UUIDStr | None = None,
+ debit_account_uuid: UUIDStr | None = None,
+ account_reference_type: str | None = None,
+ account_reference_uuid: UUIDStr | None = None,
+ cashout_method_uuid: UUIDStr | None = None,
+ description: str | None = None,
+ created: AwareDatetimeISO | None = None,
+ amount: PositiveInt | None = None,
+ status: PayoutStatus | None = None,
+ ext_ref_id: str | None = None,
+ payout_type: PayoutType | None = None,
+ request_data: dict[str, Any] | None = None,
+ order_data: dict[str, Any] | CashMailOrderData | None = None,
) -> UserPayoutEvent:
debit_account_uuid = debit_account_uuid or uuid4().hex
@@ -398,7 +395,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
self,
pe_uuid: UUIDStr,
# --- Support resources ---
- account_product_mapping: Optional[Dict[UUIDStr, UUIDStr]] = None,
+ account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
) -> BrokerageProductPayoutEvent:
res = self.pg_config.execute_sql_query(
@@ -422,7 +419,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
# 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")
+ 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"]]
@@ -477,17 +474,17 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
def create(
self,
- uuid: Optional[UUIDStr] = None,
- debit_account_uuid: Optional[UUIDStr] = None,
+ uuid: UUIDStr | None = None,
+ debit_account_uuid: UUIDStr | None = None,
created: AwareDatetimeISO = None,
amount: PositiveInt = None,
- status: Optional[PayoutStatus] = None,
- ext_ref_id: Optional[str] = None,
+ status: PayoutStatus | None = None,
+ ext_ref_id: str | None = None,
payout_type: PayoutType = None,
- request_data: Optional[Dict[str, Any]] = None,
- order_data: Optional[Union[Dict[str, Any], CashMailOrderData]] = None,
+ request_data: dict[str, Any] | None = None,
+ order_data: dict[str, Any] | CashMailOrderData | None = None,
# --- Support resources ---
- account_product_mapping: Optional[Dict[UUIDStr, UUIDStr]] = None,
+ account_product_mapping: dict[UUIDStr, UUIDStr] | None = None,
) -> BrokerageProductPayoutEvent:
if request_data is None:
@@ -497,7 +494,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
# 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")
+ 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]
@@ -535,17 +532,17 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
def filter_by(
self,
- reference_uuid: Optional[str] = None,
- ext_ref_id: Optional[str] = None,
- debit_account_uuids: Optional[Collection[UUIDStr]] = None,
- amount: Optional[int] = None,
- created: Optional[datetime] = None,
- created_after: Optional[datetime] = None,
- product_ids: Collection[str] = None,
- bp_user_ids: Optional[Collection[str]] = None,
- cashout_types: Optional[Collection[PayoutType]] = None,
- statuses: Optional[Collection[PayoutStatus]] = None,
- ) -> List[BrokerageProductPayoutEvent]:
+ 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,
+ 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.
@@ -645,7 +642,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
def get_bp_payout_events_for_accounts(
self, accounts: Collection[LedgerAccount]
- ) -> List[BrokerageProductPayoutEvent]:
+ ) -> list[BrokerageProductPayoutEvent]:
return self.filter_by(
debit_account_uuids=[i.uuid for i in accounts],
cashout_types=[PayoutType.ACH],
@@ -655,8 +652,8 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
self,
thl_ledger_manager: ThlLedgerManager,
product_uuids: Collection[UUIDStr],
- order_by: Optional[OrderBy] = OrderBy.ASC,
- ) -> List["BrokerageProductPayoutEvent"]:
+ 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.
@@ -673,7 +670,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
rc = self.redis_client
account_product_mapping = rc.hgetall(name="pem:account_to_product")
- payout_events: List[BrokerageProductPayoutEvent] = (
+ payout_events: list[BrokerageProductPayoutEvent] = (
self.get_bp_payout_events_for_accounts(
accounts=accounts,
)
@@ -724,8 +721,8 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
product: Product,
amount: USDCent,
payout_type: PayoutType = PayoutType.ACH,
- ext_ref_id: Optional[str] = None,
- created: Optional[AwareDatetime] = 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:
@@ -803,7 +800,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
bp_pe: BrokerageProductPayoutEvent,
product: Product,
amount: USDCent,
- created: Optional[AwareDatetime] = None,
+ created: AwareDatetime | None = None,
skip_wallet_balance_check: bool = False,
skip_one_per_day_check: bool = False,
) -> BrokerageProductPayoutEvent:
@@ -846,20 +843,20 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
self,
thl_ledger_manager: ThlLedgerManager,
product: Product,
- ) -> List[BrokerageProductPayoutEvent]:
+ ) -> list[BrokerageProductPayoutEvent]:
account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
return self.get_bp_payout_events_for_accounts(accounts=[account])
def get_bp_payout_events_for_account(
self, account: LedgerAccount
- ) -> List[BrokerageProductPayoutEvent]:
+ ) -> list[BrokerageProductPayoutEvent]:
return self.get_bp_payout_events_for_accounts(accounts=[account])
def get_bp_payout_events_for_products(
self,
thl_ledger_manager: ThlLedgerManager,
product_uuids: Collection[UUIDStr],
- ) -> List[BrokerageProductPayoutEvent]:
+ ) -> list[BrokerageProductPayoutEvent]:
accounts = thl_ledger_manager.get_accounts_bp_wallet_for_products(
product_uuids=product_uuids
)
@@ -871,7 +868,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
def update_ext_reference_ids(
self,
new_value: str,
- current_value: Optional[str] = None,
+ current_value: str | None = None,
) -> None:
"""
There are scenarios where an ACH/Wire payout event was saved with
@@ -901,8 +898,6 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
assert c.rowcount < 10000
conn.commit()
- return None
-
def delete_failed_business_payout(self, ext_ref_id: str, thl_lm: ThlLedgerManager):
"""
Sometimes ACH/Wire payouts fail due to multiple reasons (timeouts,
@@ -992,8 +987,8 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
self,
thl_ledger_manager: ThlLedgerManager,
product_uuids: Collection[UUIDStr],
- order_by: Optional[OrderBy] = OrderBy.ASC,
- ) -> List["BusinessPayoutEvent"]:
+ 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,
@@ -1005,7 +1000,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
@staticmethod
def from_bp_payout_events(
bp_payout_events: Collection["BrokerageProductPayoutEvent"],
- ) -> List["BusinessPayoutEvent"]:
+ ) -> list["BusinessPayoutEvent"]:
if len(bp_payout_events) == 0:
return []
@@ -1022,7 +1017,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
@staticmethod
def recoup_proportional(
df: pd.DataFrame,
- target_amount: Union[USDCent, NonNegativeInt],
+ target_amount: USDCent | NonNegativeInt,
) -> pd.DataFrame:
"""
Recoup a target amount from rows proportionally based on a numeric column.
@@ -1151,9 +1146,9 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
amount: USDCent,
pm: ProductManager,
thl_lm: ThlLedgerManager,
- created: Optional[datetime] = None,
- transaction_id: Optional[str] = None,
- ) -> Optional[BusinessPayoutEvent]:
+ 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
@@ -1215,7 +1210,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
# Can't pay any Products that don't have an issue amount
res = res[res["issue_amount"] > 0]
- recouped_amounts: List[Dict[str, int]] = res[
+ recouped_amounts: list[dict[str, int]] = res[
["product_id", "remaining_balance", "issue_amount"]
].to_dict(orient="records")
@@ -1224,7 +1219,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
product_uuids=[i["product_id"] for i in recouped_amounts]
)
- bp_payouts: List[BrokerageProductPayoutEvent] = []
+ 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
diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py
index a3a8da6..36e8fc7 100644
--- a/generalresearch/managers/thl/product.py
+++ b/generalresearch/managers/thl/product.py
@@ -1,10 +1,13 @@
+from __future__ import annotations
+
import json
import logging
import operator
+from collections.abc import Collection
from datetime import datetime, timezone
from decimal import Decimal
from threading import Lock
-from typing import TYPE_CHECKING, Collection, List, Optional, Union
+from typing import TYPE_CHECKING
from uuid import UUID, uuid4
from cachetools import TTLCache, cachedmethod, keys
@@ -41,7 +44,7 @@ class ProductManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
):
super().__init__(pg_config=pg_config, permissions=permissions)
self.uuid_cache = TTLCache(maxsize=1024, ttl=5 * 60)
@@ -70,8 +73,8 @@ class ProductManager(PostgresManager):
def get_by_uuids(
self,
- product_uuids: List[UUIDStr],
- ) -> List["Product"]:
+ product_uuids: list[UUIDStr],
+ ) -> list["Product"]:
res = self.fetch_uuids(
product_uuids=product_uuids,
@@ -85,7 +88,7 @@ class ProductManager(PostgresManager):
def get_by_uuid_if_exists(
self,
product_uuid: UUIDStr,
- ) -> Optional["Product"]:
+ ) -> "Product" | None:
# many=False, raise_on_error=False
try:
return self.fetch_uuids(
@@ -98,18 +101,18 @@ class ProductManager(PostgresManager):
def get_by_uuids_if_exists(
self,
- product_uuids: List[UUIDStr],
- ) -> List["Product"]:
+ product_uuids: list[UUIDStr],
+ ) -> list["Product"]:
# Same as .get_by_uuids but doesn't raise Exception if len(product_uuids) != len(res)
return self.fetch_uuids(
product_uuids=product_uuids,
)
- def get_all(self, rand_limit: Optional[int]) -> List["Product"]:
+ def get_all(self, rand_limit: int | None) -> list["Product"]:
product_uuids = self.get_all_uuids(rand_limit=rand_limit)
return self.fetch_uuids(product_uuids=product_uuids)
- def get_all_uuids(self, rand_limit: Optional[int]) -> List[UUIDStr]:
+ def get_all_uuids(self, rand_limit: int | None) -> list[UUIDStr]:
if rand_limit:
res = self.pg_config.execute_sql_query(
@@ -123,20 +126,18 @@ class ProductManager(PostgresManager):
)
else:
- res = self.pg_config.execute_sql_query(
- query="""
+ res = self.pg_config.execute_sql_query(query="""
SELECT p.id::uuid
FROM userprofile_brokerageproduct AS p
- """
- )
+ """)
return [i["id"] for i in res]
def fetch_uuids(
self,
- product_uuids: Optional[List[UUIDStr]] = None,
- business_uuids: Optional[List[UUIDStr]] = None,
- team_uuids: Optional[List[UUIDStr]] = None,
- ) -> List["Product"]:
+ product_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ ) -> list["Product"]:
LOG.debug(f"PM.fetch_uuids({product_uuids=}, {business_uuids=}, {team_uuids=})")
assert (
@@ -179,8 +180,8 @@ class ProductManager(PostgresManager):
return res
def fetch_uuids_(
- self, c: Cursor, filter_uuids: List[UUIDStr], filter_column: str
- ) -> List["Product"]:
+ self, c: Cursor, filter_uuids: list[UUIDStr], filter_column: str
+ ) -> list["Product"]:
from generalresearch.models.thl.product import Product
assert len(filter_uuids) <= 500, "chunk me"
@@ -265,20 +266,20 @@ class ProductManager(PostgresManager):
def create_dummy(
self,
- product_id: Optional[UUIDStr] = None,
- team_id: Optional[UUIDStr] = None,
- business_id: Optional[UUIDStr] = None,
- name: Optional[str] = None,
- redirect_url: Optional[str] = None,
- harmonizer_domain: Optional[str] = None,
+ product_id: UUIDStr | None = None,
+ team_id: UUIDStr | None = None,
+ business_id: UUIDStr | None = None,
+ name: str | None = None,
+ redirect_url: str | None = None,
+ harmonizer_domain: str | None = None,
commission_pct: Decimal = Decimal("0.05000"),
- sources_config: Optional[Union["SourcesConfig", "SupplyConfigs"]] = None,
- payout_config: Optional["PayoutConfig"] = None,
- session_config: Optional["SessionConfig"] = None,
- profiling_config: Optional["ProfilingConfig"] = None,
- user_wallet_config: Optional["UserWalletConfig"] = None,
- user_create_config: Optional["UserCreateConfig"] = None,
- user_health_config: Optional["UserHealthConfig"] = None,
+ sources_config: "SourcesConfig" | "SupplyConfigs" | None = None,
+ payout_config: "PayoutConfig" | None = None,
+ session_config: "SessionConfig" | None = None,
+ profiling_config: "ProfilingConfig" | None = None,
+ user_wallet_config: "UserWalletConfig" | None = None,
+ user_create_config: "UserCreateConfig" | None = None,
+ 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
@@ -309,16 +310,16 @@ class ProductManager(PostgresManager):
team_id: UUIDStr,
name: str,
redirect_url: str,
- business_id: Optional[UUIDStr] = None,
- harmonizer_domain: Optional[str] = None,
+ business_id: UUIDStr | None = None,
+ harmonizer_domain: str | None = None,
commission_pct: Decimal = Decimal("0.05"),
- sources_config: Optional[Union["SourcesConfig", "SupplyConfigs"]] = None,
- payout_config: Optional["PayoutConfig"] = None,
- session_config: Optional["SessionConfig"] = None,
- profiling_config: Optional["ProfilingConfig"] = None,
- user_wallet_config: Optional["UserWalletConfig"] = None,
- user_create_config: Optional["UserCreateConfig"] = None,
- user_health_config: Optional["UserHealthConfig"] = None,
+ sources_config: "SourcesConfig" | "SupplyConfigs" | None = None,
+ payout_config: "PayoutConfig" | None = None,
+ session_config: "SessionConfig" | None = None,
+ profiling_config: "ProfilingConfig" | None = None,
+ user_wallet_config: "UserWalletConfig" | None = None,
+ user_create_config: "UserCreateConfig" | None = None,
+ user_health_config: "UserHealthConfig" | None = None,
) -> "Product":
"""Create a Product with all the basic defaults and return the instance"""
from generalresearch.models.thl.product import (
@@ -400,14 +401,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"] = list(self.pg_config.execute_sql_query(query="""
SELECT COALESCE(MAX(id_int), 0) + 1 as id_int
FROM userprofile_brokerageproduct
- """
- )
- )[0]["id_int"]
+ """))[0]["id_int"]
instance.id_int = insert_data["id_int"]
query = """
@@ -494,7 +491,7 @@ class ProductManager(PostgresManager):
raise ValueError(f"Not allowed to change: {keys_to_update & not_allowed}")
if not keys_to_update:
- return None
+ return
in_bp_keys = {
"name",
diff --git a/generalresearch/managers/thl/profiling/question.py b/generalresearch/managers/thl/profiling/question.py
index 7b2a7ad..10a9e32 100644
--- a/generalresearch/managers/thl/profiling/question.py
+++ b/generalresearch/managers/thl/profiling/question.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
import random
import threading
-from typing import Any, Collection, Dict, List, Tuple
+from collections.abc import Collection
+from typing import Any
from cachetools import TTLCache, cached
from pydantic import ValidationError
@@ -15,13 +18,13 @@ from generalresearch.models.thl.profiling.upk_question import (
class QuestionManager(PostgresManager):
- def get_multi_upk(self, question_ids: Collection[str]) -> List[UpkQuestion]:
+ def get_multi_upk(self, question_ids: Collection[str]) -> list[UpkQuestion]:
query = """
SELECT data, property_code, explanation_template, explanation_fragment_template
FROM marketplace_question
WHERE id = ANY(%(question_ids)s);
"""
- res: List[Dict[str, Any]] = self.pg_config.execute_sql_query(
+ res: list[dict[str, Any]] = self.pg_config.execute_sql_query(
query=query, params={"question_ids": list(question_ids)}
)
for x in res:
@@ -41,7 +44,7 @@ class QuestionManager(PostgresManager):
)
def get_questions_ranked(
self, country_iso: str, language_iso: str
- ) -> List[UpkQuestion]:
+ ) -> list[UpkQuestion]:
query = """
SELECT data, property_code, explanation_template, explanation_fragment_template
FROM marketplace_question
@@ -51,11 +54,11 @@ class QuestionManager(PostgresManager):
AND property_code NOT LIKE 'g:%%'
AND is_live
"""
- res: List[Dict[str, Any]] = self.pg_config.execute_sql_query(
+ res: list[dict[str, Any]] = self.pg_config.execute_sql_query(
query=query,
params={"country_iso": country_iso, "language_iso": language_iso},
)
- qs: List[UpkQuestion] = []
+ qs: list[UpkQuestion] = []
for x in res:
x["data"]["ext_question_id"] = x["property_code"]
x["data"]["explanation_template"] = x["explanation_template"]
@@ -95,7 +98,7 @@ class QuestionManager(PostgresManager):
"country_iso": country_iso,
"language_iso": language_iso,
}
- res: List[Dict[str, Any]] = self.pg_config.execute_sql_query(
+ res: list[dict[str, Any]] = self.pg_config.execute_sql_query(
query=query, params=params
)
assert len(res) == 1, f"expected 1, got {len(res)} results"
@@ -108,8 +111,8 @@ class QuestionManager(PostgresManager):
return UpkQuestion.model_validate(x["data"])
def filter_by_property(
- self, lookup: Collection[Tuple[str, str, str]]
- ) -> List[UpkQuestion]:
+ self, lookup: Collection[tuple[str, str, str]]
+ ) -> list[UpkQuestion]:
"""
lookup is [(property_code, country_iso, language_iso)]
"""
@@ -123,7 +126,7 @@ class QuestionManager(PostgresManager):
WHERE {where_str}
"""
flat_params = [item for tup in lookup for item in tup]
- res: List[Dict[str, Any]] = self.pg_config.execute_sql_query(
+ res: list[dict[str, Any]] = self.pg_config.execute_sql_query(
query=query, params=flat_params
)
for x in res:
@@ -141,7 +144,7 @@ class QuestionManager(PostgresManager):
LOG.warning(e)
return res2
- def update_question_explanation(self, q: UpkQuestion):
+ def update_question_explanation(self, q: UpkQuestion) -> None:
# Assuming the question already exists in the db, and we're updating
# the fields explanation_template and explanation_fragment_template
assert q.id, "q.id must be set"
@@ -160,4 +163,3 @@ class QuestionManager(PostgresManager):
c.execute(query, params)
assert c.rowcount == 1
conn.commit()
- return None
diff --git a/generalresearch/managers/thl/profiling/schema.py b/generalresearch/managers/thl/profiling/schema.py
index e209067..1e3dece 100644
--- a/generalresearch/managers/thl/profiling/schema.py
+++ b/generalresearch/managers/thl/profiling/schema.py
@@ -1,5 +1,6 @@
+from __future__ import annotations
+
from threading import RLock
-from typing import List
from uuid import UUID
from cachetools import TTLCache, cached
@@ -13,7 +14,7 @@ from generalresearch.models.thl.profiling.upk_property import (
class UpkSchemaManager(PostgresManager):
@cached(cache=TTLCache(maxsize=1, ttl=18 * 60), lock=RLock())
- def get_props_info(self) -> List[UpkProperty]:
+ def get_props_info(self) -> list[UpkProperty]:
query = """
SELECT
p.id AS property_id,
@@ -68,7 +69,7 @@ class UpkSchemaManager(PostgresManager):
c["id"] = UUID(c["id"]).hex
return [UpkProperty.model_validate(x) for x in res]
- def get_props_info_for_country(self, country_iso: str) -> List[UpkProperty]:
+ def get_props_info_for_country(self, country_iso: str) -> list[UpkProperty]:
assert country_iso.lower() == country_iso
res = self.get_props_info()
res = [x for x in res if x.country_iso == country_iso].copy()
diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py
index 1cab6c2..cbe39e7 100644
--- a/generalresearch/managers/thl/profiling/uqa.py
+++ b/generalresearch/managers/thl/profiling/uqa.py
@@ -1,6 +1,8 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
-from typing import Collection, List, Optional
from generalresearch.managers.base import PostgresManagerWithRedis
from generalresearch.models.thl.profiling.user_question_answer import (
@@ -25,7 +27,7 @@ class UQAManager(PostgresManagerWithRedis):
def update_cache(
self,
user: User,
- uqas: List[UserQuestionAnswer],
+ uqas: list[UserQuestionAnswer],
):
"""
Adds new answers to the redis cache for this user. If the cache
@@ -82,8 +84,8 @@ class UQAManager(PostgresManagerWithRedis):
def _dedupe_and_clean_uqas(
self,
- uqas: List[UserQuestionAnswer],
- ) -> List[UserQuestionAnswer]:
+ uqas: list[UserQuestionAnswer],
+ ) -> list[UserQuestionAnswer]:
# Remove anything older than 30 days
uqas = [uqa for uqa in uqas if not uqa.is_stale()]
@@ -98,14 +100,14 @@ class UQAManager(PostgresManagerWithRedis):
return sorted(new_uqas, key=lambda x: x.timestamp, reverse=True)
- def get(self, user: User) -> List[UserQuestionAnswer]:
+ def get(self, user: User) -> list[UserQuestionAnswer]:
uqas = self.get_from_cache(user=user)
if uqas is None:
uqas = self.recreate_cache(user)
return self._dedupe_and_clean_uqas(uqas)
- def get_from_cache(self, user: User) -> Optional[List[UserQuestionAnswer]]:
+ def get_from_cache(self, user: User) -> list[UserQuestionAnswer] | None:
redis_key = self.redis_key(user)
# Do the exists check and the list retrieval in a single transaction
@@ -123,7 +125,7 @@ class UQAManager(PostgresManagerWithRedis):
logger.info(f"{redis_key} exists")
return uqas
- def get_from_db(self, user: User) -> List[UserQuestionAnswer]:
+ 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)
@@ -174,14 +176,12 @@ class UQAManager(PostgresManagerWithRedis):
redis_key = f"thl-grpc:user-demographics:{user.user_id}"
self.redis_client.delete(redis_key)
- return None
-
def create(
self,
user: User,
- uqas: List[UserQuestionAnswer],
- session_id: Optional[str] = None,
- ):
+ uqas: list[UserQuestionAnswer],
+ session_id: str | None = None,
+ ) -> None:
for uqa in uqas:
if uqa.user_id is None:
uqa.user_id = user.user_id
@@ -189,11 +189,10 @@ class UQAManager(PostgresManagerWithRedis):
assert uqa.user_id == user.user_id
self.create_in_db(uqas=uqas, session_id=session_id)
self.update_cache(user=user, uqas=uqas)
- return None
def create_in_db(
- self, uqas: List[UserQuestionAnswer], session_id: Optional[str] = None
- ):
+ self, uqas: list[UserQuestionAnswer], session_id: str | None = None
+ ) -> None:
values = [uqa.model_dump_mysql(session_id=session_id) for uqa in uqas]
query = """
INSERT INTO marketplace_userquestionanswer
@@ -208,4 +207,3 @@ class UQAManager(PostgresManagerWithRedis):
with conn.cursor() as c:
c.executemany(query=query, params_seq=values)
conn.commit()
- return None
diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py
index b3ea52e..1820103 100644
--- a/generalresearch/managers/thl/profiling/user_upk.py
+++ b/generalresearch/managers/thl/profiling/user_upk.py
@@ -1,7 +1,10 @@
+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, Collection, Dict, List, Optional, Set, Tuple, Union
+from typing import Any
from uuid import UUID
from psycopg import Cursor
@@ -29,8 +32,8 @@ class UserUpkManager(PostgresManagerWithRedis):
self,
pg_config: PostgresConfig,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
):
super().__init__(
pg_config=pg_config,
@@ -42,9 +45,8 @@ class UserUpkManager(PostgresManagerWithRedis):
def clear_upk_cache(self, user_id: int) -> None:
self.redis_client.delete(f"thl-grpc:user-upk:{user_id}")
- return None
- def get_user_upk(self, user_id: int) -> List[UpkQuestionAnswer]:
+ def get_user_upk(self, user_id: int) -> list[UpkQuestionAnswer]:
res = self.redis_client.get(f"thl-grpc:user-upk:{user_id}")
if res:
return [UpkQuestionAnswer.model_validate(x) for x in json.loads(res)]
@@ -53,7 +55,7 @@ class UserUpkManager(PostgresManagerWithRedis):
self.redis_client.set(f"thl-grpc:user-upk:{user_id}", value, ex=60 * 60 * 24)
return res
- def get_user_upk_mysql(self, user_id: int) -> List[UpkQuestionAnswer]:
+ def get_user_upk_mysql(self, user_id: int) -> list[UpkQuestionAnswer]:
since = datetime.now(tz=timezone.utc) - timedelta(days=89)
query = """
@@ -119,11 +121,11 @@ class UserUpkManager(PostgresManagerWithRedis):
def get_user_upk_simple(
self, user_id: PositiveInt, country_iso: str = "us"
- ) -> Dict[str, Union[Set[str], str, float]]:
+ ) -> dict[str, set[str] | str | float]:
res = self.get_user_upk(user_id=user_id)
res = [x for x in res if x.country_iso == country_iso]
- d: Dict[str, Union[Set[str], str, float]] = defaultdict(set)
+ d: dict[str, set[str] | str | float] = defaultdict(set)
for x in res:
if x.cardinality == Cardinality.ZERO_OR_ONE:
d[x.property_label] = x.value
@@ -134,7 +136,7 @@ class UserUpkManager(PostgresManagerWithRedis):
def get_age_gender(
self, user_id: PositiveInt, country_iso: str = "us"
- ) -> Tuple[Optional[int], Optional[str]]:
+ ) -> tuple[int | None, str | None]:
# Returns an integer year for age, and {'male', 'female', 'other_gender'}
d = self.get_user_upk_simple(user_id, country_iso)
@@ -145,14 +147,14 @@ class UserUpkManager(PostgresManagerWithRedis):
gender = d.get("gender")
return age, gender
- def get_upk_schema(self, country_iso: str) -> List[UpkProperty]:
+ def get_upk_schema(self, country_iso: str) -> list[UpkProperty]:
return self.upk_schema_manager.get_props_info_for_country(
country_iso=country_iso
)
def populate_user_upk_from_dict(
- self, upk_ans_dict: List[Dict[str, Any]]
- ) -> List[UpkQuestionAnswer]:
+ self, upk_ans_dict: list[dict[str, Any]]
+ ) -> list[UpkQuestionAnswer]:
country_isos = {x["country_iso"] for x in upk_ans_dict}
assert len(country_isos) == 1
@@ -193,7 +195,7 @@ class UserUpkManager(PostgresManagerWithRedis):
upk_ans = [UpkQuestionAnswer.model_validate(x) for x in upk_ans_dict]
return upk_ans
- def upsert_user_profile_knowledge(self, c: Cursor, row: UpkQuestionAnswer):
+ def upsert_user_profile_knowledge(self, c: Cursor, row: UpkQuestionAnswer) -> None:
prop_type_table = {
PropertyType.UPK_ITEM: "marketplace_userprofileknowledgeitem",
PropertyType.UPK_NUMERICAL: "marketplace_userprofileknowledgenumerical",
@@ -244,7 +246,7 @@ class UserUpkManager(PostgresManagerWithRedis):
def upsert_user_profile_knowledge_multi_item(
self, c: Cursor, row: UpkQuestionAnswer
- ):
+ ) -> None:
args = row.model_dump_mysql()
c.execute(
@@ -299,9 +301,7 @@ class UserUpkManager(PostgresManagerWithRedis):
args,
)
- return None
-
- def set_user_upk(self, upk_ans: List[UpkQuestionAnswer]):
+ 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]
diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py
index 6e77031..bb467a2 100644
--- a/generalresearch/managers/thl/session.py
+++ b/generalresearch/managers/thl/session.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
from decimal import Decimal
-from typing import Any, Collection, Dict, List, Optional, Tuple
+from typing import Any
from uuid import UUID, uuid4
from faker import Faker
@@ -44,12 +47,12 @@ class SessionManager(PostgresManager):
self,
started: datetime,
user: User,
- country_iso: Optional[str] = None,
- device_type: Optional[DeviceType] = None,
- ip: Optional[str] = None,
- bucket: Optional[Bucket] = None,
- url_metadata: Optional[Dict[str, str]] = None,
- uuid_id: Optional[str] = None,
+ 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:
"""Creates a Session. Prefer to use this rather than instantiating the
model directly, because we're explicitly defining here which keys
@@ -70,8 +73,7 @@ class SessionManager(PostgresManager):
)
d = session.model_dump_mysql()
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO thl_session (
uuid, user_id, started, loi_min, loi_max,
user_payout_min, user_payout_max, country_iso,
@@ -81,8 +83,7 @@ class SessionManager(PostgresManager):
%(user_payout_min)s, %(user_payout_max)s, %(country_iso)s,
%(device_type)s, %(ip)s, %(url_metadata_json)s
) RETURNING id;
- """
- )
+ """)
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=d)
@@ -93,20 +94,20 @@ class SessionManager(PostgresManager):
def create_dummy(
self,
# -- Create Dummy "optional" -- #
- started: Optional[datetime] = None,
- user: Optional[User] = None,
+ started: datetime | None = None,
+ user: User | None = None,
# -- Optional -- #
- country_iso: Optional[str] = None,
- device_type: Optional[DeviceType] = None,
- ip: Optional[str] = None,
- bucket: Optional[Bucket] = None,
- url_metadata: Optional[Dict[str, str]] = None,
- uuid_id: Optional[str] = None,
+ 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),
- end_date=datetime(year=2000, month=1, day=1),
+ 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(
@@ -170,7 +171,7 @@ class SessionManager(PostgresManager):
assert len(res) == 1
return self.session_from_mysql(res[0])
- def session_from_mysql(self, d: Dict) -> Session:
+ def session_from_mysql(self, d: dict) -> Session:
d["id"] = d.pop("session_id")
d["uuid"] = UUID(d.pop("session_uuid")).hex
d["user"] = User(
@@ -210,12 +211,12 @@ class SessionManager(PostgresManager):
def finish_with_status(
self,
session: Session,
- finished: Optional[datetime] = None,
- status: Optional[Status] = None,
- status_code_1: Optional[StatusCode1] = None,
- status_code_2: Optional[SessionStatusCode2] = None,
- payout: Optional[Decimal] = None,
- user_payout: Optional[Decimal] = None,
+ finished: datetime | None = None,
+ status: Status | None = None,
+ status_code_1: StatusCode1 | None = None,
+ status_code_2: SessionStatusCode2 | None = None,
+ payout: Decimal | None = None,
+ user_payout: Decimal | None = None,
) -> Session:
# We have to update all the fields at once, or else we'll get
# validation errors. There doesn't seem to be a clean way of doing this.
@@ -249,7 +250,7 @@ class SessionManager(PostgresManager):
assert session.user.product, "prefetch product"
modified = session.adjust_status()
if not modified:
- return None
+ return
d = {
"adjusted_status": (
@@ -286,18 +287,18 @@ class SessionManager(PostgresManager):
def filter_paginated(
self,
- user_id: Optional[PositiveInt] = None,
- session_uuids: Optional[List[UUIDStr]] = None,
- product_uuids: Optional[List[UUIDStr]] = None,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
- status: Optional[Status] = None,
- adjusted_after: Optional[datetime] = None,
- adjusted_before: Optional[datetime] = None,
+ user_id: PositiveInt | None = None,
+ session_uuids: list[UUIDStr] | None = None,
+ product_uuids: list[UUIDStr] | None = None,
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
+ status: Status | None = None,
+ adjusted_after: datetime | None = None,
+ adjusted_before: datetime | None = None,
page: int = 1,
size: int = 100,
- order_by: Optional[str] = "-started",
- ) -> Tuple[List[Session], int]:
+ order_by: str | None = "-started",
+ ) -> tuple[list[Session], int]:
"""
Sessions are filtered using user, product_uuids, started_after, &
started_before (if set).
@@ -419,8 +420,8 @@ class SessionManager(PostgresManager):
def session_from_mysql_rows_json(
self,
- rows: Collection[Dict],
- ) -> List[Session]:
+ rows: Collection[dict],
+ ) -> list[Session]:
"""Columns: thl_session.*, thl_user.*, walls_json
- walls_json: list of objects, containing keys: thl_wall.*
"""
@@ -472,16 +473,16 @@ class SessionManager(PostgresManager):
@staticmethod
def make_filter_str(
- user_id: Optional[int] = None,
- session_uuids: Optional[List[UUIDStr]] = None,
- product_uuids: Optional[List[UUIDStr]] = None,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
- status: Optional[Status] = None,
- adjusted_after: Optional[datetime] = None,
- adjusted_before: Optional[datetime] = None,
- extra_filters: Optional[str] = None,
- ) -> Tuple[str, Dict[str, Any]]:
+ user_id: int | None = None,
+ session_uuids: list[UUIDStr] | None = None,
+ product_uuids: list[UUIDStr] | None = None,
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
+ status: Status | None = None,
+ adjusted_after: datetime | None = None,
+ adjusted_before: datetime | None = None,
+ extra_filters: str | None = None,
+ ) -> tuple[str, dict[str, Any]]:
filters = []
params = {}
@@ -547,15 +548,15 @@ class SessionManager(PostgresManager):
def filter(
self,
- started_since: Optional[datetime] = None,
- started_between: Optional[Tuple[datetime, datetime]] = None,
- user: Optional[User] = None,
- product_uuids: Optional[List[UUIDStr]] = None,
- team_uuids: Optional[List[UUIDStr]] = None,
- business_uuids: Optional[List[UUIDStr]] = None,
+ started_since: datetime | None = None,
+ started_between: tuple[datetime, datetime] | None = None,
+ user: User | None = None,
+ product_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
order_by: str = "-started",
- limit: Optional[int] = None,
- ) -> List[Session]:
+ limit: int | None = None,
+ ) -> list[Session]:
# to be deprecated ...
if team_uuids:
@@ -584,15 +585,15 @@ class SessionManager(PostgresManager):
def filter_count(
self,
- user_id: Optional[int] = None,
- session_uuids: Optional[List[UUIDStr]] = None,
- product_uuids: Optional[List[UUIDStr]] = None,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
- status: Optional[Status] = None,
- adjusted_after: Optional[datetime] = None,
- adjusted_before: Optional[datetime] = None,
- extra_filters: Optional[str] = None,
+ user_id: int | None = None,
+ session_uuids: list[UUIDStr] | None = None,
+ product_uuids: list[UUIDStr] | None = None,
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
+ status: Status | None = None,
+ adjusted_after: datetime | None = None,
+ adjusted_before: datetime | None = None,
+ extra_filters: str | None = None,
) -> NonNegativeInt:
filter_str, params = self.make_filter_str(
user_id=user_id,
@@ -620,7 +621,7 @@ class SessionManager(PostgresManager):
def get_task_status_response(
self, session_uuid: UUIDStr
- ) -> Optional[TaskStatusResponse]:
+ ) -> TaskStatusResponse | None:
res, total = self.filter_paginated(session_uuids=[session_uuid])
if total == 0:
return None
@@ -632,16 +633,16 @@ class SessionManager(PostgresManager):
def get_tasks_status_response(
self,
product_uuid: UUIDStr,
- user_id: Optional[int] = None,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
- status: Optional[Status] = None,
- adjusted_after: Optional[datetime] = None,
- adjusted_before: Optional[datetime] = None,
+ user_id: int | None = None,
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
+ status: Status | None = None,
+ adjusted_after: datetime | None = None,
+ adjusted_before: datetime | None = None,
page: int = 1,
size: int = 100,
- order_by: Optional[str] = "-started",
- ) -> Optional[TasksStatusResponse]:
+ order_by: str | None = "-started",
+ ) -> TasksStatusResponse | None:
PM = ProductManager(pg_config=self.pg_config, permissions=[Permission.READ])
product = PM.get_by_uuid(product_uuid=product_uuid)
diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py
index 1ce81d4..c7671ce 100644
--- a/generalresearch/managers/thl/survey.py
+++ b/generalresearch/managers/thl/survey.py
@@ -1,6 +1,9 @@
+from __future__ import annotations
+
from collections import defaultdict
+from collections.abc import Collection
from datetime import datetime, timezone
-from typing import Any, Collection, Dict, List, Optional, Tuple
+from typing import Any
import pandas as pd
from more_itertools import chunked
@@ -24,13 +27,13 @@ class SurveyManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
):
super().__init__(pg_config=pg_config, permissions=permissions)
self.buyer_manager = BuyerManager(pg_config=pg_config, permissions=permissions)
self.category_manager = CategoryManager(pg_config=pg_config)
- def create_or_update(self, surveys: List[Survey]):
+ def create_or_update(self, surveys: list[Survey]):
"""
The only field that is checked for a possible update is `is_live`!
"""
@@ -76,12 +79,11 @@ class SurveyManager(PostgresManager):
"survey_updated_count": len(to_update),
}
- def create_bulk(self, surveys: List[Survey]):
+ def create_bulk(self, surveys: list[Survey]) -> None:
for chunk in chunked(surveys, 500):
self.create_bulk_chunk(chunk)
- return None
- def create_bulk_chunk(self, surveys: List[Survey]):
+ def create_bulk_chunk(self, surveys: list[Survey]) -> None:
assert len(surveys) <= 500, "chunk me"
query = """
@@ -97,9 +99,8 @@ class SurveyManager(PostgresManager):
with conn.cursor() as c:
c.executemany(query=query, params_seq=params)
conn.commit()
- return None
- def update_is_live(self, surveys: List[Survey]):
+ def update_is_live(self, surveys: list[Survey]):
ids_ON = [s.id for s in surveys if s.is_live]
ids_OFF = [s.id for s in surveys if not s.is_live]
query_ON = """
@@ -127,7 +128,7 @@ class SurveyManager(PostgresManager):
self,
survey_keys: Collection[SurveyKey],
include_categories: bool = False,
- ) -> List[Survey]:
+ ) -> list[Survey]:
assert len(survey_keys) <= 1000
if len(survey_keys) == 0:
@@ -202,7 +203,7 @@ class SurveyManager(PostgresManager):
def filter_by_natural_key(
self, source: Source, survey_ids: Collection[str]
- ) -> List[Survey]:
+ ) -> list[Survey]:
res = []
for chunk in chunked(survey_ids, 1000):
res.extend(self.filter_by_natural_key_chunk(source, chunk))
@@ -210,7 +211,7 @@ class SurveyManager(PostgresManager):
def filter_by_natural_key_chunk(
self, source: Source, survey_ids: Collection[str]
- ) -> List[Survey]:
+ ) -> list[Survey]:
query = """
SELECT id, source, survey_id, created_at, updated_at,
is_live, is_recontact, buyer_id, eligibility_criteria
@@ -224,7 +225,7 @@ class SurveyManager(PostgresManager):
)
return [Survey.model_validate(x) for x in res]
- def filter_by_source_live(self, source: Source) -> List[Survey]:
+ def filter_by_source_live(self, source: Source) -> list[Survey]:
"""
Return all live surveys for this source
"""
@@ -237,7 +238,7 @@ class SurveyManager(PostgresManager):
res = self.pg_config.execute_sql_query(query, params={"source": source.value})
return [Survey.model_validate(x) for x in res]
- def filter_by_live(self, fields: Optional[List[str]] = None) -> List[Survey]:
+ def filter_by_live(self, fields: list[str] | None = None) -> list[Survey]:
"""
Return all live surveys
"""
@@ -281,29 +282,27 @@ class SurveyManager(PostgresManager):
)
return None
- def update_surveys_categories(self, surveys: List[Survey] = None) -> None:
+ def update_surveys_categories(self, surveys: list[Survey] | None = None) -> None:
for chunk in chunked(surveys, 500):
self.update_surveys_categories_chunk(chunk)
- return None
- def update_surveys_categories_chunk(self, surveys: List[Survey] = None) -> None:
+ def update_surveys_categories_chunk(
+ self, surveys: list[Survey] | None = None
+ ) -> None:
assert len(surveys) <= 500, "chunk me"
- temp_table_sql = sql.SQL(
- """
+ temp_table_sql = sql.SQL("""
CREATE TEMP TABLE tmp_survey_categories (
survey_id bigint,
category_id int,
strength float8
) ON COMMIT DROP;
- """
- )
+ """)
# noinspection SqlResolve
insert_values_sql = sql.SQL(
"INSERT INTO tmp_survey_categories VALUES (%s, %s, %s)"
)
# noinspection SqlResolve
- delete_sql = sql.SQL(
- """
+ delete_sql = sql.SQL("""
DELETE FROM marketplace_surveycategory sc
WHERE NOT EXISTS (
SELECT 1
@@ -313,18 +312,15 @@ class SurveyManager(PostgresManager):
)
AND sc.survey_id IN (
SELECT DISTINCT survey_id FROM tmp_survey_categories
- );"""
- )
+ );""")
# noinspection SqlResolve
- upsert_sql = sql.SQL(
- """
+ upsert_sql = sql.SQL("""
INSERT INTO marketplace_surveycategory (survey_id, category_id, strength)
SELECT survey_id, category_id, strength
FROM tmp_survey_categories
ON CONFLICT (survey_id, category_id)
DO UPDATE SET
- strength = EXCLUDED.strength;"""
- )
+ strength = EXCLUDED.strength;""")
rows = [
(survey.id, c.category.id, c.strength)
@@ -417,7 +413,7 @@ class SurveyStatManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Collection[Permission] = None,
+ permissions: Collection[Permission] | None = None,
):
super().__init__(pg_config=pg_config, permissions=permissions)
self.survey_manager = SurveyManager(
@@ -457,8 +453,8 @@ class SurveyStatManager(PostgresManager):
# info.register(conn)
def update_or_create(
- self, survey_stats: List[SurveyStat]
- ) -> Optional[List[SurveyStat]]:
+ self, survey_stats: list[SurveyStat]
+ ) -> list[SurveyStat] | None:
"""
This manager is NOT responsible for creating surveys or buyers.
It will check to make sure they exist
@@ -502,10 +498,9 @@ class SurveyStatManager(PostgresManager):
# survey_stats = sorted(survey_stats, key=lambda s: s.natural_key)
# return survey_stats
- def upsert_sql(self, survey_stats: List[SurveyStat]) -> None:
+ def upsert_sql(self, survey_stats: list[SurveyStat]) -> None:
for chunk in chunked(survey_stats, 1000):
self.upsert_sql_chunk(survey_stats=chunk)
- return None
# def insert_sql(self, survey_stats: List[SurveyStat]):
# for chunk in chunked(survey_stats, 1000):
@@ -532,7 +527,7 @@ class SurveyStatManager(PostgresManager):
# conn.commit()
# return None
- def upsert_sql_chunk(self, survey_stats: List[SurveyStat]) -> None:
+ def upsert_sql_chunk(self, survey_stats: list[SurveyStat]) -> None:
assert len(survey_stats) <= 1000, "chunk me"
keys = self.KEYS
keys_str = ", ".join(keys)
@@ -557,16 +552,14 @@ class SurveyStatManager(PostgresManager):
c.executemany(query=query, params_seq=params)
conn.commit()
- return None
-
- def filter_by_unique_keys(self, keys: Collection[Tuple]) -> List[SurveyStat]:
+ def filter_by_unique_keys(self, keys: Collection[tuple]) -> list[SurveyStat]:
res = []
for chunk in chunked(keys, 5000):
res.extend(self.filter_by_unique_keys_chunk(chunk))
return res
- def filter_by_unique_keys_chunk(self, keys: Collection[Tuple]):
+ def filter_by_unique_keys_chunk(self, keys: Collection[tuple]):
values_sql = ", ".join(["(%s, %s, %s, %s)"] * len(keys))
query = f"""
SELECT
@@ -590,8 +583,8 @@ class SurveyStatManager(PostgresManager):
def update_surveystats_for_source(
self,
source: Source,
- surveys: List[Survey],
- survey_stats: List[SurveyStat],
+ surveys: list[Survey],
+ survey_stats: list[SurveyStat],
):
"""
What ym-survey-stats actually calls.
@@ -634,13 +627,13 @@ class SurveyStatManager(PostgresManager):
def make_filter_str(
self,
- is_live: Optional[bool] = True,
- updated_after: Optional[datetime] = None,
- min_score: Optional[float] = None,
- survey_keys: Optional[Collection[SurveyKey]] = None,
- sources: Optional[Collection[Source]] = None,
- country_iso: Optional[str] = None,
- ) -> Tuple[str, Dict[str, Any]]:
+ is_live: bool | None = True,
+ updated_after: datetime | None = None,
+ min_score: float | None = None,
+ survey_keys: Collection[SurveyKey] | None = None,
+ sources: Collection[Source] | None = None,
+ country_iso: str | None = None,
+ ) -> tuple[str, dict[str, Any]]:
filters = []
params = dict()
if updated_after is not None:
@@ -686,12 +679,12 @@ class SurveyStatManager(PostgresManager):
def filter_count(
self,
- is_live: Optional[bool] = True,
- updated_after: Optional[datetime] = None,
- min_score: Optional[float] = None,
- survey_keys: Optional[Collection[SurveyKey]] = None,
- sources: Optional[Collection[Source]] = None,
- country_iso: Optional[str] = None,
+ is_live: bool | None = True,
+ updated_after: datetime | None = None,
+ min_score: float | None = None,
+ survey_keys: Collection[SurveyKey] | None = None,
+ sources: Collection[Source] | None = None,
+ country_iso: str | None = None,
) -> NonNegativeInt:
filter_str, params = self.make_filter_str(
is_live=is_live,
@@ -710,17 +703,17 @@ class SurveyStatManager(PostgresManager):
def filter(
self,
- is_live: Optional[bool] = True,
- updated_after: Optional[datetime] = None,
- min_score: Optional[float] = None,
- survey_keys: Optional[Collection[SurveyKey]] = None,
- sources: Optional[Collection[Source]] = None,
- country_iso: Optional[str] = None,
- page: Optional[int] = None,
- size: Optional[int] = None,
- order_by: Optional[str] = None,
- debug: Optional[bool] = False,
- ) -> List[SurveyStat]:
+ is_live: bool | None = True,
+ updated_after: datetime | None = None,
+ min_score: float | None = None,
+ survey_keys: Collection[SurveyKey] | None = None,
+ sources: Collection[Source] | None = None,
+ country_iso: str | None = None,
+ page: int | None = None,
+ size: int | None = None,
+ order_by: str | None = None,
+ debug: bool | None = False,
+ ) -> list[SurveyStat]:
filter_str, params = self.make_filter_str(
is_live=is_live,
updated_after=updated_after,
@@ -778,10 +771,10 @@ class SurveyStatManager(PostgresManager):
def filter_to_merge_table(
self,
- is_live: Optional[bool] = True,
- updated_after: Optional[datetime] = None,
- min_score: Optional[float] = 0.0001,
- ) -> Optional[pd.DataFrame]:
+ is_live: bool | None = True,
+ updated_after: datetime | None = None,
+ min_score: float | None = 0.0001,
+ ) -> pd.DataFrame | None:
survey_stats = self.filter(
is_live=is_live, updated_after=updated_after, min_score=min_score
diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py
index 3c402ee..efaa930 100644
--- a/generalresearch/managers/thl/survey_penalty.py
+++ b/generalresearch/managers/thl/survey_penalty.py
@@ -1,8 +1,9 @@
+from __future__ import annotations
+
import json
import threading
from collections import defaultdict
from datetime import timedelta
-from typing import Dict, List, Optional, Tuple
from cachetools import TTLCache, cachedmethod
@@ -36,7 +37,7 @@ class SurveyPenaltyManager(RedisManager):
def __init__(
self,
redis_config: RedisConfig,
- cache_prefix: Optional[str] = None,
+ cache_prefix: str | None = None,
**kwargs,
):
super().__init__(redis_config=redis_config, cache_prefix=cache_prefix, **kwargs)
@@ -59,7 +60,7 @@ class SurveyPenaltyManager(RedisManager):
def get_redis_key_for_id(self, uuid_id: UUIDStr):
return f"{self.redis_prefix}:{uuid_id}"
- def set_penalties(self, penalties: List[Penalty]):
+ def set_penalties(self, penalties: list[Penalty]):
""" """
if len(penalties) > 1000:
LOG.warning("SurveyPenaltyManager.set_penalties batch me!")
@@ -81,7 +82,7 @@ class SurveyPenaltyManager(RedisManager):
def _load_penalties(
self, product_id: UUIDStr, team_id: UUIDStr
- ) -> Tuple[List[BPSurveyPenalty], List[TeamSurveyPenalty]]:
+ ) -> tuple[list[BPSurveyPenalty], list[TeamSurveyPenalty]]:
pipe = self.redis_client.pipeline(transaction=False)
bp_res, team_res = (
pipe.hgetall(self.get_redis_key_for_id(product_id))
@@ -99,7 +100,7 @@ class SurveyPenaltyManager(RedisManager):
@cachedmethod(lambda self: self.cache, lock=lambda self: self.cache_lock)
def get_penalties_for(
self, product_id: UUIDStr, team_id: UUIDStr
- ) -> Dict[str, float]:
+ ) -> dict[str, float]:
"""
Returns a dict with keys survey sids ({source}:{survey_id}) and values penalties.
e.g. {'s:1234': 0.8}
diff --git a/generalresearch/managers/thl/tango_api.py b/generalresearch/managers/thl/tango_api.py
index 5c1706e..657224e 100644
--- a/generalresearch/managers/thl/tango_api.py
+++ b/generalresearch/managers/thl/tango_api.py
@@ -1,5 +1,5 @@
from decimal import Decimal
-from typing import Any, Dict, List
+from typing import Any
import requests
from pydantic import BaseModel
@@ -104,7 +104,7 @@ class TangoClient:
f"/accounts/{account_identifier}",
)
- def get_catalog(self) -> Dict[str, List[Dict[str, Any]]]:
+ def get_catalog(self) -> dict[str, list[dict[str, Any]]]:
"""
Replacement for:
api_client.catalog.get_catalog()
diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py
index b1b509b..e4736d4 100644
--- a/generalresearch/managers/thl/task_adjustment.py
+++ b/generalresearch/managers/thl/task_adjustment.py
@@ -1,8 +1,9 @@
+from __future__ import annotations
+
import logging
from datetime import datetime, timezone
from decimal import Decimal
from functools import cached_property
-from typing import List, Optional
from generalresearch.managers import parse_order_by
from generalresearch.managers.base import (
@@ -39,8 +40,8 @@ class TaskAdjustmentManager(PostgresManager):
wall_uuid: UUIDStr,
page: int = 1,
size: int = 100,
- order_by: Optional[str] = "-created",
- ) -> List[TaskAdjustmentEvent]:
+ order_by: str | None = "-created",
+ ) -> list[TaskAdjustmentEvent]:
params = {"wall_uuid": wall_uuid}
order_by_str = parse_order_by(order_by)
paginated_filter_str = "LIMIT %(limit)s OFFSET %(offset)s"
@@ -106,9 +107,9 @@ class TaskAdjustmentManager(PostgresManager):
ledger_manager: ThlLedgerManager,
wall_uuid: str,
adjusted_status: WallAdjustedStatus,
- alert_time: Optional[datetime] = None,
- ext_status_code: Optional[str] = None,
- adjusted_cpi: Optional[Decimal] = None,
+ alert_time: datetime | None = None,
+ ext_status_code: str | None = None,
+ adjusted_cpi: Decimal | None = None,
) -> None:
"""
We just got an adjustment notification from a marketplace.
@@ -172,7 +173,7 @@ class TaskAdjustmentManager(PostgresManager):
)
except AssertionError as e:
logging.warning(e)
- return None
+ return
event = TaskAdjustmentEvent(
adjusted_status=adjusted_status,
@@ -198,4 +199,4 @@ class TaskAdjustmentManager(PostgresManager):
self.session_manager.adjust_status(session)
ledger_manager.create_tx_bp_adjustment(session, created=alert_time)
- return None
+ return
diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py
index 7cd16fe..543de87 100644
--- a/generalresearch/managers/thl/user_compensate.py
+++ b/generalresearch/managers/thl/user_compensate.py
@@ -1,6 +1,7 @@
+from __future__ import annotations
+
from datetime import datetime, timezone
from decimal import Decimal
-from typing import Optional
from uuid import uuid4
from pydantic import NonNegativeInt
@@ -16,9 +17,9 @@ def user_compensate(
ledger_manager: ThlLedgerManager,
user: User,
amount_int: NonNegativeInt,
- ext_ref: Optional[str] = None,
- description: Optional[str] = None,
- skip_flag_check: Optional[bool] = False,
+ ext_ref: str | None = None,
+ description: str | None = None,
+ skip_flag_check: bool | None = False,
) -> UUIDStr:
"""
Compensate a user. aka "bribe". The money is paid out of the BP's
diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py
index 8875de8..0392edb 100644
--- a/generalresearch/managers/thl/user_manager/__init__.py
+++ b/generalresearch/managers/thl/user_manager/__init__.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import csv
import logging
import os
@@ -5,7 +7,7 @@ import threading
import time
from pathlib import Path
from threading import RLock
-from typing import Any, Dict, Union
+from typing import Any
from cachetools import TTLCache, cached
@@ -53,7 +55,7 @@ def get_bp_trust_df():
convert_int = lambda x: int(float(x))
-def parse_bp_trust_df(fp: Union[str, Path]) -> Dict[str, Any]:
+def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]:
dtype = {
"bp_trust": float,
"team_trust": float,
diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
index 299b167..dbed5de 100644
--- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py
+++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
@@ -1,7 +1,9 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Collection
from datetime import datetime, timezone
from functools import lru_cache
-from typing import Collection, List, Optional
from uuid import uuid4
import psycopg
@@ -37,12 +39,12 @@ class MysqlUserManager:
def get_user_from_mysql(
self,
*,
- product_id: Optional[str] = None,
- product_user_id: Optional[str] = None,
- user_id: Optional[int] = None,
- user_uuid: Optional[UUIDStr] = None,
+ product_id: str | None = None,
+ product_user_id: str | None = None,
+ user_id: int | None = None,
+ user_uuid: UUIDStr | None = None,
can_use_read_replica: bool = True,
- ) -> Optional[User]:
+ ) -> User | None:
logger.info(
f"get_user_from_mysql: {product_id}, {product_user_id}, {user_id}, {user_uuid}"
@@ -108,7 +110,7 @@ class MysqlUserManager:
self,
product_user_id: str,
product_id: str,
- created: Optional[datetime] = None,
+ created: datetime | None = None,
) -> User:
"""Creates a thl_user record for a new user."""
assert self.is_read_replica is False
@@ -127,16 +129,14 @@ class MysqlUserManager:
}
# in postgres, you do not include the auto-increment id column
- query = sql.SQL(
- """
+ query = sql.SQL("""
INSERT INTO thl_user
(uuid, product_id, product_user_id, created,
last_seen, blocked, last_country_iso, last_geoname_id, last_ip)
VALUES (%(user_uuid)s, %(product_id)s, %(product_user_id)s, %(created)s,
%(last_seen)s, FALSE, NULL, NULL, NULL)
RETURNING id;
- """
- )
+ """)
try:
with self.pg_config.make_connection() as conn:
@@ -221,7 +221,7 @@ class MysqlUserManager:
*,
product_id: str,
product_user_ids: Collection[str],
- ) -> List[User]:
+ ) -> list[User]:
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"
@@ -247,9 +247,9 @@ class MysqlUserManager:
def fetch(
self,
*,
- user_ids: Collection[int] = None,
- user_uuids: Collection[str] = None,
- ) -> List[User]:
+ 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"
diff --git a/generalresearch/managers/thl/user_manager/redis_user_manager.py b/generalresearch/managers/thl/user_manager/redis_user_manager.py
index 6c869bc..2190e6e 100644
--- a/generalresearch/managers/thl/user_manager/redis_user_manager.py
+++ b/generalresearch/managers/thl/user_manager/redis_user_manager.py
@@ -1,4 +1,4 @@
-from typing import Optional
+from __future__ import annotations
import redis
from pydantic import RedisDsn
@@ -10,8 +10,8 @@ class RedisUserManager:
def __init__(
self,
redis_dsn: RedisDsn,
- cache_prefix: Optional[str] = None,
- redis_timeout: Optional[float] = None,
+ cache_prefix: str | None = None,
+ redis_timeout: float | None = None,
):
self.redis = redis_dsn
self.redis_timeout = redis_timeout if redis_timeout else 0.10
@@ -32,11 +32,11 @@ class RedisUserManager:
def get_user(
self,
*,
- product_id: Optional[str] = None,
- product_user_id: Optional[str] = None,
- user_id: Optional[int] = None,
- user_uuid: Optional[str] = None,
- ) -> Optional[User]:
+ product_id: str | None = None,
+ product_user_id: str | None = None,
+ user_id: int | None = None,
+ user_uuid: str | None = None,
+ ) -> User | None:
# assume we did input validation in user_manager.get_user() function
if user_uuid:
d = self.client.get(f"{self.cache_prefix}:uuid:{user_uuid}")
@@ -75,8 +75,6 @@ class RedisUserManager:
p.execute()
- return None
-
def clear_user(self, user: User) -> None:
# this should only be used by tests
with self.client.pipeline(transaction=False) as p:
@@ -86,5 +84,3 @@ class RedisUserManager:
f"{self.cache_prefix}:ubp:{user.product_id}:{user.product_user_id}"
)
p.execute()
-
- return None
diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py
index ffd8191..a7bbd7e 100644
--- a/generalresearch/managers/thl/user_manager/user_manager.py
+++ b/generalresearch/managers/thl/user_manager/user_manager.py
@@ -1,7 +1,9 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Collection
from datetime import datetime
from functools import lru_cache
-from typing import Collection, List, Optional
from uuid import uuid4
from pydantic import RedisDsn
@@ -32,12 +34,12 @@ auditlog = logging.getLogger("auditlog")
class UserManager:
def __init__(
self,
- redis: Optional[RedisDsn] = None,
- pg_config: Optional[PostgresConfig] = None,
- pg_config_rr: Optional[PostgresConfig] = None,
- sql_permissions: Optional[Collection[Permission]] = None,
- cache_prefix: Optional[str] = None,
- redis_timeout: Optional[float] = None,
+ redis: RedisDsn | None = None,
+ pg_config: PostgresConfig | None = None,
+ pg_config_rr: PostgresConfig | None = None,
+ sql_permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
+ redis_timeout: float | None = None,
):
if sql_permissions is None:
@@ -87,8 +89,8 @@ class UserManager:
user: User,
level: int,
event_type: str,
- event_msg: Optional[str] = None,
- event_value: Optional[float] = None,
+ event_msg: str | None = None,
+ event_value: float | None = None,
) -> None:
from generalresearch.managers.thl.userhealth import AuditLogManager
from generalresearch.models.thl.userhealth import AuditLogLevel
@@ -115,10 +117,10 @@ class UserManager:
def get_user(
self,
*,
- product_id: Optional[str] = None,
- product_user_id: Optional[str] = None,
- user_id: Optional[int] = None,
- user_uuid: Optional[UUIDStr] = None,
+ product_id: str | None = None,
+ product_user_id: str | None = None,
+ user_id: int | None = None,
+ user_uuid: UUIDStr | None = None,
) -> User:
"""
Retrieve User from (product_id & product_user_id) or (user_id), or (uuid).
@@ -171,9 +173,9 @@ class UserManager:
def get_user_if_exists(
self,
- product_id: Optional[str] = None,
- product_user_id: Optional[str] = None,
- ) -> Optional[User]:
+ product_id: str | None = None,
+ product_user_id: str | None = None,
+ ) -> User | None:
"""
Look up User from (product_id & product_user_id). Returns
None if user does not exist.
@@ -186,11 +188,11 @@ class UserManager:
def get_user_inmemory_cache(
self,
*,
- product_id: Optional[str] = None,
- product_user_id: Optional[str] = None,
- user_id: Optional[int] = None,
- user_uuid: Optional[UUIDStr] = None,
- ) -> Optional[User]:
+ product_id: str | None = None,
+ product_user_id: str | None = None,
+ user_id: int | None = None,
+ user_uuid: UUIDStr | None = None,
+ ) -> User | None:
input_str = f"{product_id}, {product_user_id}, {user_id}, {user_uuid}"
if self.redis_user_manager:
@@ -226,15 +228,11 @@ class UserManager:
) as e:
logger.info(f"redis.set_user failed: {user}, {e}")
- return None
-
def clear_user_inmemory_cache(self, user: User) -> None:
if self.redis_user_manager:
# this should only be used by tests
self.redis_user_manager.clear_user(user)
- return None
-
def get_or_create_user(self, product_user_id: str, product_id: str) -> User:
"""
Given a bp_user_id and a product_id, get or create a User
@@ -266,9 +264,9 @@ class UserManager:
def create_user(
self,
product_user_id: str,
- product_id: Optional[UUIDStr] = None,
- product: Optional[Product] = None,
- created: Optional[datetime] = None,
+ product_id: UUIDStr | None = None,
+ product: Product | None = None,
+ created: datetime | None = None,
) -> User:
assert (
@@ -299,11 +297,11 @@ class UserManager:
def create_dummy(
self,
# --- Create dummy "optional" --- #
- product_user_id: Optional[str] = None,
+ product_user_id: str | None = None,
# --- Optional --- #
- product_id: Optional[UUIDStr] = None,
- product: Optional[Product] = None,
- created: Optional[datetime] = None,
+ product_id: UUIDStr | None = None,
+ product: Product | None = None,
+ created: datetime | None = None,
) -> User:
product_user_id = product_user_id or uuid4().hex
@@ -357,7 +355,7 @@ class UserManager:
*,
product_id: str,
product_user_ids: Collection[str],
- ) -> List[User]:
+ ) -> list[User]:
assert product_id, "must pass product_id"
assert len(product_user_ids) > 0, "must pass 1 or more product_user_ids"
return self.mysql_user_manager_rr.fetch_by_bpuids(
@@ -367,9 +365,9 @@ class UserManager:
def fetch(
self,
*,
- user_ids: Collection[int] = None,
- user_uuids: Collection[str] = None,
- ) -> List[User]:
+ 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"
diff --git a/generalresearch/managers/thl/user_manager/user_metadata_manager.py b/generalresearch/managers/thl/user_manager/user_metadata_manager.py
index a97b214..ffa8b44 100644
--- a/generalresearch/managers/thl/user_manager/user_metadata_manager.py
+++ b/generalresearch/managers/thl/user_manager/user_metadata_manager.py
@@ -1,4 +1,6 @@
-from typing import Collection, List, Optional
+from __future__ import annotations
+
+from collections.abc import Collection
from generalresearch.managers.base import PostgresManager
from generalresearch.models.thl.user_profile import UserMetadata
@@ -7,12 +9,12 @@ from generalresearch.models.thl.user_profile import UserMetadata
class UserMetadataManager(PostgresManager):
def filter(
self,
- user_ids: Optional[Collection[int]] = None,
- email_addresses: Optional[Collection[str]] = None,
- email_sha256s: Optional[Collection[str]] = None,
- email_sha1s: Optional[Collection[str]] = None,
- email_md5s: Optional[Collection[str]] = None,
- ) -> List[UserMetadata]:
+ user_ids: Collection[int] | None = None,
+ email_addresses: Collection[str] | None = None,
+ email_sha256s: Collection[str] | None = None,
+ email_sha1s: Collection[str] | None = None,
+ email_md5s: Collection[str] | None = None,
+ ) -> list[UserMetadata]:
for arg in [
user_ids,
email_addresses,
@@ -57,12 +59,12 @@ class UserMetadataManager(PostgresManager):
def get_if_exists(
self,
- user_id: Optional[int] = None,
- email_address: Optional[str] = None,
- email_sha256: Optional[str] = None,
- email_sha1: Optional[str] = None,
- email_md5: Optional[str] = None,
- ) -> Optional[UserMetadata]:
+ user_id: int | None = None,
+ email_address: str | None = None,
+ email_sha256: str | None = None,
+ email_sha1: str | None = None,
+ email_md5: str | None = None,
+ ) -> UserMetadata | None:
filters = {
"user_ids": user_id,
"email_addresses": email_address,
@@ -81,11 +83,11 @@ class UserMetadataManager(PostgresManager):
def get(
self,
- user_id: Optional[int] = None,
- email_address: Optional[str] = None,
- email_sha256: Optional[str] = None,
- email_sha1: Optional[str] = None,
- email_md5: Optional[str] = None,
+ user_id: int | None = None,
+ email_address: str | None = None,
+ email_sha256: str | None = None,
+ email_sha1: str | None = None,
+ email_md5: str | None = None,
) -> UserMetadata:
res = self.get_if_exists(
user_id=user_id,
diff --git a/generalresearch/managers/thl/user_streak.py b/generalresearch/managers/thl/user_streak.py
index 4bc7e70..e5d17cf 100644
--- a/generalresearch/managers/thl/user_streak.py
+++ b/generalresearch/managers/thl/user_streak.py
@@ -1,5 +1,6 @@
+from __future__ import annotations
+
from datetime import date, datetime
-from typing import List, Optional, Tuple
import pandas as pd
@@ -50,8 +51,8 @@ class UserStreakManager(PostgresManager):
return self.pg_config.execute_sql_query(query, params)
def get_user_streaks(
- self, user_id: int, country_iso: Optional[str] = None
- ) -> List[UserStreak]:
+ self, user_id: int, country_iso: str | None = None
+ ) -> list[UserStreak]:
country_iso = country_iso or self.get_user_country(user_id=user_id)
if country_iso is None:
return []
@@ -60,7 +61,7 @@ class UserStreakManager(PostgresManager):
active_days = [x["d"] for x in res]
complete_days = [x["d"] for x in res if x["is_complete"]]
- streaks: List[UserStreak] = []
+ streaks: list[UserStreak] = []
for period in StreakPeriod:
for fulfillment, days in [
@@ -92,11 +93,11 @@ class UserStreakManager(PostgresManager):
def compute_streaks_from_days(
- days: List[date],
+ days: list[date],
country_iso: str,
period: StreakPeriod,
- today: Optional[date] = None,
-) -> Tuple[int, int, StreakState, Optional[date]]:
+ today: date | None = None,
+) -> tuple[int, int, StreakState, date | None]:
"""
:returns: (current_streak, longest_streak, streak_state, last_period_start)
"""
diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py
index dab35d1..5924938 100644
--- a/generalresearch/managers/thl/userhealth.py
+++ b/generalresearch/managers/thl/userhealth.py
@@ -1,9 +1,12 @@
+from __future__ import annotations
+
import ipaddress
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
from itertools import zip_longest
from random import choice as rchoice
from random import random
-from typing import Any, Collection, Dict, List, Optional, Tuple
+from typing import Any
import faker
from pydantic import NonNegativeInt, PositiveInt
@@ -35,8 +38,8 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
self,
pg_config: PostgresConfig,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
):
super().__init__(
pg_config=pg_config,
@@ -53,7 +56,7 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
def get_redis_key(self, user_id: int) -> str:
return f"py-utils:user-ip-history:{user_id}"
- def get_user_ip_records_sql(self, user_id: int) -> List[UserIPRecord]:
+ def get_user_ip_records_sql(self, user_id: int) -> list[UserIPRecord]:
# The IP metadata is ONLY for the 'ip', NOT for any forwarded ips.
# This might get called immediately after a write, so use the non-rr
res = self.pg_config.execute_sql_query(
@@ -78,7 +81,7 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
res = [UserIPRecord.model_validate(x) for x in res]
return res
- def get_user_ip_history_cache(self, user_id: int) -> Optional[UserIPHistory]:
+ def get_user_ip_history_cache(self, user_id: int) -> UserIPHistory | None:
res = self.redis_client.get(self.get_redis_key(user_id))
if res:
return UserIPHistory.model_validate_json(res)
@@ -86,12 +89,10 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
def delete_user_ip_history_cache(self, user_id: int) -> None:
self.redis_client.delete(self.get_redis_key(user_id))
- return None
def set_user_ip_history_cache(self, user_id: int, iph: UserIPHistory) -> None:
value = iph.model_dump_json()
self.redis_client.set(self.get_redis_key(user_id), value, ex=3 * 24 * 3600)
- return None
def recreate_user_ip_history_cache(self, user_id: int) -> None:
self.delete_user_ip_history_cache(user_id=user_id)
@@ -99,7 +100,6 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
# todo: we may get dns records from somewhere else here ...
iph = UserIPHistory(user_id=user_id, ips=records)
self.set_user_ip_history_cache(user_id=user_id, iph=iph)
- return None
def get_user_ip_history(self, user_id: int) -> UserIPHistory:
assert isinstance(user_id, int)
@@ -117,9 +117,7 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
iph.enrich_ips(pg_config=self.pg_config, redis_config=self.redis_config)
return iph
- def get_user_latest_ip(
- self, user: User, exclude_anon: bool = False
- ) -> Optional[str]:
+ def get_user_latest_ip(self, user: User, exclude_anon: bool = False) -> str | None:
record = self.get_user_latest_ip_record(user=user, exclude_anon=exclude_anon)
if record:
return record.ip
@@ -127,7 +125,7 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
def get_user_latest_ip_record(
self, user: User, exclude_anon: bool = False
- ) -> Optional[UserIPRecord]:
+ ) -> UserIPRecord | None:
iphistory = self.get_user_ip_history(user_id=user.user_id)
if iphistory.ips:
@@ -146,14 +144,14 @@ class UserIpHistoryManager(PostgresManagerWithRedis):
def get_user_latest_country(
self, user: User, exclude_anon: bool = False
- ) -> Optional[str]:
+ ) -> str | None:
"""Get the country the user is in, based off their latest ip."""
ipr = self.get_user_latest_ip_record(user, exclude_anon=exclude_anon)
# The ipr.information should exist, but it is possible the user has
# no IP history at all, so the record is None
return ipr.country_iso if ipr is not None else None
- def is_user_anonymous(self, user: User) -> Optional[bool]:
+ def is_user_anonymous(self, user: User) -> bool | None:
# Get the user's latest ip. is it marked as anonymous?
# Note: it is possible we only did a "basic" lookup of this IP so
# we don't know if they are anonymous. Default to False
@@ -170,8 +168,8 @@ class IPRecordManager(PostgresManagerWithRedis):
self,
pg_config: PostgresConfig,
redis_config: RedisConfig,
- permissions: Collection[Permission] = None,
- cache_prefix: Optional[str] = None,
+ permissions: Collection[Permission] | None = None,
+ cache_prefix: str | None = None,
):
super().__init__(
pg_config=pg_config,
@@ -189,13 +187,13 @@ class IPRecordManager(PostgresManagerWithRedis):
def create_dummy(
self,
user_id: PositiveInt,
- ip: Optional[IPvAnyAddressStr] = None,
- forwarded_ip1: Optional[IPvAnyAddressStr] = None,
- forwarded_ip2: Optional[IPvAnyAddressStr] = None,
- forwarded_ip3: Optional[IPvAnyAddressStr] = None,
- forwarded_ip4: Optional[IPvAnyAddressStr] = None,
- forwarded_ip5: Optional[IPvAnyAddressStr] = None,
- forwarded_ip6: Optional[IPvAnyAddressStr] = None,
+ 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 self.create(
user_id=user_id,
@@ -214,7 +212,7 @@ class IPRecordManager(PostgresManagerWithRedis):
self,
user_id: PositiveInt,
ip: IPvAnyAddressStr,
- forwarded_ips: List[str],
+ forwarded_ips: list[str],
) -> IPRecord:
if len(forwarded_ips) > 6:
raise ValueError("A maximum of 6 forwarded IPs is allowed.")
@@ -282,7 +280,7 @@ class IPRecordManager(PostgresManagerWithRedis):
return IPRecord.from_mysql(data)
- def get_user_latest_ip_record(self, user: User) -> Optional[IPRecord]:
+ def get_user_latest_ip_record(self, user: User) -> IPRecord | None:
res = self.filter_ip_records(user_ids=[user.user_id], limit=1)
if res:
return res[0]
@@ -290,11 +288,11 @@ class IPRecordManager(PostgresManagerWithRedis):
def filter_ip_records(
self,
- filter_ips: Optional[List[IPvAnyAddressStr]] = None,
- user_ids: Optional[List[PositiveInt]] = None,
- created_from: Optional[datetime] = None,
- limit: Optional[int] = None,
- ) -> List[IPRecord]:
+ filter_ips: list[IPvAnyAddressStr] | None = None,
+ user_ids: list[PositiveInt] | None = None,
+ created_from: datetime | None = None,
+ limit: int | None = None,
+ ) -> list[IPRecord]:
assert any([filter_ips, user_ids, created_from]), "Must provide filter criteria"
@@ -353,10 +351,10 @@ class AuditLogManager(PostgresManager):
def create_dummy(
self,
user_id: PositiveInt,
- level: Optional[AuditLogLevel] = None,
- event_type: Optional[str] = None,
- event_msg: Optional[str] = None,
- event_value: Optional[float] = None,
+ level: AuditLogLevel | None = None,
+ event_type: str | None = None,
+ event_msg: str | None = None,
+ event_value: float | None = None,
) -> AuditLog:
event_types = {
@@ -378,8 +376,8 @@ class AuditLogManager(PostgresManager):
user_id: PositiveInt,
level: AuditLogLevel,
event_type: str,
- event_msg: Optional[str] = None,
- event_value: Optional[float] = None,
+ event_msg: str | None = None,
+ event_value: float | None = None,
) -> AuditLog:
"""AuditLogs may exist with the same event_type, and with different levels"""
@@ -433,7 +431,7 @@ class AuditLogManager(PostgresManager):
return AuditLog.from_mysql(res[0])
- def filter_by_product(self, product: Product) -> List[AuditLog]:
+ def filter_by_product(self, product: Product) -> list[AuditLog]:
res = self.pg_config.execute_sql_query(
query="""
@@ -450,7 +448,7 @@ class AuditLogManager(PostgresManager):
return [AuditLog.from_mysql(i) for i in res]
- def filter_by_user_id(self, user_id: PositiveInt) -> List[AuditLog]:
+ def filter_by_user_id(self, user_id: PositiveInt) -> list[AuditLog]:
res = self.pg_config.execute_sql_query(
query="""
SELECT *
@@ -467,13 +465,13 @@ class AuditLogManager(PostgresManager):
def filter(
self,
user_ids: Collection[int],
- level: Optional[int] = None,
- level_ge: Optional[int] = None,
- event_type: Optional[str] = None,
- event_type_like: Optional[str] = None,
- event_msg: Optional[str] = None,
- created_after: Optional[datetime] = None,
- ) -> List[AuditLog]:
+ level: int | None = None,
+ level_ge: int | None = None,
+ event_type: str | None = None,
+ event_type_like: str | None = None,
+ event_msg: str | None = None,
+ created_after: datetime | None = None,
+ ) -> list[AuditLog]:
filter_str, args = self.make_filter_str(
user_ids=user_ids,
@@ -500,12 +498,12 @@ class AuditLogManager(PostgresManager):
def filter_count(
self,
user_ids: Collection[int],
- level: Optional[int] = None,
- level_ge: Optional[int] = None,
- event_type: Optional[str] = None,
- event_type_like: Optional[str] = None,
- event_msg: Optional[str] = None,
- created_after: Optional[datetime] = None,
+ level: int | None = None,
+ level_ge: int | None = None,
+ event_type: str | None = None,
+ event_type_like: str | None = None,
+ event_msg: str | None = None,
+ created_after: datetime | None = None,
) -> NonNegativeInt:
filter_str, args = self.make_filter_str(
@@ -534,13 +532,13 @@ class AuditLogManager(PostgresManager):
@staticmethod
def make_filter_str(
user_ids: Collection[int],
- level: Optional[int] = None,
- level_ge: Optional[int] = None,
- event_type: Optional[str] = None,
- event_type_like: Optional[str] = None,
- event_msg: Optional[str] = None,
- created_after: Optional[datetime] = None,
- ) -> Tuple[str, Dict[str, Any]]:
+ level: int | None = None,
+ level_ge: int | None = None,
+ event_type: str | None = None,
+ event_type_like: str | None = None,
+ event_msg: str | Nond = None,
+ 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]
diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py
index 3ddbf51..de7b599 100644
--- a/generalresearch/managers/thl/wall.py
+++ b/generalresearch/managers/thl/wall.py
@@ -1,10 +1,12 @@
+from __future__ import annotations
+
import logging
from collections import defaultdict
+from collections.abc import Collection
from datetime import datetime, timedelta, timezone
from decimal import ROUND_DOWN, Decimal
from functools import cached_property
from random import choice as rchoice
-from typing import Collection, List, Optional
from uuid import uuid4
from faker import Faker
@@ -44,7 +46,7 @@ class WallManager(PostgresManager):
def __init__(
self,
pg_config: PostgresConfig,
- permissions: Optional[Collection[Permission]] = None,
+ permissions: Collection[Permission] | None = None,
):
assert pg_config.row_factory == dict_row
super().__init__(pg_config=pg_config, permissions=permissions)
@@ -57,8 +59,8 @@ class WallManager(PostgresManager):
source: Source,
req_survey_id: str,
req_cpi: Decimal,
- buyer_id: Optional[str] = None,
- uuid_id: Optional[str] = None,
+ buyer_id: str | None = None,
+ uuid_id: str | None = None,
) -> Wall:
"""
Creates a Wall event. Prefer to use this rather than instantiating
@@ -94,14 +96,14 @@ class WallManager(PostgresManager):
def create_dummy(
self,
- session_id: Optional[int] = None,
- user_id: Optional[int] = None,
- started: Optional[datetime] = None,
- source: Optional[Source] = None,
- req_survey_id: Optional[str] = None,
- req_cpi: Optional[Decimal] = None,
- buyer_id: Optional[str] = None,
- uuid_id: Optional[str] = None,
+ 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"""
@@ -158,7 +160,7 @@ class WallManager(PostgresManager):
assert len(res) == 1, f"Expected 1 result, got {len(res)}"
return Wall.model_validate(res[0])
- def get_from_uuid_if_exists(self, wall_uuid: UUIDStr) -> Optional[Wall]:
+ def get_from_uuid_if_exists(self, wall_uuid: UUIDStr) -> Wall | None:
try:
return self.get_from_uuid(wall_uuid=wall_uuid)
except AssertionError:
@@ -170,12 +172,12 @@ class WallManager(PostgresManager):
status: Status,
status_code_1: StatusCode1,
finished: datetime,
- ext_status_code_1: Optional[str] = None,
- ext_status_code_2: Optional[str] = None,
- ext_status_code_3: Optional[str] = None,
- status_code_2: Optional[WallStatusCode2] = None,
- survey_id: Optional[str] = None,
- cpi: Optional[Decimal] = None,
+ ext_status_code_1: str | None = None,
+ ext_status_code_2: str | None = None,
+ ext_status_code_3: str | None = None,
+ status_code_2: WallStatusCode2 | None = None,
+ survey_id: str | None = None,
+ cpi: Decimal | None = None,
) -> None:
"""This wall event is finished. This would be called if/when we get a
callback for this wall event. Some other code is responsible for
@@ -232,10 +234,10 @@ class WallManager(PostgresManager):
def get_wall_events(
self,
- session_id: Optional[PositiveInt] = None,
- session_ids: Optional[List[PositiveInt]] = None,
+ session_id: PositiveInt | None = None,
+ session_ids: list[PositiveInt] | None = None,
order_by: OrderBy = OrderBy.ASC,
- ) -> List[Wall]:
+ ) -> list[Wall]:
if session_id is not None and session_ids is not None:
raise ValueError("Cannot provide both session_id and session_ids")
@@ -271,8 +273,8 @@ class WallManager(PostgresManager):
self,
wall: Wall,
adjusted_timestamp: AwareDatetime,
- adjusted_status: Optional[WallAdjustedStatus] = None,
- adjusted_cpi: Optional[Decimal] = None,
+ adjusted_status: WallAdjustedStatus | None = None,
+ adjusted_cpi: Decimal | None = None,
) -> None:
assert wall.status, "Wall must have an existing Status"
@@ -317,15 +319,13 @@ class WallManager(PostgresManager):
"uuid": wall.uuid,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE thl_wall
SET adjusted_status = %(adjusted_status)s,
adjusted_timestamp = %(adjusted_timestamp)s,
adjusted_cpi = %(adjusted_cpi)s
WHERE uuid = %(uuid)s;
- """
- )
+ """)
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
@@ -333,14 +333,12 @@ class WallManager(PostgresManager):
assert c.rowcount == 1
conn.commit()
- return None
-
def report(
self,
wall: Wall,
report_value: ReportValue,
- report_notes: Optional[str] = None,
- report_timestamp: Optional[AwareDatetime] = None,
+ report_notes: str | None = None,
+ report_timestamp: AwareDatetime | None = None,
) -> None:
wall.report(
report_value=report_value,
@@ -354,16 +352,14 @@ class WallManager(PostgresManager):
"finished": wall.finished,
"report_notes": report_notes,
}
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE thl_wall
SET report_value = %(report_value)s,
report_notes = %(report_notes)s,
status = %(status)s,
finished = %(finished)s
WHERE uuid = %(uuid)s;
- """
- )
+ """)
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=params)
@@ -400,12 +396,12 @@ class WallManager(PostgresManager):
def filter_wall_attempts_paginated(
self,
user_id: int,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
page: int = 1,
size: int = 100,
- order_by: Optional[str] = "-started",
- ) -> List[WallAttempt]:
+ order_by: str | None = "-started",
+ ) -> list[WallAttempt]:
"""
Returns WallAttempt
"""
@@ -461,10 +457,10 @@ class WallManager(PostgresManager):
def filter_wall_attempts(
self,
user_id: int,
- started_after: Optional[datetime] = None,
- started_before: Optional[datetime] = None,
- order_by: Optional[str] = "-started",
- ) -> List[WallAttempt]:
+ started_after: datetime | None = None,
+ started_before: datetime | None = None,
+ order_by: str | None = "-started",
+ ) -> list[WallAttempt]:
started_before = started_before or datetime.now(tz=timezone.utc)
res = []
page = 1
@@ -485,8 +481,8 @@ class WallManager(PostgresManager):
return res
def get_survey_activities(
- self, survey_keys: Collection[SurveyKey], product_id: Optional[str] = None
- ) -> List[TaskActivity]:
+ self, survey_keys: Collection[SurveyKey], product_id: str | None = None
+ ) -> list[TaskActivity]:
query_base = """
row_stats AS (
SELECT
@@ -617,16 +613,16 @@ class WallCacheManager(PostgresManagerWithRedis):
assert type(user_id) is int, "user_id must be int"
self.redis_client.delete(self.get_flag_key_(user_id=user_id))
- def get_attempts_redis_(self, user_id: int) -> List[WallAttempt]:
+ def get_attempts_redis_(self, user_id: int) -> list[WallAttempt]:
redis_key = self.get_cache_key_(user_id=user_id)
# Returns a list even if there is nothing set
res = self.redis_client.lrange(redis_key, 0, 5000)
attempts = [WallAttempt.model_validate_json(x) for x in res]
return attempts
- def update_attempts_redis_(self, attempts: List[WallAttempt], user_id: int) -> None:
+ def update_attempts_redis_(self, attempts: list[WallAttempt], user_id: int) -> None:
if not attempts:
- return None
+ return
redis_key = self.get_cache_key_(user_id=user_id)
# Make sure attempts is ordered, so the most recent is last
@@ -639,9 +635,8 @@ class WallCacheManager(PostgresManagerWithRedis):
# So this doesn't grow forever, keep only the most recent 5k
self.redis_client.ltrim(redis_key, 0, 4999)
- return None
- def get_attempts(self, user_id: PositiveInt) -> List[WallAttempt]:
+ def get_attempts(self, user_id: PositiveInt) -> list[WallAttempt]:
"""
This is used in the GetOpportunityIDs call to get a list of surveys
(& surveygroups) which should be excluded for this user. We don't
@@ -658,7 +653,7 @@ class WallCacheManager(PostgresManagerWithRedis):
# Attempt to get the most recent wall attempt
redis_key = self.get_cache_key_(user_id=user_id)
- res: Optional[str] = self.redis_client.lindex(redis_key, 0) # type: ignore[assignment]
+ res: str | None = self.redis_client.lindex(redis_key, 0) # type: ignore[assignment]
if res is None:
# Nothing in the cache, query for all from db
attempts = self.wall_manager.filter_wall_attempts(user_id=user_id)
diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py
index 6e57d56..ebe839a 100644
--- a/generalresearch/models/admin/__init__.py
+++ b/generalresearch/models/admin/__init__.py
@@ -1,11 +1,12 @@
+from __future__ import annotations
+
from datetime import datetime, timezone
-from typing import Optional
import pandas as pd
from dateutil import relativedelta
-def get_date_list(start_datetime: datetime, end_datetime: Optional[datetime] = None):
+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)
return (
diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py
index 30e9fcd..67bd263 100644
--- a/generalresearch/models/admin/request.py
+++ b/generalresearch/models/admin/request.py
@@ -1,9 +1,11 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import datetime, timedelta, timezone
from enum import Enum
-from typing import Literal, List, Tuple
+from typing import Literal
import pandas as pd
-from pydantic import BaseModel, Field, model_validator, computed_field
+from pydantic import BaseModel, Field, computed_field, model_validator
from generalresearch.models.custom_types import AwareDatetimeISO
@@ -33,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)],
+ examples=[datetime(year=2025, month=5, day=1, tzinfo=timezone.utc)],
return_type=datetime,
)
@property
@@ -151,7 +153,7 @@ class ReportRequest(BaseModel):
tz=timezone.utc,
)
- def bucket_ranges(self) -> List[Tuple[pd.Timestamp, pd.Timestamp]]:
+ def bucket_ranges(self) -> list[tuple[pd.Timestamp, pd.Timestamp]]:
"""Returns list of (start, end) tuples for each bucket."""
starts = self.buckets()
return [(s, s + self.interval_timedelta) for s in starts]
diff --git a/generalresearch/models/cint/__init__.py b/generalresearch/models/cint/__init__.py
index ea3e500..2c1be7e 100644
--- a/generalresearch/models/cint/__init__.py
+++ b/generalresearch/models/cint/__init__.py
@@ -1,5 +1,4 @@
from pydantic import Field
-
from typing_extensions import Annotated
CintQuestionIdType = Annotated[
diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py
index 77212a2..b870246 100644
--- a/generalresearch/models/cint/question.py
+++ b/generalresearch/models/cint/question.py
@@ -1,7 +1,9 @@
+from __future__ import annotations
+
import json
from datetime import datetime, timezone
from enum import Enum
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
+from typing import TYPE_CHECKING, Any, Literal
from uuid import UUID
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -48,7 +50,7 @@ class CintQuestionType(str, Enum):
class CintUserQuestionAnswer(MarketplaceUserQuestionAnswer):
question_id: CintQuestionIdType = Field()
- question_type: Optional[CintQuestionType] = Field(default=None)
+ question_type: CintQuestionType | None = Field(default=None)
# Did this answer come from us asking, or was it passed back from the marketplace
from_thl: bool = Field(default=True)
@@ -89,11 +91,11 @@ class CintQuestion(MarketplaceQuestion):
frozen=True,
examples=[CintQuestionType.MULTI_SELECT],
)
- options: Optional[List[CintQuestionOption]] = Field(
+ options: list[CintQuestionOption] | None = Field(
default=None, min_length=1, frozen=True
)
option_mask: str = Field(examples=["000000000000000000"])
- classification_code: Optional[str] = Field(examples=["ELE"], default=None)
+ classification_code: str | None = Field(examples=["ELE"], default=None)
# This comes from the API! not us
created_at: AwareDatetimeISO = Field(description="Called create_date in API")
@@ -104,7 +106,7 @@ class CintQuestion(MarketplaceQuestion):
return self.question_id
@field_validator("question_name", "question_text", mode="after")
- def remove_nbsp(cls, s: Optional[str]) -> Optional[str]:
+ def remove_nbsp(cls, s: str | None) -> str | None:
return string_utils.remove_nbsp(s)
@model_validator(mode="after")
@@ -124,8 +126,8 @@ class CintQuestion(MarketplaceQuestion):
@field_validator("options")
@classmethod
def order_options(
- cls, options: Optional[List[CintQuestionOption]]
- ) -> Optional[List[CintQuestionOption]]:
+ cls, options: list[CintQuestionOption] | None
+ ) -> list[CintQuestionOption] | None:
if options:
options.sort(key=lambda x: x.order)
@@ -134,8 +136,8 @@ class CintQuestion(MarketplaceQuestion):
@field_validator("options")
@classmethod
def validate_options(
- cls, options: Optional[List[CintQuestionOption]]
- ) -> Optional[List[CintQuestionOption]]:
+ cls, options: list[CintQuestionOption] | None
+ ) -> list[CintQuestionOption] | None:
if options:
ids = {x.id for x in options}
assert len(ids) == len(options), "options.id must be unique"
@@ -145,7 +147,7 @@ class CintQuestion(MarketplaceQuestion):
return options
@classmethod
- def from_api(cls, d: Dict[str, Any], country_iso: str, language_iso: str) -> Self:
+ def from_api(cls, d: dict[str, Any], country_iso: str, language_iso: str) -> Self:
options = None
created_at = datetime.strptime(
d["create_date"], "%Y-%m-%dT%H:%M:%S%z"
@@ -178,7 +180,7 @@ class CintQuestion(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 = [
@@ -203,7 +205,7 @@ class CintQuestion(MarketplaceQuestion):
category_id=UUID(d["category_id"]).hex if d["category_id"] else None,
)
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(mode="json")
d["options"] = json.dumps(d["options"])
if self.created_at:
diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py
index be07e04..56384e3 100644
--- a/generalresearch/models/cint/survey.py
+++ b/generalresearch/models/cint/survey.py
@@ -4,7 +4,7 @@ import json
import logging
from datetime import datetime, timezone
from decimal import Decimal
-from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Type
+from typing import Any, Literal, Type
from more_itertools import flatten
from pydantic import (
@@ -46,7 +46,7 @@ class CintCondition(MarketplaceCondition):
question_id: CintQuestionIdType = Field()
@classmethod
- def from_api(cls, d: Dict[str, Any]) -> Self:
+ def from_api(cls, d: dict[str, Any]) -> Self:
d["question_id"] = str(d["question_id"])
d["values"] = list(map(str.lower, d["precodes"]))
d["value_type"] = ConditionValueType.LIST
@@ -61,11 +61,11 @@ class CintQuota(BaseModel):
model_config = ConfigDict(populate_by_name=True, frozen=True)
quota_id: CoercedStr = Field(validation_alias="survey_quota_id")
quota_type: Literal["total", "client"] = Field(validation_alias="survey_quota_type")
- conversion: Optional[float] = Field(ge=0, le=1, default=None)
+ conversion: float | None = Field(ge=0, le=1, default=None)
number_of_respondents: NonNegativeInt = Field(
description="Number of completes available"
)
- condition_hashes: Optional[List[str]] = Field(min_length=1, default=None)
+ condition_hashes: list[str] | None = Field(min_length=1, default=None)
def __hash__(self):
return hash(tuple((tuple(self.condition_hashes), self.quota_id)))
@@ -85,23 +85,23 @@ class CintQuota(BaseModel):
return self.number_of_respondents >= 2
@classmethod
- def from_api(cls, d: Dict) -> Self:
+ def from_api(cls, d: dict) -> Self:
d["survey_quota_type"] = d["survey_quota_type"].lower()
return cls.model_validate(d)
- 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)
- 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.
# We can "match" a quota that is closed. In that case, we would
# not be eligible for the survey.
return all(criteria_evaluation.get(c) for c in self.condition_hashes)
def matches_optional(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Optional[bool]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> bool | None:
# We need to know if any conditions are unknown to avoid matching a full quota. If any fail,
# then we know we fail regardless of any being unknown.
evals = [criteria_evaluation.get(c) for c in self.condition_hashes]
@@ -112,8 +112,8 @@ class CintQuota(BaseModel):
return True
def matches_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]]:
# Passes back "matches" (T/F/none) and a list of unknown criterion hashes
hash_evals = {
cell: criteria_evaluation.get(cell) for cell in self.condition_hashes
@@ -138,12 +138,12 @@ class CintSurvey(MarketplaceTask):
is_live_raw: bool = Field(alias="is_live")
- bid_loi: Optional[int] = Field(
+ bid_loi: int | None = Field(
ge=60, le=90 * 60, validation_alias="bid_length_of_interview"
)
- bid_ir: Optional[float] = Field(ge=0, le=1, validation_alias="bid_incidence")
- collects_pii: Optional[bool] = Field()
- survey_group_ids: Set[CoercedStr] = Field()
+ bid_ir: float | None = Field(ge=0, le=1, validation_alias="bid_incidence")
+ collects_pii: bool | None = Field()
+ survey_group_ids: set[CoercedStr] = Field()
calculation_type: TaskCalculationType = Field(
description="Indicates whether quotas are calculated based on completes or prescreens",
@@ -186,50 +186,50 @@ class CintSurvey(MarketplaceTask):
description="Number of completes still available to the supplier"
)
completion_percentage: float = Field()
- conversion: Optional[float] = Field(
+ conversion: float | None = Field(
ge=0,
le=1,
description="Percentage of respondents who complete the survey after qualifying",
)
- mobile_conversion: Optional[float] = Field(
+ mobile_conversion: float | None = Field(
ge=0,
le=1,
description="Percentage of respondents on a mobile device who complete the survey after qualifying.",
)
- length_of_interview: Optional[NonNegativeInt] = Field(
+ length_of_interview: NonNegativeInt | None = Field(
description="Median time for a respondent to complete the survey, excluding prescreener, in minutes. This "
"value will be zero until 6 completes are achieved."
)
overall_completes: NonNegativeInt = Field(
description="Number of completes already achieved across all suppliers on the survey."
)
- revenue_per_click: Optional[float] = Field(
+ revenue_per_click: float | None = Field(
description="The Revenue Per Click value of the survey. RPC = (RPI * completes) / system entrants",
default=None,
)
- termination_length_of_interview: Optional[NonNegativeInt] = Field(
+ termination_length_of_interview: NonNegativeInt | None = Field(
description="Median time for a respondent to be termed, in minutes. This value is calculated after six survey "
"entrants and rounded to the nearest whole number. Until six survey entrants are achieved the "
"value will be zero."
)
- respondent_pids: Set[str] = Field(default_factory=set)
+ respondent_pids: set[str] = Field(default_factory=set)
- qualifications: List[str] = Field(default_factory=list)
- quotas: List[CintQuota] = Field(default_factory=list)
+ qualifications: list[str] = Field(default_factory=list)
+ quotas: list[CintQuota] = Field(default_factory=list)
source: Literal[Source.CINT] = Field(default=Source.CINT)
- used_question_ids: Set[AlphaNumStr] = Field(default_factory=set)
+ used_question_ids: set[AlphaNumStr] = 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, CintCondition]] = Field(default=None)
+ conditions: dict[str, CintCondition] | None = Field(default=None)
# These do not come from the API. We set it when we update/create in the db.
- created_at: Optional[AwareDatetimeISO] = Field(default=None)
- last_updated: Optional[AwareDatetimeISO] = Field(default=None)
+ created_at: AwareDatetimeISO | None = Field(default=None)
+ last_updated: AwareDatetimeISO | None = Field(default=None)
@property
def internal_id(self) -> str:
@@ -246,14 +246,14 @@ class CintSurvey(MarketplaceTask):
def is_live(self) -> bool:
return self.is_live_raw
- def model_dump(self, **kwargs: Any) -> Dict[str, Any]:
+ def model_dump(self, **kwargs: Any) -> dict[str, Any]:
data = super().model_dump(**kwargs)
data["is_live"] = data.pop("is_live_raw", None)
return data
@computed_field
@property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
s = set(self.qualifications)
for q in self.quotas:
s.update(set(q.condition_hashes)) if q.condition_hashes else None
@@ -299,7 +299,7 @@ class CintSurvey(MarketplaceTask):
return "42"
@property
- def marketplace_genders(self) -> Dict[Gender, Optional[MarketplaceCondition]]:
+ def marketplace_genders(self) -> dict[Gender, MarketplaceCondition | None]:
return {
Gender.MALE: CintCondition(
question_id="43",
@@ -315,7 +315,7 @@ class CintSurvey(MarketplaceTask):
}
@classmethod
- def from_api(cls, d: Dict[str, Any]) -> Optional[Self]:
+ def from_api(cls, d: dict[str, Any]) -> Self | None:
try:
return cls._from_api(d)
except Exception as e:
@@ -323,7 +323,7 @@ class CintSurvey(MarketplaceTask):
return None
@classmethod
- def _from_api(cls, d: Dict[str, Any]) -> Self:
+ def _from_api(cls, d: dict[str, Any]) -> Self:
if "cpi" in d:
d["gross_cpi"] = Decimal(d.pop("cpi"))
if "revenue_per_interview" in d:
@@ -396,7 +396,7 @@ class CintSurvey(MarketplaceTask):
return cls.model_validate(d)
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(
mode="json",
exclude={
@@ -428,14 +428,14 @@ class CintSurvey(MarketplaceTask):
return cls.model_validate(d)
def passes_qualifications(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
# We have to match all quals
return all(criteria_evaluation.get(q) for q in self.qualifications)
def passes_qualifications_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]]:
# Passes back "passes" (T/F/none) and a list of unknown criterion hashes
hash_evals = {q: criteria_evaluation.get(q) for q in self.qualifications}
evals = set(hash_evals.values())
@@ -447,7 +447,7 @@ class CintSurvey(MarketplaceTask):
return None, {cell for cell, ev in hash_evals.items() if ev is None}
return True, set()
- def passes_quotas(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Many surveys have 0 quotas. Quotas are exclusionary.
# They can NOT match a quota where currently_open=0
any_pass = True
@@ -462,8 +462,8 @@ class CintSurvey(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]]:
# 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]
@@ -506,7 +506,7 @@ class CintSurvey(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
@@ -515,8 +515,8 @@ class CintSurvey(MarketplaceTask):
)
def determine_eligibility_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]]:
# We check is_open when putting the survey in global_vars. Don't need to check again.
# if self.is_open is False:
# return False, set()
diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py
index b43250a..afd9fc4 100644
--- a/generalresearch/models/cint/task_collection.py
+++ b/generalresearch/models/cint/task_collection.py
@@ -1,4 +1,6 @@
-from typing import List, Set
+from __future__ import annotations
+
+from typing import List
import pandas as pd
from pandera import Check, Column, DataFrameSchema, Index
@@ -10,8 +12,8 @@ from generalresearch.models.thl.survey.task_collection import (
create_empty_df_from_schema,
)
-COUNTRY_ISOS: Set[str] = Localelator().get_all_countries()
-LANGUAGE_ISOS: Set[str] = Localelator().get_all_languages()
+COUNTRY_ISOS: set[str] = Localelator().get_all_countries()
+LANGUAGE_ISOS: set[str] = Localelator().get_all_languages()
CintTaskCollectionSchema = DataFrameSchema(
columns={
@@ -50,7 +52,7 @@ CintTaskCollectionSchema = DataFrameSchema(
class CintTaskCollection(TaskCollection):
- items: List[CintSurvey]
+ items: list[CintSurvey]
_schema = CintTaskCollectionSchema
def to_row(self, s: CintSurvey):
diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py
index 4d675c0..7288588 100644
--- a/generalresearch/models/dynata/question.py
+++ b/generalresearch/models/dynata/question.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
# https://developers.dynata.com/docs/rex-respondent-gateway/dc5b33f20a1c9-get-attribute-info
import json
import logging
@@ -5,7 +7,7 @@ import re
from datetime import timedelta
from enum import Enum
from functools import cached_property
-from typing import Any, Dict, List, Literal, Optional, Set
+from typing import Any, Literal
from pydantic import BaseModel, Field, PositiveInt, field_validator, model_validator
@@ -87,11 +89,11 @@ class DynataUserQuestionAnswer(BaseModel):
# This is optional b/c this model can be used for eligibility checks for "anonymous" users, which are represented
# by a list of question answers not associated with an actual user. No default b/c we must explicitly set
# the field to None.
- user_id: Optional[PositiveInt] = Field(lt=MAX_INT32)
+ user_id: PositiveInt | None = Field(lt=MAX_INT32)
question_id: str = Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
# This is optional b/c we do not need it when writing these to the db. When these are fetched from the db
# for use in yield-management, we read this field from the question table.
- question_type: Optional[DynataQuestionType] = Field(default=None)
+ question_type: DynataQuestionType | None = Field(default=None)
# 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()
@@ -105,10 +107,10 @@ class DynataUserQuestionAnswer(BaseModel):
)
@cached_property
- def options_ids(self) -> Set[str]:
+ def options_ids(self) -> set[str]:
return set(self.option_id.split("|"))
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(mode="json", exclude={"question_type"})
d["created"] = self.created.replace(tzinfo=None)
return d
@@ -118,7 +120,7 @@ class DynataQuestionDependency(BaseModel, frozen=True):
# This is not explained or documented. Going to just store it for now
question_id: str = Field(min_length=1, max_length=16, pattern=r"^[0-9]+$")
# Some are an empty list. Unclear if this means "any option" or it is broken.
- option_ids: List[str] = Field()
+ option_ids: list[str] = Field()
class DynataQuestion(MarketplaceQuestion):
@@ -139,10 +141,10 @@ class DynataQuestion(MarketplaceQuestion):
max_length=1024, min_length=1, description="The text shown to respondents"
)
question_type: DynataQuestionType = Field(frozen=True)
- options: Optional[List[DynataQuestionOption]] = Field(default=None, min_length=1)
+ options: list[DynataQuestionOption] | None = Field(default=None, min_length=1)
# This does not mean that it doesn't expire, it means undefined.
- expiration_duration: Optional[timedelta] = Field(default=None)
- parent_dependencies: List[DynataQuestionDependency] = Field(default_factory=list)
+ expiration_duration: timedelta | None = Field(default=None)
+ parent_dependencies: list[DynataQuestionDependency] = Field(default_factory=list)
source: Literal[Source.DYNATA] = Source.DYNATA
@@ -215,7 +217,7 @@ class DynataQuestion(MarketplaceQuestion):
expiration_duration=expiration_duration,
)
- 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"])
d["parent_dependencies"] = json.dumps(d["parent_dependencies"])
diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py
index ae12436..d65b55d 100644
--- a/generalresearch/models/dynata/survey.py
+++ b/generalresearch/models/dynata/survey.py
@@ -5,7 +5,7 @@ import logging
from datetime import timezone
from decimal import Decimal
from functools import cached_property
-from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Type
+from typing import Any, Literal, Type
from more_itertools import flatten
from pydantic import (
@@ -86,7 +86,7 @@ class DynataRequirements(BaseModel):
class DynataCondition(MarketplaceCondition):
- question_id: Optional[CoercedStr] = Field(
+ question_id: CoercedStr | None = Field(
min_length=1,
max_length=16,
pattern=r"^[0-9]+$",
@@ -95,10 +95,10 @@ class DynataCondition(MarketplaceCondition):
# This comes in the API and is used to match "cells" to quotas they're associated with. Once
# we parse the API response, we don't need this tag id anymore.
- tag: Optional[str] = Field(default=None, max_length=36)
+ tag: str | None = Field(default=None, max_length=36)
@classmethod
- def from_api(cls, cell: Dict[str, Any]) -> "DynataCondition":
+ def from_api(cls, cell: dict[str, Any]) -> "DynataCondition":
"""
We perform some preprocessing before calling this to pull in the data from COLLECTION cells.
"""
@@ -167,7 +167,7 @@ class DynataQuota(BaseModel):
count: int = Field(description="Limit of completes available")
# Each condition_hash is called in Dynata a "Quota Cell"
# Some quotas have no conditions. I'm not sure how eligibility is supposed to work for this.
- condition_hashes: List[str] = Field(min_length=0, default_factory=list)
+ condition_hashes: list[str] = Field(min_length=0, default_factory=list)
status: DynataStatus = Field()
def __hash__(self):
@@ -180,13 +180,13 @@ class DynataQuota(BaseModel):
min_open_spots = 3
return self.status == DynataStatus.OPEN and (self.count >= min_open_spots)
- def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# We have to match all conditions (aka cells) within the quota (aka quota object).
return self.is_open and all(
criteria_evaluation.get(c) for c in self.condition_hashes
)
- def passes_verbose(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_verbose(self, criteria_evaluation: dict[str, bool | None]) -> bool:
print(f"quota.is_open: {self.is_open}")
print(
", ".join(
@@ -198,8 +198,8 @@ class DynataQuota(BaseModel):
)
def passes_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]]:
# Passes back "passes" (T/F/none) and a list of unknown criterion hashes
if self.is_open is False:
return False, set()
@@ -218,7 +218,7 @@ class DynataQuota(BaseModel):
class DynataQuotaGroup(RootModel):
- root: List[DynataQuota] = Field()
+ root: list[DynataQuota] = Field()
def __iter__(self):
return iter(self.root)
@@ -226,11 +226,11 @@ class DynataQuotaGroup(RootModel):
def __hash__(self):
return hash(tuple(self.root))
- def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Qualify for ANY quota object within a quota group
return any(quota.passes(criteria_evaluation) for quota in self.root)
- def passes_verbose(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_verbose(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Qualify for ANY quota object within a quota group
for quota in self.root:
print("---")
@@ -243,8 +243,8 @@ class DynataQuotaGroup(RootModel):
return any(cell.is_open for cell in self.root)
def passes_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]]:
# 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())
@@ -262,7 +262,7 @@ class DynataQuotaGroup(RootModel):
class DynataFilterObject(RootModel):
- root: List[str] = Field() # list of criterion hashes
+ root: list[str] = Field() # list of criterion hashes
def __iter__(self):
return iter(self.root)
@@ -270,19 +270,19 @@ class DynataFilterObject(RootModel):
def __hash__(self):
return hash(tuple(self.root))
- def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# We have to match all cells within an object.
return all(criteria_evaluation.get(cell) for cell in self.root)
- def passes_verbose(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_verbose(self, criteria_evaluation: dict[str, bool | None]) -> bool:
for cell in self.root:
print(f"{cell}: {criteria_evaluation.get(cell)}")
# We have to match all cells within an object.
return all(criteria_evaluation.get(cell) for cell in self.root)
def passes_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]]:
# Passes back "passes" (T/F/none) and a list of unknown criterion hashes
cell_evals = {cell: criteria_evaluation.get(cell) for cell in self.root}
evals = set(cell_evals.values())
@@ -297,7 +297,7 @@ class DynataFilterObject(RootModel):
class DynataFilterGroup(RootModel):
- root: List[DynataFilterObject] = Field()
+ root: list[DynataFilterObject] = Field()
def __iter__(self):
return iter(self.root)
@@ -305,11 +305,11 @@ class DynataFilterGroup(RootModel):
def __hash__(self):
return hash(tuple(self.root))
- def passes(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# A filter group is matched if we match at least 1 filter objs in the group.
return any(obj.passes(criteria_evaluation) for obj in self.root)
- def passes_verbose(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_verbose(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# A filter group is matched if we match at least 1 filter objs in the group.
for obj in self.root:
print("---")
@@ -318,8 +318,8 @@ class DynataFilterGroup(RootModel):
return any(obj.passes(criteria_evaluation) for obj in self.root)
def passes_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]]:
# 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())
@@ -360,14 +360,14 @@ class DynataSurvey(MarketplaceTask):
)
# There are 91 min surveys. We'll filter them out later
- bid_loi: Optional[int] = Field(
+ bid_loi: int | None = Field(
default=None,
le=120 * 60,
description="Docs says 'Estimated length of interview', but this is "
"really the bid LOI'",
validation_alias="length_of_interview",
)
- bid_ir: Optional[float] = Field(validation_alias="incidence_rate", ge=0, le=1)
+ bid_ir: float | None = Field(validation_alias="incidence_rate", ge=0, le=1)
cpi: Decimal = Field(gt=0, le=100, validation_alias="cost_per_interview")
days_in_field: int = Field(description="Expected duration of opportunity in days")
# This isn't checked for eligibility determination
@@ -399,20 +399,20 @@ class DynataSurvey(MarketplaceTask):
category_exclusions: AlphaNumStrSet = Field(default_factory=set)
requirements: DynataRequirements = Field()
- filters: List[DynataFilterGroup] = Field(default_factory=list)
- quotas: List[DynataQuotaGroup] = Field(default_factory=list)
+ filters: list[DynataFilterGroup] = Field(default_factory=list)
+ quotas: list[DynataQuotaGroup] = Field(default_factory=list)
source: Literal[Source.DYNATA] = Field(default=Source.DYNATA)
- used_question_ids: Set[AlphaNumStr] = Field(default_factory=set)
+ used_question_ids: set[AlphaNumStr] = 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, DynataCondition]] = Field(default=None)
+ conditions: dict[str, DynataCondition] | None = Field(default=None)
# These do not come from the API. We set them ourselves
- last_updated: Optional[AwareDatetimeISO] = Field(default=None)
+ last_updated: AwareDatetimeISO | None = Field(default=None)
@property
def internal_id(self) -> str:
@@ -431,7 +431,7 @@ class DynataSurvey(MarketplaceTask):
@computed_field
@cached_property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
s = set()
for fg in self.filters:
for f in fg.root:
@@ -474,7 +474,7 @@ class DynataSurvey(MarketplaceTask):
return self
@property
- def filters_verbose(self) -> List[List[str]]:
+ def filters_verbose(self) -> list[list[str]]:
assert self.conditions is not None, "conditions must be set"
res = []
for filter_group in self.filters:
@@ -485,7 +485,7 @@ class DynataSurvey(MarketplaceTask):
return res
@property
- def quotas_verbose(self) -> List[List[Dict[str, Any]]]:
+ def quotas_verbose(self) -> list[list[dict[str, Any]]]:
assert self.conditions is not None, "conditions must be set"
res = []
for quota_group in self.quotas:
@@ -531,7 +531,7 @@ class DynataSurvey(MarketplaceTask):
exclude={"created", "last_updated", "conditions"}
) == other.model_dump(exclude={"created", "last_updated", "conditions"})
- def to_mysql(self) -> Dict[str, Any]:
+ def to_mysql(self) -> dict[str, Any]:
d = self.model_dump(
mode="json",
exclude={
@@ -560,12 +560,12 @@ class DynataSurvey(MarketplaceTask):
d["requirements"] = json.loads(d["requirements"])
return cls.model_validate(d)
- def passes_filters(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_filters(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# We have to match all filter groups
return all(group.passes(criteria_evaluation) for group in self.filters)
def passes_filters_verbose(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
# We have to match all filter groups
for group in self.filters:
@@ -575,8 +575,8 @@ class DynataSurvey(MarketplaceTask):
return all(group.passes(criteria_evaluation) for group in self.filters)
def passes_filters_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]]:
# We have to match all filter groups
group_eval = {
group: group.passes_soft(criteria_evaluation) for group in self.filters
@@ -592,14 +592,14 @@ class DynataSurvey(MarketplaceTask):
else:
return True, set()
- 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 all quota groups
return all(
quota_group.passes(criteria_evaluation) for quota_group in self.quotas
)
def passes_quotas_verbose(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
# We have to match all quota groups
for quota_group in self.quotas:
@@ -611,8 +611,8 @@ class DynataSurvey(MarketplaceTask):
)
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]]:
# We have to match all quota groups
group_eval = {
quota: quota.passes_soft(criteria_evaluation) for quota in self.quotas
@@ -629,7 +629,7 @@ class DynataSurvey(MarketplaceTask):
return True, set()
def determine_eligibility(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
return (
self.is_open
@@ -638,7 +638,7 @@ class DynataSurvey(MarketplaceTask):
)
def determine_eligibility_verbose(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
print(f"is_open: {self.is_open}")
print("passes_filters")
@@ -652,8 +652,8 @@ class DynataSurvey(MarketplaceTask):
)
def determine_eligibility_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]]:
if self.is_open is False:
return False, set()
pass_filters, h_filters = self.passes_filters_soft(criteria_evaluation)
diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py
index a10fcf0..b1b7c3e 100644
--- a/generalresearch/models/dynata/task_collection.py
+++ b/generalresearch/models/dynata/task_collection.py
@@ -1,4 +1,6 @@
-from typing import Any, Dict, List
+from __future__ import annotations
+
+from typing import Any
import pandas as pd
from pandera import Check, Column, DataFrameSchema, Index
@@ -35,8 +37,8 @@ DynataTaskCollectionSchema = DataFrameSchema(
"requirements": Column(str), # json dumped str
"created": 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(
@@ -55,7 +57,7 @@ class DynataTaskCollection(TaskCollection):
items: List[DynataSurvey]
_schema = DynataTaskCollectionSchema
- def to_row(self, s: DynataSurvey) -> Dict[str, Any]:
+ def to_row(self, s: DynataSurvey) -> dict[str, Any]:
d = s.model_dump(
mode="json",
exclude={
diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py
index 764c694..63e31c4 100644
--- a/generalresearch/models/gr/authentication.py
+++ b/generalresearch/models/gr/authentication.py
@@ -4,7 +4,7 @@ import binascii
import json
import os
from datetime import datetime, timezone
-from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
+from typing import TYPE_CHECKING, Any
from pydantic import (
AnyHttpUrl,
@@ -29,62 +29,62 @@ if TYPE_CHECKING:
class Claims(BaseModel):
- iss: Optional[str] = Field(
+ iss: str | None = Field(
default=None,
description="Issuer: https://www.rfc-editor.org/rfc/rfc7519.html#section-4.1.1",
)
- sub: Optional[str] = Field(
+ sub: str | None = Field(
default=None,
description="Subject: https://www.rfc-editor.org/rfc/rfc7519.html#section-4.1.2",
)
- aud: Optional[str] = Field(
+ aud: str | None = Field(
default=None,
description="Audience: https://www.rfc-editor.org/rfc/rfc7519.html#section-4.1.3",
)
- exp: Optional[NonNegativeInt] = Field(
+ exp: NonNegativeInt | None = Field(
default=None,
description="Expiration time: https://www.rfc-editor.org/rfc/rfc7519.html#section-4.1.4",
)
- iat: Optional[NonNegativeInt] = Field(
+ iat: NonNegativeInt | None = Field(
default=None,
description="Issued at: https://www.rfc-editor.org/rfc/rfc7519.html#section-4.1.6",
)
- auth_time: Optional[NonNegativeInt] = Field(
+ auth_time: NonNegativeInt | None = Field(
default=None,
description="When authentication occured: https://openid.net/specs/openid-connect-core-1_0.html#IDToken",
)
- acr: Optional[str] = Field(
+ acr: str | None = Field(
default=None,
description="Authentication Context Class Reference: https://openid.net/specs/openid-connect-core-1_0.html#IDToken",
)
- amr: Optional[List[str]] = Field(
+ amr: list[str] | None = Field(
default=None,
description="Authentication Methods References: https://openid.net/specs/openid-connect-core-1_0.html#IDToken",
)
- c_hash: Optional[str] = Field(
+ c_hash: str | None = Field(
default=None,
description="Code hash value: http://openid.net/specs/openid-connect-core-1_0.html",
)
- nonce: Optional[str] = Field(
+ nonce: str | None = Field(
default=None,
description="Value used to associate a Client session with an ID Token: http://openid.net/specs/openid-connect-core-1_0.html",
)
- at_hash: Optional[str] = Field(
+ at_hash: str | None = Field(
default=None,
description="Access Token hash value: http://openid.net/specs/openid-connect-core-1_0.html",
)
- sid: Optional[str] = Field(
+ sid: str | None = Field(
default=None,
description="Session ID: https://openid.net/specs/openid-connect-frontchannel-1_0.html#ClaimsContents",
)
@@ -103,8 +103,8 @@ class GRUser(BaseModel):
arbitrary_types_allowed=True
)
- id: Optional[PositiveInt] = Field(default=None)
- sub: Optional[str] = Field(max_length=200)
+ id: PositiveInt | None = Field(default=None)
+ sub: str | None = Field(max_length=200)
is_superuser: bool = Field(default=False)
date_joined: AwareDatetimeISO = Field(
@@ -112,14 +112,14 @@ class GRUser(BaseModel):
)
# prefetch attributes
- businesses: Optional[List["Business"]] = Field(default=None)
- teams: Optional[List["Team"]] = Field(default=None)
- products: Optional[List["Product"]] = Field(default=None)
- token: Optional["GRToken"] = Field(default=None)
- claims: Optional["Claims"] = Field(default=None)
+ businesses: list["Business"] | None = Field(default=None)
+ teams: list["Team"] | None = Field(default=None)
+ products: list["Product"] | None = Field(default=None)
+ token: "GRToken" | None = Field(default=None)
+ claims: "Claims" | None = Field(default=None)
def prefetch_claims(
- self, token: str, key: Dict[str, Any], audience: str, issuer: AnyHttpUrl
+ self, token: str, key: dict[str, Any], audience: str, issuer: AnyHttpUrl
) -> None:
from jose import jwt
@@ -170,7 +170,7 @@ class GRUser(BaseModel):
if len(business_uuids + team_uuids) == 0:
self.products = []
- return None
+ return
from generalresearch.managers.thl.product import ProductManager
@@ -207,7 +207,7 @@ class GRUser(BaseModel):
return f"gr_user:{self.id}"
@property
- def business_uuids(self) -> Optional[List[UUIDStr]]:
+ def business_uuids(self) -> list[UUIDStr] | None:
if self.businesses is None:
LOG.warning("prefetch not run")
return None
@@ -215,7 +215,7 @@ class GRUser(BaseModel):
return [b.uuid for b in self.businesses]
@property
- def business_ids(self) -> Optional[List[PositiveInt]]:
+ def business_ids(self) -> list[PositiveInt] | None:
if self.businesses is None:
LOG.warning("prefetch not run")
return None
@@ -223,7 +223,7 @@ class GRUser(BaseModel):
return [b.id for b in self.businesses if b.id is not None]
@property
- def team_uuids(self) -> Optional[List[UUIDStr]]:
+ def team_uuids(self) -> list[UUIDStr] | None:
if self.teams is None:
LOG.warning("prefetch not run")
return None
@@ -231,7 +231,7 @@ class GRUser(BaseModel):
return [t.uuid for t in self.teams]
@property
- def team_ids(self) -> Optional[List[PositiveInt]]:
+ def team_ids(self) -> list[PositiveInt] | None:
if self.teams is None:
LOG.warning("prefetch not run")
return None
@@ -239,7 +239,7 @@ class GRUser(BaseModel):
return [t.id for t in self.teams if t.id is not None]
@property
- def product_uuids(self) -> Optional[List[UUIDStr]]:
+ def product_uuids(self) -> list[UUIDStr] | None:
if self.products is None:
LOG.warning("prefetch not run")
return None
@@ -294,7 +294,7 @@ class GRUser(BaseModel):
return GRUser.model_validate(d)
@classmethod
- def from_redis(cls, d: Union[str, Dict[str, Any]]) -> Self:
+ def from_redis(cls, d: str | dict[str, Any]) -> Self:
if isinstance(d, str):
d = json.loads(d)
assert isinstance(d, dict)
@@ -327,7 +327,7 @@ class GRToken(BaseModel):
user_id: PositiveInt = Field()
# --- prefetch field ---
- user: Optional["GRUser"] = Field(default=None)
+ user: "GRUser" | None = Field(default=None)
@property
def sso(self) -> bool:
@@ -359,13 +359,13 @@ class GRToken(BaseModel):
# --- Properties ---
@property
- def auth_header(self, key_name: str = "Authorization") -> Dict[str, str]:
+ def auth_header(self, key_name: str = "Authorization") -> dict[str, str]:
return {key_name: self.key}
# --- ORM ---
@classmethod
- def from_redis(cls, d: Union[str, Dict[str, Any]]) -> Self:
+ def from_redis(cls, d: str | dict[str, Any]) -> Self:
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 8a37f74..2df7d61 100644
--- a/generalresearch/models/gr/business.py
+++ b/generalresearch/models/gr/business.py
@@ -6,7 +6,7 @@ import os
from datetime import datetime, timezone
from enum import Enum
from pathlib import Path
-from typing import TYPE_CHECKING, List, Optional, Union
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
@@ -76,7 +76,7 @@ class BusinessBankAccount(BaseModel):
json_encoders={TransferMethod: lambda tm: tm.value},
)
- id: SkipJsonSchema[Optional[PositiveInt]] = Field(default=None)
+ id: SkipJsonSchema[PositiveInt | None] = Field(default=None)
uuid: UUIDStrCoerce = Field(examples=[uuid4().hex])
business_id: PositiveInt = Field()
@@ -84,7 +84,7 @@ class BusinessBankAccount(BaseModel):
# 'business' is a Class with values that are fetched from the DB.
# Initialization is deferred until it is actually needed
# (see .prefetch_business())
- business: SkipJsonSchema[Optional["Business"]] = Field(default=None)
+ business: SkipJsonSchema["Business" | None] = Field(default=None)
transfer_method: TransferMethod = Field(
description=TransferMethod.as_openapi(),
@@ -92,14 +92,14 @@ class BusinessBankAccount(BaseModel):
)
# ACH requirements
- account_number: Optional[str] = Field(
+ account_number: str | None = Field(
default=None,
max_length=16,
description="ACH requirements",
examples=[f"{'*' * 9}1234"],
)
- routing_number: Optional[str] = Field(
+ routing_number: str | None = Field(
default=None,
max_length=9,
description="ACH requirements",
@@ -107,13 +107,13 @@ class BusinessBankAccount(BaseModel):
)
# Wire requirements
- iban: Optional[str] = Field(
+ iban: str | None = Field(
default=None,
max_length=50,
description="Wire requirements",
examples=[None],
)
- swift: Optional[str] = Field(
+ swift: str | None = Field(
default=None,
max_length=50,
description="Wire requirements",
@@ -133,20 +133,16 @@ class BusinessBankAccount(BaseModel):
class BusinessAddress(BaseModel):
model_config = ConfigDict(extra="ignore")
- id: SkipJsonSchema[Optional[PositiveInt]] = Field(default=None)
+ id: SkipJsonSchema[PositiveInt | None] = Field(default=None)
uuid: UUIDStrCoerce = Field(examples=[uuid4().hex])
- line_1: Optional[str] = Field(
- default=None, max_length=255, examples=["540 Mariposa"]
- )
+ line_1: str | None = Field(default=None, max_length=255, examples=["540 Mariposa"])
- line_2: Optional[str] = Field(default=None, max_length=255, examples=[None])
+ line_2: str | None = Field(default=None, max_length=255, examples=[None])
- city: Optional[str] = Field(
- default=None, max_length=255, examples=["Mountain View"]
- )
+ city: str | None = Field(default=None, max_length=255, examples=["Mountain View"])
- state: Optional[str] = Field(
+ state: str | None = Field(
default=None,
max_length=255,
description="This can only be more than len=2 if it's a state or"
@@ -154,11 +150,11 @@ class BusinessAddress(BaseModel):
examples=["CA"],
)
- postal_code: Optional[str] = Field(default=None, max_length=12, examples=["94041"])
+ postal_code: str | None = Field(default=None, max_length=12, examples=["94041"])
- phone_number: Optional[PhoneNumber] = Field(default=None)
+ phone_number: PhoneNumber | None = Field(default=None)
- country: Optional[str] = Field(default=None, max_length=2, examples=["US"])
+ country: str | None = Field(default=None, max_length=2, examples=["US"])
business_id: PositiveInt = Field()
@@ -166,10 +162,10 @@ class BusinessAddress(BaseModel):
class BusinessContact(BaseModel):
model_config = ConfigDict(extra="ignore")
- name: Optional[str] = Field(default=None)
- email: Optional[str] = Field(default=None)
+ name: str | None = Field(default=None)
+ email: str | None = Field(default=None)
- phone_number: Optional[str] = Field(
+ phone_number: str | None = Field(
default=None,
min_length=10,
max_length=31,
@@ -182,7 +178,7 @@ class Business(BaseModel):
model_config = ConfigDict(extra="ignore")
- id: SkipJsonSchema[Optional[PositiveInt]] = Field(default=None)
+ id: SkipJsonSchema[PositiveInt | None] = Field(default=None)
uuid: UUIDStrCoerce = Field(examples=[uuid4().hex])
name: str = Field(
@@ -197,23 +193,23 @@ class Business(BaseModel):
examples=[BusinessType.COMPANY.value],
)
- tax_number: Optional[str] = Field(default=None, max_length=20)
- contact: Optional["BusinessContact"] = Field(default=None)
+ tax_number: str | None = Field(default=None, max_length=20)
+ contact: "BusinessContact" | None = Field(default=None)
# Initialization is deferred until it is actually needed
# (see .prefetch_***())
- addresses: Optional[List["BusinessAddress"]] = Field(default=None)
- teams: Optional[List["Team"]] = Field(default=None)
- products: Optional[List["Product"]] = Field(default=None)
- bank_accounts: Optional[List["BusinessBankAccount"]] = Field(default=None)
+ addresses: list["BusinessAddress"] | None = Field(default=None)
+ teams: list["Team"] | None = Field(default=None)
+ products: list["Product"] | None = Field(default=None)
+ bank_accounts: list["BusinessBankAccount"] | None = Field(default=None)
# Initialization is deferred until unless it's called
# (see .prebuild_***())
- balance: Optional["BusinessBalances"] = Field(default=None, name="Business Balance")
+ balance: "BusinessBalances" | None = Field(default=None, name="Business Balance")
- payouts_total_str: Optional[str] = Field(default=None)
- payouts_total: Optional[USDCent] = Field(default=None)
- payouts: Optional[List[BusinessPayoutEvent]] = Field(
+ 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"
@@ -221,8 +217,8 @@ class Business(BaseModel):
"child Products",
)
- pop_financial: Optional[List[POPFinancial]] = Field(default=None)
- bp_accounts: Optional[List[LedgerAccount]] = Field(default=None)
+ pop_financial: list[POPFinancial] | None = Field(default=None)
+ bp_accounts: list[LedgerAccount] | None = Field(default=None)
def __str__(self) -> str:
return (
@@ -327,8 +323,8 @@ class Business(BaseModel):
lm: "LedgerManager",
ds: "GRLDatasets",
client: Client,
- pop_ledger: Optional["PopLedgerMerge"] = None,
- at_timestamp: Optional[AwareDatetime] = None,
+ pop_ledger: "PopLedgerMerge" | None = None,
+ at_timestamp: AwareDatetime | None = None,
) -> None:
"""
This returns the Business's Balances that are calculated across
@@ -356,7 +352,7 @@ class Business(BaseModel):
self.prefetch_products(thl_pg_config=thl_pg_config)
- accounts: List[LedgerAccount] = lm.get_accounts_if_exists(
+ accounts: list[LedgerAccount] = lm.get_accounts_if_exists(
qualified_names=(
[f"{lm.currency.value}:bp_wallet:{bpid}" for bpid in self.product_uuids]
if self.product_uuids
@@ -414,7 +410,7 @@ class Business(BaseModel):
input_data=df, accounts=accounts, thl_pg_config=thl_pg_config
)
- return None
+ return
def prebuild_payouts(
self,
@@ -439,7 +435,7 @@ 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 None
+ return
def prebuild_pop_financial(
self,
@@ -447,7 +443,7 @@ class Business(BaseModel):
lm: "LedgerManager",
ds: "GRLDatasets",
client: Client,
- pop_ledger: Optional["PopLedgerMerge"] = None,
+ pop_ledger: "PopLedgerMerge" | None = None,
) -> None:
"""This is very similar to the Product POP Financial endpoint; however,
it returns more than one item for a single time interval. This is
@@ -480,13 +476,13 @@ class Business(BaseModel):
)
if ddf is None:
self.pop_financial = []
- return None
+ return
df = client.compute(collections=ddf, sync=True)
if df.empty:
self.pop_financial = []
- return None
+ return
df = df.groupby(
[pd.Grouper(key="time_idx", freq=rr.interval), "account_id"]
@@ -502,7 +498,7 @@ class Business(BaseModel):
ds: "GRLDatasets",
client: Client,
mnt_gr_api: Path,
- enriched_session: Optional["EnrichedSessionMerge"] = None,
+ enriched_session: "EnrichedSessionMerge" | None = None,
) -> None:
self.prefetch_products(thl_pg_config=thl_pg_config)
@@ -547,7 +543,7 @@ class Business(BaseModel):
ds: "GRLDatasets",
client: Client,
mnt_gr_api: Path,
- enriched_wall: Optional["EnrichedWallMerge"] = None,
+ enriched_wall: "EnrichedWallMerge" | None = None,
) -> None:
self.prefetch_products(thl_pg_config=thl_pg_config)
@@ -587,7 +583,7 @@ class Business(BaseModel):
return None
@classmethod
- def required_fields(cls) -> List[str]:
+ def required_fields(cls) -> list[str]:
return [
field_name
for field_name, field_info in cls.model_fields.items()
@@ -597,7 +593,7 @@ class Business(BaseModel):
# --- Properties ---
@property
- def product_uuids(self) -> Optional[List[UUIDStr]]:
+ def product_uuids(self) -> list[UUIDStr] | None:
if self.products is None:
LOG.warning("prefetch not run")
return None
@@ -624,10 +620,10 @@ class Business(BaseModel):
lm: "LedgerManager",
thl_lm: "ThlLedgerManager",
bpem: "BusinessPayoutEventManager",
- mnt_gr_api: Union[Path, str],
- pop_ledger: Optional["PopLedgerMerge"] = None,
- enriched_session: Optional["EnrichedSessionMerge"] = None,
- enriched_wall: Optional["EnrichedWallMerge"] = None,
+ mnt_gr_api: Path | str,
+ pop_ledger: "PopLedgerMerge" | None = None,
+ enriched_session: "EnrichedSessionMerge" | None = None,
+ enriched_wall: "EnrichedWallMerge" | None = None,
) -> None:
LOG.debug(f"Business.set_cache({self.uuid=})")
@@ -700,18 +696,16 @@ class Business(BaseModel):
enriched_wall=enriched_wall,
)
- return None
-
# --- ORM ---
@classmethod
def from_redis(
cls,
uuid: UUIDStr,
- fields: List[str],
+ fields: list[str],
gr_redis_config: RedisConfig,
- ) -> Optional[Self]:
- keys: List[str] = Business.required_fields() + fields
+ ) -> Self | 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,
@@ -721,7 +715,7 @@ class Business(BaseModel):
rc = gr_redis_config.create_redis_client()
try:
- res: List = rc.hmget(name=f"business:{uuid}", keys=keys)
+ res: list = 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)
diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py
index a3ac9cf..2c02f6f 100644
--- a/generalresearch/models/gr/team.py
+++ b/generalresearch/models/gr/team.py
@@ -1,9 +1,11 @@
+from __future__ import annotations
+
import json
import os
from datetime import datetime, timezone
from enum import Enum
from pathlib import Path
-from typing import TYPE_CHECKING, List, Optional, Union
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
@@ -58,7 +60,7 @@ class Membership(BaseModel):
model_config = ConfigDict(use_enum_values=True)
- id: SkipJsonSchema[Optional[PositiveInt]] = Field(
+ id: SkipJsonSchema[PositiveInt | None] = Field(
default=None,
)
uuid: UUIDStrCoerce = Field(examples=[uuid4().hex])
@@ -82,13 +84,13 @@ class Membership(BaseModel):
team_id: SkipJsonSchema[PositiveInt] = Field()
# prefetch attributes
- team: SkipJsonSchema[Optional["Team"]] = Field(default=None)
+ team: SkipJsonSchema["Team" | None] = Field(default=None)
# --- Validators ---
@field_validator("created", mode="before")
@classmethod
- def created_utc(cls, v: Union[datetime, str]) -> Union[datetime, str]:
+ def created_utc(cls, v: datetime | str) -> datetime | str:
if isinstance(v, datetime):
return v.replace(tzinfo=timezone.utc)
return v
@@ -105,15 +107,15 @@ class Membership(BaseModel):
class Team(BaseModel):
- id: SkipJsonSchema[Optional[PositiveInt]] = Field(default=None)
+ id: SkipJsonSchema[PositiveInt | None] = Field(default=None)
uuid: UUIDStrCoerce = Field(examples=[uuid4().hex])
name: str = Field(max_length=255, examples=["Team ABC"])
# prefetch attributes
- memberships: SkipJsonSchema[Optional[List["Membership"]]] = Field(default=None)
- gr_users: SkipJsonSchema[Optional[List["GRUser"]]] = Field(default=None)
- businesses: SkipJsonSchema[Optional[List["Business"]]] = Field(default=None)
- products: SkipJsonSchema[Optional[List["Product"]]] = Field(default=None)
+ memberships: SkipJsonSchema[list["Membership"] | None] = Field(default=None)
+ gr_users: SkipJsonSchema[list["GRUser"] | None] = Field(default=None)
+ businesses: SkipJsonSchema[list["Business"] | None] = Field(default=None)
+ products: SkipJsonSchema[list["Product"] | None] = Field(default=None)
# --- Prefetch Methods ---
@@ -156,7 +158,7 @@ class Team(BaseModel):
ds: "GRLDatasets",
client: Client,
mnt_gr_api: Path,
- enriched_session: Optional["EnrichedSessionMerge"] = None,
+ enriched_session: "EnrichedSessionMerge" | None = None,
) -> None:
self.prefetch_products(thl_pg_config=thl_pg_config)
@@ -193,7 +195,7 @@ class Team(BaseModel):
except Exception as e:
raise IOError(f"Parquet verification failed: {e}")
- return None
+ return
def prebuild_enriched_wall_parquet(
self,
@@ -201,7 +203,7 @@ class Team(BaseModel):
ds: "GRLDatasets",
client: Client,
mnt_gr_api: Path,
- enriched_wall: Optional["EnrichedWallMerge"] = None,
+ enriched_wall: "EnrichedWallMerge" | None = None,
) -> None:
self.prefetch_products(thl_pg_config=thl_pg_config)
@@ -241,7 +243,7 @@ class Team(BaseModel):
return None
@classmethod
- def required_fields(cls) -> List[str]:
+ def required_fields(cls) -> list[str]:
return [
field_name
for field_name, field_info in cls.model_fields.items()
@@ -258,7 +260,7 @@ class Team(BaseModel):
return f"team-{self.uuid}"
@property
- def product_ids(self) -> Optional[List[UUIDStr]]:
+ def product_ids(self) -> list[UUIDStr] | None:
if self.products is None:
LOG.warning("prefetch not run")
return None
@@ -266,7 +268,7 @@ class Team(BaseModel):
return [p.uuid for p in self.products]
@property
- def product_uuids(self) -> Optional[List[UUIDStr]]:
+ def product_uuids(self) -> list[UUIDStr] | None:
return self.product_ids
# --- Methods ---
@@ -278,9 +280,9 @@ class Team(BaseModel):
redis_config: RedisConfig,
client: "Client",
ds: "GRLDatasets",
- mnt_gr_api: Union[Path, str],
- enriched_session: Optional["EnrichedSessionMerge"] = None,
- enriched_wall: Optional["EnrichedWallMerge"] = None,
+ mnt_gr_api: Path | str,
+ enriched_session: "EnrichedSessionMerge" | None = None,
+ enriched_wall: "EnrichedWallMerge" | None = None,
) -> None:
ex_secs = 60 * 60 * 24 * 3 # 3 days
@@ -324,7 +326,7 @@ class Team(BaseModel):
enriched_wall=enriched_wall,
)
- return None
+ return
# --- ORM ---
@@ -332,14 +334,14 @@ class Team(BaseModel):
def from_redis(
cls,
uuid: UUIDStr,
- fields: List[str],
+ fields: list[str],
gr_redis_config: RedisConfig,
- ) -> Optional[Self]:
- keys: List = Team.required_fields() + fields
+ ) -> Self | 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 = 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)
diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py
index 0d47ba7..402b35d 100644
--- a/generalresearch/models/innovate/question.py
+++ b/generalresearch/models/innovate/question.py
@@ -4,7 +4,7 @@ from __future__ import annotations
import json
import logging
from enum import Enum
-from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
+from typing import TYPE_CHECKING, Any, Literal
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -28,7 +28,7 @@ logger.setLevel(logging.INFO)
class InnovateUserQuestionAnswer(MarketplaceUserQuestionAnswer):
# Note, this is referred to as the KEY in the Question model
question_id: InnovateQuestionID = Field()
- question_type: Optional[InnovateQuestionType] = Field(default=None)
+ question_type: InnovateQuestionType | None = Field(default=None)
# Did this answer come from us asking, or was it passed back from the marketplace
from_thl: bool = Field(default=True)
@@ -104,8 +104,8 @@ class InnovateQuestion(MarketplaceQuestion):
# This comes from the API field "Category". There are some useful categories in here, but a bunch have
# categories that are not (e.g. NFX - Adhoc, Testing_Cat). We'll store it as a comma-separated string
# here to use it to aid our own real categorization.
- tags: Optional[str] = Field(default=None, frozen=True)
- options: Optional[List[InnovateQuestionOption]] = Field(
+ tags: str | None = Field(default=None, frozen=True)
+ options: list[InnovateQuestionOption] | None = Field(
default=None, min_length=1, frozen=True
)
@@ -141,8 +141,8 @@ class InnovateQuestion(MarketplaceQuestion):
@classmethod
def from_api(
- cls, d: Dict, country_iso: str, language_iso: str
- ) -> Optional["InnovateQuestion"]:
+ cls, d: dict, country_iso: str, language_iso: str
+ ) -> "InnovateQuestion" | None:
"""
:param d: Raw response from API
:param country_iso:
@@ -185,7 +185,7 @@ class InnovateQuestion(MarketplaceQuestion):
)
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> "InnovateQuestion":
+ def from_db(cls, d: dict[str, Any]) -> "InnovateQuestion":
options = None
if d["options"]:
@@ -207,7 +207,7 @@ class InnovateQuestion(MarketplaceQuestion):
tags=d["tags"],
)
- 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/innovate/survey.py b/generalresearch/models/innovate/survey.py
index 3a3d9e2..8b5d24c 100644
--- a/generalresearch/models/innovate/survey.py
+++ b/generalresearch/models/innovate/survey.py
@@ -8,12 +8,7 @@ from functools import cached_property
from typing import (
Annotated,
Any,
- Dict,
- List,
Literal,
- Optional,
- Set,
- Tuple,
Type,
)
@@ -62,16 +57,16 @@ locale_helper = Localelator()
class InnovateCondition(MarketplaceCondition):
model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore")
# store everything lowercase !
- question_id: Optional[CoercedStr] = Field(
+ question_id: CoercedStr | None = Field(
min_length=1, max_length=64, pattern=r"^[^A-Z]+$"
)
# There isn't really a hard limit, but their API is inconsistent and
# sometimes returns all the options comma-separated instead of as a list.
# Try to catch that.
- values: List[Annotated[str, Field(max_length=128)]] = Field()
+ values: list[Annotated[str, Field(max_length=128)]] = Field()
@classmethod
- def from_api(cls, d: Dict[str, Any]) -> "InnovateCondition":
+ def from_api(cls, d: dict[str, Any]) -> "InnovateCondition":
d["logical_operator"] = LogicalOperator.OR
d["value_type"] = ConditionValueType.LIST
d["negate"] = False
@@ -91,7 +86,7 @@ class InnovateQuota(BaseModel):
task_calculation_type: TaskCalculationType = Field()
hard_stop: bool = Field()
- condition_hashes: List[str] = Field(min_length=0, default_factory=list)
+ condition_hashes: list[str] = Field(min_length=0, default_factory=list)
def __hash__(self):
return hash(tuple((tuple(self.condition_hashes), self.remaining_count)))
@@ -105,21 +100,21 @@ class InnovateQuota(BaseModel):
)
@classmethod
- def from_api(cls, d: Dict):
+ def from_api(cls, d: dict):
return cls.model_validate(d)
- 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)
- 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.
# We can "match" a quota that is closed. In that case, we would not be eligible for the survey.
return all(criteria_evaluation.get(c) for c in self.condition_hashes)
def matches_optional(
- self, criteria_evaluation: Dict[str, Optional[bool]]
- ) -> Optional[bool]:
+ self, criteria_evaluation: dict[str, bool | None]
+ ) -> bool | None:
# We need to know if any conditions are unknown to avoid matching a full quota. If any fail,
# then we know we fail regardless of any being unknown.
evals = [criteria_evaluation.get(c) for c in self.condition_hashes]
@@ -130,8 +125,8 @@ class InnovateQuota(BaseModel):
return True
def matches_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]]:
# Passes back "matches" (T/F/none) and a list of unknown criterion hashes
hash_evals = {
cell: criteria_evaluation.get(cell) for cell in self.condition_hashes
@@ -171,11 +166,11 @@ class InnovateSurvey(MarketplaceTask):
supplier_completes_achieved: int = Field()
global_completes: int = Field()
global_starts: int = Field()
- global_median_loi: Optional[int] = Field(le=120 * 60)
- global_conversion: Optional[float] = Field(ge=0, le=1)
+ global_median_loi: int | None = Field(le=120 * 60)
+ global_conversion: float | None = Field(ge=0, le=1)
- bid_loi: Optional[int] = Field(default=None, le=120 * 60)
- bid_ir: Optional[float] = Field(default=None, ge=0, le=1)
+ bid_loi: int | None = Field(default=None, le=120 * 60)
+ bid_ir: float | None = Field(default=None, ge=0, le=1)
allowed_devices: DeviceTypes = Field(min_length=1)
@@ -183,31 +178,31 @@ class InnovateSurvey(MarketplaceTask):
category: str = Field()
requires_pii: bool = Field(default=False)
- excluded_surveys: Optional[AlphaNumStrSet] = Field(
+ excluded_surveys: AlphaNumStrSet | None = Field(
description="list of excluded survey ids", default=None
)
duplicate_check_level: InnovateDuplicateCheckLevel = Field()
- exclude_pids: Optional[AlphaNumStrSet] = Field(default=None)
- include_pids: Optional[AlphaNumStrSet] = Field(default=None)
+ exclude_pids: AlphaNumStrSet | None = Field(default=None)
+ include_pids: AlphaNumStrSet | None = Field(default=None)
# idk what these mean
is_revenue_sharing: bool = Field()
group_type: str = Field()
# undocumented, not sure how we use this
- off_hour_traffic: Optional[Dict] = Field(default=None)
+ off_hour_traffic: dict | None = Field(default=None)
- qualifications: List[str] = Field(default_factory=list)
- quotas: List[InnovateQuota] = Field(default_factory=list)
+ qualifications: list[str] = Field(default_factory=list)
+ quotas: list[InnovateQuota] = Field(default_factory=list)
source: Literal[Source.INNOVATE] = Field(default=Source.INNOVATE)
- used_question_ids: Set[InnovateQuestionID] = Field(default_factory=set)
+ used_question_ids: set[InnovateQuestionID] = 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, InnovateCondition]] = Field(default=None)
+ conditions: dict[str, InnovateCondition] | None = Field(default=None)
# These come from the API
created_api: AwareDatetimeISO = Field(
@@ -219,8 +214,8 @@ class InnovateSurvey(MarketplaceTask):
expected_end_date: date = Field()
# 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:
@@ -239,7 +234,7 @@ class InnovateSurvey(MarketplaceTask):
@computed_field
@cached_property
- def all_hashes(self) -> Set[str]:
+ def all_hashes(self) -> set[str]:
s = set(self.qualifications)
for q in self.quotas:
s.update(set(q.condition_hashes))
@@ -266,7 +261,7 @@ class InnovateSurvey(MarketplaceTask):
return data
@classmethod
- def from_api(cls, d: Dict[str, Any]) -> Optional["InnovateSurvey"]:
+ def from_api(cls, d: dict[str, Any]) -> "InnovateSurvey" | None:
try:
return cls._from_api(d)
except Exception as e:
@@ -274,7 +269,7 @@ class InnovateSurvey(MarketplaceTask):
return None
@classmethod
- def _from_api(cls, d: Dict[str, Any]) -> "InnovateSurvey":
+ def _from_api(cls, d: dict[str, Any]) -> "InnovateSurvey":
d["conditions"] = dict()
# If we haven't hit the "detail" endpoint, we won't get this
@@ -343,7 +338,7 @@ class InnovateSurvey(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={
@@ -365,7 +360,7 @@ class InnovateSurvey(MarketplaceTask):
return d
@classmethod
- def from_db(cls, d: Dict[str, Any]) -> Self:
+ 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)
@@ -377,7 +372,7 @@ class InnovateSurvey(MarketplaceTask):
return cls.model_validate(d)
def participation_allowed(
- self, att_survey_ids: Set[str], att_job_ids: Set[str]
+ self, att_survey_ids: set[str], att_job_ids: set[str]
) -> bool:
"""
Checks if this user can participate in this survey based on the 'duplicate_check_level'-dictated requirements
@@ -397,14 +392,14 @@ class InnovateSurvey(MarketplaceTask):
return True
def passes_qualifications(
- self, criteria_evaluation: Dict[str, Optional[bool]]
+ self, criteria_evaluation: dict[str, bool | None]
) -> bool:
# We have to match all quals
return all(criteria_evaluation.get(q) for q in self.qualifications)
def passes_qualifications_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]]:
# Passes back "passes" (T/F/none) and a list of unknown criterion hashes
hash_evals = {q: criteria_evaluation.get(q) for q in self.qualifications}
evals = set(hash_evals.values())
@@ -416,7 +411,7 @@ class InnovateSurvey(MarketplaceTask):
return None, {cell for cell, ev in hash_evals.items() if ev is None}
return True, set()
- def passes_quotas(self, criteria_evaluation: Dict[str, Optional[bool]]) -> bool:
+ def passes_quotas(self, criteria_evaluation: dict[str, bool | None]) -> bool:
# Many surveys have 0 quotas. Quotas are exclusionary.
# They can NOT match a quota where currently_open=0
any_pass = True
@@ -428,8 +423,8 @@ class InnovateSurvey(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]]:
# Many surveys have 0 quotas. Quotas are exclusionary.
# They can NOT match a quota where currently_open=0
if len(self.quotas) == 0:
@@ -469,7 +464,7 @@ class InnovateSurvey(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
@@ -478,8 +473,8 @@ class InnovateSurvey(MarketplaceTask):
)
def determine_eligibility_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]]:
if self.is_open is False:
return False, set()
pass_quals, h_quals = self.passes_qualifications_soft(criteria_evaluation)
diff --git a/generalresearch/models/innovate/task_collection.py b/generalresearch/models/innovate/task_collection.py
index 4d7fe51..88cac8c 100644
--- a/generalresearch/models/innovate/task_collection.py
+++ b/generalresearch/models/innovate/task_collection.py
@@ -1,4 +1,4 @@
-from typing import List, Set
+from __future__ import annotations
import pandas as pd
from pandera import Check, Column, DataFrameSchema, Index
@@ -11,8 +11,8 @@ from generalresearch.models.thl.survey.task_collection import (
create_empty_df_from_schema,
)
-COUNTRY_ISOS: Set[str] = Localelator().get_all_countries()
-LANGUAGE_ISOS: Set[str] = Localelator().get_all_languages()
+COUNTRY_ISOS: set[str] = Localelator().get_all_countries()
+LANGUAGE_ISOS: set[str] = Localelator().get_all_languages()
InnovateTaskCollectionSchema = DataFrameSchema(
columns={
@@ -42,8 +42,8 @@ InnovateTaskCollectionSchema = DataFrameSchema(
"created_api": Column(dtype=pd.DatetimeTZDtype(tz="UTC")),
"modified_api": 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(
@@ -59,7 +59,7 @@ InnovateTaskCollectionSchema = DataFrameSchema(
class InnovateTaskCollection(TaskCollection):
- items: List[InnovateSurvey]
+ items: list[InnovateSurvey]
_schema = InnovateTaskCollectionSchema
def to_row(self, s: InnovateSurvey):
diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py
index f7c9675..b9a7d79 100644
--- a/generalresearch/pg_helper.py
+++ b/generalresearch/pg_helper.py
@@ -1,5 +1,6 @@
+from __future__ import annotations
+
from datetime import timezone
-from typing import Optional
import psycopg
from psycopg.adapt import Buffer
@@ -51,7 +52,7 @@ class PostgresConfig:
dsn: PostgresDsn,
connect_timeout: int,
statement_timeout: float,
- schema: Optional[str] = None,
+ schema: str | None = None,
row_factory: RowFactory = dict_row,
):
"""
diff --git a/generalresearch/priority_thread_pool.py b/generalresearch/priority_thread_pool.py
index 9d254e4..b33df87 100644
--- a/generalresearch/priority_thread_pool.py
+++ b/generalresearch/priority_thread_pool.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
import time
from concurrent.futures.thread import ThreadPoolExecutor, _WorkItem
from queue import PriorityQueue
@@ -46,7 +48,7 @@ class PriorityThreadPoolExecutor(ThreadPoolExecutor):
>> q.submit(do_nothing, 'high', priority=-1)
"""
- def __init__(self, max_workers=None, thread_name_prefix=""):
+ def __init__(self, max_workers: int | None = None, thread_name_prefix=""):
super().__init__(max_workers, thread_name_prefix)
self._work_queue = WorkItemPriorityQueue()
diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py
index 175830c..efdf6f5 100644
--- a/generalresearch/sql_helper.py
+++ b/generalresearch/sql_helper.py
@@ -1,16 +1,18 @@
+from __future__ import annotations
+
import logging
-from typing import Any, Dict, List, Optional, Tuple, Union
+from typing import Any, Optional
from uuid import UUID
from pydantic import MariaDBDsn, MySQLDsn, PostgresDsn
from pymysql import Connection
-ListOrTupleOfStrings = Union[List[str], Tuple[str, ...]]
-ListOrTupleOfListOrTuple = Union[
- List[List], List[Tuple], Tuple[List, ...], Tuple[Tuple, ...]
-]
+ListOrTupleOfStrings = list[str] | tuple[str, ...] | None
+ListOrTupleOfListOrTuple = (
+ list[list] | list[tuple] | tuple[list, ...] | tuple[tuple, ...]
+)
-DataBaseDsn = Union[MySQLDsn, MariaDBDsn, PostgresDsn]
+DataBaseDsn = MySQLDsn | MariaDBDsn | PostgresDsn | None
class MultipleObjectsReturned(Exception):
@@ -26,7 +28,7 @@ class SqlConnector:
# For connection and cursor handling, and any difference between mysql
# and postgresql
- def __init__(self, dsn: Optional[DataBaseDsn] = None, **kwargs):
+ def __init__(self, dsn: DataBaseDsn | None = None, **kwargs):
"""
Anything in kwargs gets passed into the engine_module's connect
function. To be used for e.g.:
@@ -119,7 +121,7 @@ def is_uuid4(s: Any) -> bool:
return False
-def decode_uuids(row: Dict[str, Any]) -> Dict[str, Any]:
+def decode_uuids(row: dict[str, Any]) -> dict[str, Any]:
return {
key: (UUID(value, version=4).hex if is_uuid4(value) else value)
for key, value in row.items()
@@ -132,8 +134,8 @@ class SqlHelper(SqlConnector):
super(SqlHelper, self).__init__(dsn, **kwargs)
def execute_sql_query(
- self, query: str, params: Optional[Dict[str, Any]] = None, commit: bool = False
- ) -> List[Dict[str, Any]]:
+ self, query: str, params: dict[str, Any] | None = None, commit: bool = False
+ ) -> 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")
@@ -207,7 +209,7 @@ class SqlHelper(SqlConnector):
cursor=None,
) -> None:
if len(values_to_insert) == 0:
- return None
+ return
assert len(set([len(x) for x in values_to_insert])) == 1
if cursor is None:
@@ -240,7 +242,7 @@ class SqlHelper(SqlConnector):
lookup_dict: dict,
update_dict: dict,
cursor=None,
- ) -> Tuple[Union[str, int], bool]:
+ ) -> tuple[str | int, bool]:
"""
returns the value of the primary key ONLY, and bool (created)
"""
@@ -310,7 +312,7 @@ class SqlHelper(SqlConnector):
filter_d=None,
limit=None,
cursor=None,
- ) -> List[Dict[Any, Any]]:
+ ) -> list[dict[Any, Any]]:
if cursor is None:
connection = self.make_connection()
diff --git a/tests/sql_helper.py b/tests/sql_helper.py
index 247f0cd..c4cc2ca 100644
--- a/tests/sql_helper.py
+++ b/tests/sql_helper.py
@@ -1,7 +1,7 @@
from uuid import uuid4
import pytest
-from pydantic import MySQLDsn, MariaDBDsn, ValidationError
+from pydantic import MariaDBDsn, MySQLDsn, ValidationError
class TestSqlHelper: