From 6cf7ccbaa8306700e64ada19d6f99807743b2865 Mon Sep 17 00:00:00 2001
From: Max Nanis
Date: Fri, 21 Aug 2026 17:19:41 -0700
Subject: Ruff auto updates to 3.14
---
tests/incite/schemas/test_admin_responses.py | 29 +++++++++++++---------------
1 file changed, 13 insertions(+), 16 deletions(-)
(limited to 'tests/incite/schemas/test_admin_responses.py')
diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py
index 43aa399..29d93fe 100644
--- a/tests/incite/schemas/test_admin_responses.py
+++ b/tests/incite/schemas/test_admin_responses.py
@@ -1,4 +1,4 @@
-from datetime import datetime, timezone, timedelta
+from datetime import UTC, datetime, timedelta, timezone
from random import sample
from typing import List
@@ -8,8 +8,8 @@ import pytest
from generalresearch.incite.schemas import empty_dataframe_from_schema
from generalresearch.incite.schemas.admin_responses import (
- AdminPOPSchema,
SIX_HOUR_SECONDS,
+ AdminPOPSchema,
)
from generalresearch.locales import Localelator
@@ -72,8 +72,7 @@ class TestAdminPOPSchema:
def test_index_tz_parser(self):
tz_dates = [
- datetime(year=2024, month=1, day=i, tzinfo=timezone.utc)
- for i in range(1, 10)
+ datetime(year=2024, month=1, day=i, tzinfo=UTC) for i in range(1, 10)
]
df = pd.DataFrame(
@@ -85,16 +84,16 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
# Initially, they're all set with a timezone
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz == timezone.utc for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all([ts.tz == UTC for ts in timestmaps])
# After validation, the timezone is removed
df = AdminPOPSchema.validate(df)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
assert all([ts.tz is None for ts in timestmaps])
def test_index_tz_no_future_beyond_one_year(self):
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
tz_dates = [now + timedelta(days=i * 365) for i in range(1, 10)]
df = pd.DataFrame(
@@ -157,9 +156,7 @@ class TestAdminPOPSchema:
def test_invalid_parsing(self):
# (1) Timezones AND as strings will still parse correctly
tz_str_dates = [
- datetime(
- year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc
- ).isoformat()
+ datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC).isoformat()
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -173,12 +170,12 @@ class TestAdminPOPSchema:
df = AdminPOPSchema.validate(df, lazy=True)
assert isinstance(df, pd.DataFrame)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
assert all([ts.tz is None for ts in timestmaps])
# (2) Timezones are removed
dates = [
- datetime(year=2024, month=1, day=1, minute=i, tzinfo=timezone.utc)
+ datetime(year=2024, month=1, day=1, minute=i, tzinfo=UTC)
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -190,12 +187,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
# Has tz before validation, and none after
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is timezone.utc for ts in timestmaps])
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ assert all([ts.tz is UTC for ts in timestmaps])
df = AdminPOPSchema.validate(df, lazy=True)
- timestmaps: List[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
+ timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
assert all([ts.tz is None for ts in timestmaps])
def test_clipping(self):
--
cgit v1.2.3
From d48f04ab7f034b9b2039c42cd616d7d7851817b9 Mon Sep 17 00:00:00 2001
From: Max Nanis
Date: Fri, 21 Aug 2026 18:04:10 -0700
Subject: "non-safe" Ruff 3.14 changes by hand.
---
generalresearch/currency.py | 4 +--
generalresearch/grliq/managers/event_plotter.py | 1 -
generalresearch/grliq/managers/forensic_events.py | 34 ++++++++--------------
generalresearch/grliq/managers/forensic_results.py | 4 +--
generalresearch/grliq/managers/forensic_summary.py | 4 +--
generalresearch/grliq/models/__init__.py | 4 +--
generalresearch/grliq/models/decider.py | 8 ++---
generalresearch/grliq/models/events.py | 3 +-
generalresearch/grliq/models/forensic_data.py | 13 ++++-----
generalresearch/grliq/models/forensic_result.py | 4 +--
generalresearch/grliq/models/forensic_summary.py | 11 +++----
generalresearch/grliq/models/useragents.py | 12 ++++----
generalresearch/incite/collections/__init__.py | 4 +--
generalresearch/incite/mergers/__init__.py | 10 +++----
generalresearch/incite/mergers/ym_wall_summary.py | 2 +-
generalresearch/incite/schemas/__init__.py | 2 --
.../mergers/foundations/enriched_task_adjust.py | 2 --
generalresearch/locales/__init__.py | 1 -
generalresearch/managers/thl/wallet/__init__.py | 2 +-
generalresearch/managers/thl/wallet/tango.py | 2 +-
generalresearch/models/__init__.py | 16 +++++-----
generalresearch/models/cint/question.py | 6 ++--
generalresearch/models/cint/survey.py | 4 +--
generalresearch/models/cint/task_collection.py | 2 --
generalresearch/models/dynata/__init__.py | 4 +--
generalresearch/models/dynata/question.py | 4 +--
generalresearch/models/dynata/survey.py | 4 +--
generalresearch/models/events.py | 11 +++----
generalresearch/models/gr/business.py | 6 ++--
generalresearch/models/innovate/__init__.py | 8 ++---
generalresearch/models/innovate/question.py | 4 +--
generalresearch/models/innovate/survey.py | 3 +-
generalresearch/models/legacy/definitions.py | 4 +--
generalresearch/models/lucid/question.py | 4 +--
generalresearch/models/morning/__init__.py | 4 +--
generalresearch/models/morning/question.py | 6 ++--
generalresearch/models/morning/survey.py | 8 +----
generalresearch/models/network/nmap/result.py | 2 +-
generalresearch/models/pollfish/question.py | 4 +--
generalresearch/models/precision/__init__.py | 4 +--
generalresearch/models/precision/question.py | 4 +--
generalresearch/models/precision/survey.py | 4 +--
.../models/precision/task_collection.py | 2 +-
generalresearch/models/prodege/__init__.py | 6 ++--
generalresearch/models/prodege/question.py | 6 ++--
generalresearch/models/prodege/survey.py | 4 +--
generalresearch/models/prodege/task_collection.py | 2 +-
generalresearch/models/repdata/__init__.py | 4 +--
generalresearch/models/repdata/question.py | 4 +--
generalresearch/models/repdata/survey.py | 4 +--
generalresearch/models/sago/__init__.py | 4 +--
generalresearch/models/sago/question.py | 4 +--
generalresearch/models/sago/survey.py | 4 +--
generalresearch/models/spectrum/question.py | 8 ++---
generalresearch/models/spectrum/survey.py | 4 +--
generalresearch/models/thl/contest/definitions.py | 16 +++++-----
generalresearch/models/thl/definitions.py | 22 +++++++-------
generalresearch/models/thl/demographics.py | 4 +--
generalresearch/models/thl/leaderboard.py | 8 ++---
generalresearch/models/thl/ledger.py | 14 ++++-----
generalresearch/models/thl/offerwall/__init__.py | 6 ++--
generalresearch/models/thl/product.py | 6 ++--
.../models/thl/profiling/upk_property.py | 6 ++--
.../models/thl/profiling/upk_question.py | 8 ++---
generalresearch/models/thl/supplier_tag.py | 4 +--
generalresearch/models/thl/survey/__init__.py | 1 -
generalresearch/models/thl/user_quality_event.py | 6 ++--
generalresearch/models/thl/user_streak.py | 10 +++----
generalresearch/models/thl/userhealth.py | 4 +--
generalresearch/models/thl/wallet/__init__.py | 6 ++--
.../models/thl/wallet/cashout_method.py | 19 +++++-------
generalresearch/utils/aggregation.py | 2 +-
generalresearch/utils/enum.py | 1 -
generalresearch/wall_status_codes/__init__.py | 2 --
generalresearch/wall_status_codes/cint.py | 2 +-
generalresearch/wall_status_codes/dynata.py | 2 +-
generalresearch/wall_status_codes/fullcircle.py | 2 +-
generalresearch/wall_status_codes/innovate.py | 2 +-
generalresearch/wall_status_codes/lucid.py | 2 +-
generalresearch/wall_status_codes/morning.py | 2 +-
generalresearch/wall_status_codes/precision.py | 2 +-
generalresearch/wall_status_codes/prodege.py | 2 +-
generalresearch/wall_status_codes/repdata.py | 2 +-
generalresearch/wall_status_codes/sago.py | 2 +-
generalresearch/wall_status_codes/spectrum.py | 2 +-
generalresearch/wall_status_codes/wxet.py | 1 -
generalresearch/wxet/models/definitions.py | 15 +++++-----
generalresearch/wxet/models/finish_type.py | 5 ++--
tests/incite/schemas/test_admin_responses.py | 3 +-
89 files changed, 222 insertions(+), 268 deletions(-)
(limited to 'tests/incite/schemas/test_admin_responses.py')
diff --git a/generalresearch/currency.py b/generalresearch/currency.py
index c76b266..9402e04 100644
--- a/generalresearch/currency.py
+++ b/generalresearch/currency.py
@@ -1,6 +1,6 @@
import warnings
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from typing import Any
from pydantic import GetCoreSchemaHandler, NonNegativeInt
@@ -9,7 +9,7 @@ from pydantic_core import CoreSchema, core_schema
from generalresearch.utils.enum import ReprEnumMeta
-class LedgerCurrency(str, Enum, metaclass=ReprEnumMeta):
+class LedgerCurrency(StrEnum, metaclass=ReprEnumMeta):
USD = "USD"
USDCent = "USDCent"
USDMill = "USDMill"
diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py
index ed01d2e..0bba7c5 100644
--- a/generalresearch/grliq/managers/event_plotter.py
+++ b/generalresearch/grliq/managers/event_plotter.py
@@ -1,6 +1,5 @@
import html
import webbrowser
-from typing import List
import numpy as np
from more_itertools import windowed
diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py
index 85e9620..c847d4d 100644
--- a/generalresearch/grliq/managers/forensic_events.py
+++ b/generalresearch/grliq/managers/forensic_events.py
@@ -1,7 +1,7 @@
import json
-from datetime import datetime
-from typing import Any, Dict, List, Optional
from collections.abc import Collection
+from datetime import datetime
+from typing import Any
from uuid import uuid4
from psycopg import sql
@@ -40,15 +40,13 @@ class GrlIqEventManager:
with conn.cursor() as c:
c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
# Try to update first
- update_query = sql.SQL(
- """
+ update_query = sql.SQL("""
UPDATE grliq_forensicevents
SET timing_data = %(timing_data)s
WHERE session_uuid = %(session_uuid)s
AND timing_data IS NULL
RETURNING id
- """
- )
+ """)
c.execute(update_query, data)
result = c.fetchone()
@@ -58,15 +56,13 @@ class GrlIqEventManager:
return pk
# No matching row to update. Do an insert
- insert_query = sql.SQL(
- """
+ insert_query = sql.SQL("""
INSERT INTO grliq_forensicevents
(uuid, session_uuid, timing_data)
VALUES
(%(uuid)s, %(session_uuid)s, %(timing_data)s)
RETURNING id
- """
- )
+ """)
c.execute(insert_query, data)
pk = c.fetchone()["id"]
conn.commit()
@@ -96,8 +92,7 @@ class GrlIqEventManager:
with conn.cursor() as c:
c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,))
# Try to update first
- update_query = sql.SQL(
- """
+ update_query = sql.SQL("""
UPDATE grliq_forensicevents
SET events = %(events)s,
mouse_events = %(mouse_events)s,
@@ -106,8 +101,7 @@ class GrlIqEventManager:
WHERE session_uuid = %(session_uuid)s
AND events IS NULL
RETURNING id
- """
- )
+ """)
c.execute(update_query, data)
result = c.fetchone()
@@ -117,8 +111,7 @@ class GrlIqEventManager:
return pk
# No matching row to update. Do an insert
- insert_query = sql.SQL(
- """
+ insert_query = sql.SQL("""
INSERT INTO grliq_forensicevents
(uuid, session_uuid, events, mouse_events,
event_start, event_end)
@@ -126,8 +119,7 @@ class GrlIqEventManager:
(%(uuid)s, %(session_uuid)s, %(events)s, %(mouse_events)s,
%(event_start)s, %(event_end)s)
RETURNING id
- """
- )
+ """)
c.execute(insert_query, data)
pk = c.fetchone()["id"]
conn.commit()
@@ -202,8 +194,7 @@ class GrlIqEventManager:
session_uuids: Collection[str],
) -> list[dict[str, Any]]:
params = {"session_uuids": list(session_uuids)}
- query = sql.SQL(
- """
+ query = sql.SQL("""
SELECT DISTINCT ON (fe.session_uuid)
timing_data,
fe.session_uuid,
@@ -214,8 +205,7 @@ class GrlIqEventManager:
WHERE fe.session_uuid = ANY(%(session_uuids)s)
AND timing_data IS NOT NULL
ORDER BY session_uuid, fe.id DESC;
- """
- )
+ """)
with self.postgres_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query, params)
diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py
index 52bde99..706a7db 100644
--- a/generalresearch/grliq/managers/forensic_results.py
+++ b/generalresearch/grliq/managers/forensic_results.py
@@ -1,6 +1,6 @@
-from datetime import datetime
-from typing import Any, Dict, List, Optional, Tuple
from collections.abc import Collection
+from datetime import datetime
+from typing import Any
from generalresearch.grliq.models.forensic_result import (
GrlIqForensicCategoryResult,
diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py
index 5039a38..98e00d8 100644
--- a/generalresearch/grliq/managers/forensic_summary.py
+++ b/generalresearch/grliq/managers/forensic_summary.py
@@ -2,8 +2,8 @@ from __future__ import annotations
import statistics
from collections import defaultdict
-from datetime import datetime, timedelta, timezone, UTC
-from typing import Any, Dict, List
+from datetime import UTC, datetime, timedelta
+from typing import Any
import numpy as np
diff --git a/generalresearch/grliq/models/__init__.py b/generalresearch/grliq/models/__init__.py
index 998de2d..fabbca2 100644
--- a/generalresearch/grliq/models/__init__.py
+++ b/generalresearch/grliq/models/__init__.py
@@ -1,10 +1,10 @@
from __future__ import annotations
import json
-from enum import Enum
+from enum import StrEnum
-class RiskWeighting(str, Enum):
+class RiskWeighting(StrEnum):
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
diff --git a/generalresearch/grliq/models/decider.py b/generalresearch/grliq/models/decider.py
index d24e150..579c39f 100644
--- a/generalresearch/grliq/models/decider.py
+++ b/generalresearch/grliq/models/decider.py
@@ -1,14 +1,14 @@
from __future__ import annotations
-from datetime import datetime, timezone, UTC
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field
from generalresearch.models.custom_types import AwareDatetimeISO
-class Decider(str, Enum):
+class Decider(StrEnum):
# This decision was made in the thl-core: pre-offerwall-entry view
PRE_ENTRY = "pre_entry"
# This decision made by grl-iq (synchronously)
@@ -17,7 +17,7 @@ class Decider(str, Enum):
YM_USER = "ym_user"
-class AttemptDecision(str, Enum):
+class AttemptDecision(StrEnum):
# This attempt should be allowed to continue
PASS = "pass"
# This attempt is deemed fraudulent
diff --git a/generalresearch/grliq/models/events.py b/generalresearch/grliq/models/events.py
index 5daa032..69b67e5 100644
--- a/generalresearch/grliq/models/events.py
+++ b/generalresearch/grliq/models/events.py
@@ -3,7 +3,7 @@ from __future__ import annotations
from collections import namedtuple
from dataclasses import dataclass, fields
from functools import cached_property
-from typing import Any, Dict, List, Optional
+from typing import Any, Self
import numpy as np
from pydantic import (
@@ -14,7 +14,6 @@ from pydantic import (
NonNegativeInt,
PositiveFloat,
)
-from typing import Self
from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr
diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py
index eda7186..8d65696 100644
--- a/generalresearch/grliq/models/forensic_data.py
+++ b/generalresearch/grliq/models/forensic_data.py
@@ -3,10 +3,10 @@ from __future__ import annotations
import hashlib
import re
from collections import Counter
-from datetime import datetime, timedelta, timezone, UTC
-from enum import Enum
+from datetime import UTC, datetime, timedelta
+from enum import StrEnum
from functools import cached_property
-from typing import Any, Literal
+from typing import Annotated, Any, Literal, Self
from uuid import uuid4
import pycountry
@@ -23,7 +23,6 @@ from pydantic import (
)
from pydantic.json_schema import SkipJsonSchema
from pydantic_extra_types.timezone_name import TimeZoneName
-from typing import Annotated, Self
from generalresearch.grliq.models import (
AUDIO_CODEC_NAMES,
@@ -60,7 +59,7 @@ from generalresearch.models.thl.session import Session
fake = Faker()
-class Platform(str, Enum):
+class Platform(StrEnum):
MAC_INTEL = "MacIntel"
ARM = "ARM"
IPAD = "iPad"
@@ -75,7 +74,7 @@ class Platform(str, Enum):
OTHER = "Other"
-class PassFailError(str, Enum):
+class PassFailError(StrEnum):
PASS = "pass"
FAIL = "fail"
ERROR = "error"
@@ -96,7 +95,7 @@ class PassFailError(str, Enum):
return {2: cls.PASS, 1: cls.FAIL, 0: cls.ERROR, -1: cls.ERROR}[int(v)]
-class SupportLevel(str, Enum):
+class SupportLevel(StrEnum):
# Used for checking if certain features are available in the browser
FULL = "full"
PARTIAL = "partial"
diff --git a/generalresearch/grliq/models/forensic_result.py b/generalresearch/grliq/models/forensic_result.py
index 3fe6481..1fd99cc 100644
--- a/generalresearch/grliq/models/forensic_result.py
+++ b/generalresearch/grliq/models/forensic_result.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from uuid import uuid4
from pydantic import BaseModel, ConfigDict, Field, computed_field
@@ -14,7 +14,7 @@ from generalresearch.grliq.models.decider import (
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-class Phase(str, Enum):
+class Phase(StrEnum):
# The 'phase' of a THL-Session experience. grliq may be collected in
# multiple places multiple times within one session
diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py
index 5ecf1a4..6f80768 100644
--- a/generalresearch/grliq/models/forensic_summary.py
+++ b/generalresearch/grliq/models/forensic_summary.py
@@ -2,13 +2,12 @@ from __future__ import annotations
import random
from typing import (
- List,
Literal,
+ Optional,
Union,
get_args,
get_origin,
get_type_hints,
- Optional,
)
import numpy as np
@@ -57,9 +56,7 @@ class UserForensicSummary(BaseModel):
)
# These must be nullable in case a user has 0 attempts!
- category_result_summary: GrlIqForensicCategorySummary | None= Field(
- default=None
- )
+ category_result_summary: GrlIqForensicCategorySummary | None = Field(default=None)
checker_result_summary: GrlIqCheckerResultsSummary | None = Field(default=None)
country_timing_data_summary: dict[CountryISO, TimingDataCountrySummary] = Field(
@@ -211,13 +208,13 @@ class CountryRTTDistribution(BaseModel):
description="Country client_ip is located in", examples=["fr"]
)
# For users marked as fraud or not
- is_fraud: bool| None = Field(
+ is_fraud: bool | None = Field(
default=None,
description="If timing data from sessions determined to be fraud are included",
)
# we could split by this optionally
- user_type: UserType|None = Field(
+ user_type: UserType | None = Field(
default=None,
description="user_type of the client_ip as determined by MaxMind",
examples=[UserType.RESIDENTIAL],
diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py
index 4bb340e..de63a67 100644
--- a/generalresearch/grliq/models/useragents.py
+++ b/generalresearch/grliq/models/useragents.py
@@ -1,15 +1,15 @@
from __future__ import annotations
import hashlib
-from enum import Enum
+from enum import StrEnum
+from typing import Self
from pydantic import BaseModel, ConfigDict, Field, field_validator
-from typing import Self
from user_agents import parse as ua_parse
from user_agents.parsers import UserAgent
-class BrowserFamily(str, Enum):
+class BrowserFamily(StrEnum):
CHROME_MOBILE = "Chrome Mobile"
CHROME = "Chrome"
CHROME_MOBILE_WEBVIEW = "Chrome Mobile WebView"
@@ -32,7 +32,7 @@ class BrowserFamily(str, Enum):
OTHER = "Other"
-class OSFamily(str, Enum):
+class OSFamily(StrEnum):
ANDROID = "Android"
WINDOWS = "Windows"
IOS = "iOS"
@@ -43,7 +43,7 @@ class OSFamily(str, Enum):
OTHER = "Other"
-class DeviceBrand(str, Enum):
+class DeviceBrand(StrEnum):
GENERIC_ANDROID = "Generic_Android"
NONE = "None"
APPLE = "Apple"
@@ -65,7 +65,7 @@ class DeviceBrand(str, Enum):
OTHER = "Other"
-class DeviceModelFamily(str, Enum):
+class DeviceModelFamily(StrEnum):
NONE = "None"
OTHER = "Other"
K = "K"
diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py
index 17c119a..38749b3 100644
--- a/generalresearch/incite/collections/__init__.py
+++ b/generalresearch/incite/collections/__init__.py
@@ -5,7 +5,7 @@ import os
import subprocess
import time
from datetime import datetime
-from enum import Enum
+from enum import StrEnum
from sys import platform
from typing import Any
@@ -56,7 +56,7 @@ LOG = logging.getLogger("incite")
DT_STR = "%Y-%m-%d %H:%M:%S"
-class DFCollectionType(str, Enum):
+class DFCollectionType(StrEnum):
TEST = "test"
USER = "thl_user"
diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py
index 9c5276c..45d3e2d 100644
--- a/generalresearch/incite/mergers/__init__.py
+++ b/generalresearch/incite/mergers/__init__.py
@@ -1,17 +1,16 @@
import logging
import os.path
import subprocess
-from datetime import datetime, timezone, UTC
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from sys import platform
-from typing import List, Optional, Type
+from typing import Self
import dask.dataframe as dd
import pandas as pd
from dask.distributed import Client
from pandera.pandas import DataFrameSchema
from pydantic import Field, ValidationInfo, field_validator, model_validator
-from typing import Self
from generalresearch.incite.base import CollectionBase, CollectionItemBase
from generalresearch.incite.schemas import PARTITION_ON
@@ -27,7 +26,6 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_wall import (
from generalresearch.incite.schemas.mergers.foundations.user_id_product import (
UserIdProductSchema,
)
-
from generalresearch.incite.schemas.mergers.pop_ledger import (
PopLedgerSchema,
)
@@ -42,7 +40,7 @@ from generalresearch.models.custom_types import AwareDatetimeISO
LOG = logging.getLogger("incite")
-class MergeType(str, Enum):
+class MergeType(StrEnum):
TEST = "test"
YM_SURVEY_WALL = "ym_survey_wall"
YM_WALL_SUMMARY = "ym_wall_summary"
diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py
index 4816c05..618b810 100644
--- a/generalresearch/incite/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/mergers/ym_wall_summary.py
@@ -1,7 +1,7 @@
from __future__ import annotations
from datetime import datetime, time, timedelta
-from typing import Literal, Type
+from typing import Literal
import dask.dataframe as dd
import pandas as pd
diff --git a/generalresearch/incite/schemas/__init__.py b/generalresearch/incite/schemas/__init__.py
index 6fc83b0..9d24f17 100644
--- a/generalresearch/incite/schemas/__init__.py
+++ b/generalresearch/incite/schemas/__init__.py
@@ -1,5 +1,3 @@
-from typing import List
-
import pandas as pd
import pandera.pandas as pa
diff --git a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
index 97e73a3..ead42d9 100644
--- a/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
+++ b/generalresearch/incite/schemas/mergers/foundations/enriched_task_adjust.py
@@ -1,5 +1,3 @@
-from typing import Set
-
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
diff --git a/generalresearch/locales/__init__.py b/generalresearch/locales/__init__.py
index 813966e..38c1832 100644
--- a/generalresearch/locales/__init__.py
+++ b/generalresearch/locales/__init__.py
@@ -12,7 +12,6 @@ https://en.wikipedia.org/wiki/List_of_ISO_639-1_codes
import json
import pkgutil
-from typing import Set
class Localelator:
diff --git a/generalresearch/managers/thl/wallet/__init__.py b/generalresearch/managers/thl/wallet/__init__.py
index 9f3ae85..05e700d 100644
--- a/generalresearch/managers/thl/wallet/__init__.py
+++ b/generalresearch/managers/thl/wallet/__init__.py
@@ -1,5 +1,5 @@
from decimal import Decimal
-from typing import Any, Dict, Optional, Union
+from typing import Any
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py
index 445719d..a4ebf22 100644
--- a/generalresearch/managers/thl/wallet/tango.py
+++ b/generalresearch/managers/thl/wallet/tango.py
@@ -1,4 +1,4 @@
-from typing import Any, Dict
+from typing import Any
from generalresearch.config import (
is_debug,
diff --git a/generalresearch/models/__init__.py b/generalresearch/models/__init__.py
index 9a2bb9c..c0348d7 100644
--- a/generalresearch/models/__init__.py
+++ b/generalresearch/models/__init__.py
@@ -1,11 +1,11 @@
from __future__ import annotations
-from enum import Enum
+from enum import IntEnum, StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class Source(str, Enum, metaclass=ReprEnumMeta):
+class Source(StrEnum, metaclass=ReprEnumMeta):
# The external marketplace, or the source of the survey / work.
# Max length of the value is 2.
GRS = "g"
@@ -31,7 +31,7 @@ class Source(str, Enum, metaclass=ReprEnumMeta):
WXET = "w"
-class DebitKey(int, Enum, metaclass=ReprEnumMeta):
+class DebitKey(IntEnum, metaclass=ReprEnumMeta):
# The debit key for marketplaces
CINT = 8
DALIA = 9
@@ -50,21 +50,21 @@ class DebitKey(int, Enum, metaclass=ReprEnumMeta):
# WXET = None
-class DeviceType(int, Enum, metaclass=ReprEnumMeta):
+class DeviceType(IntEnum, metaclass=ReprEnumMeta):
UNKNOWN = 0
MOBILE = 1
DESKTOP = 2
TABLET = 3
-class LogicalOperator(str, Enum, metaclass=ReprEnumMeta):
+class LogicalOperator(StrEnum, metaclass=ReprEnumMeta):
OR = "OR"
AND = "AND"
# There is currently no use case for NOT. See MarketplaceCondition.explain_not
NOT = "NOT"
-class TaskStatus(str, Enum, metaclass=ReprEnumMeta):
+class TaskStatus(StrEnum, metaclass=ReprEnumMeta):
# A survey is live if it is open and, given all conditions are met, is
# possible to send in traffic. All other statuses are just variants of
# NOT Live (not accepting traffic)
@@ -80,7 +80,7 @@ class TaskStatus(str, Enum, metaclass=ReprEnumMeta):
NOT_FOUND = "NOT_FOUND"
-class TaskCalculationType(str, Enum):
+class TaskCalculationType(StrEnum):
COMPLETES = "COMPLETES"
STARTS = "STARTS"
@@ -105,7 +105,7 @@ class TaskCalculationType(str, Enum):
return {0: cls.COMPLETES, 1: cls.STARTS}[v]
-class URLQueryKey(str, Enum, metaclass=ReprEnumMeta):
+class URLQueryKey(StrEnum, metaclass=ReprEnumMeta):
PRODUCT_ID = "39057c8b"
PRODUCT_USER_ID = "c184efc0"
SESSION_ID = "0bb50182"
diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py
index 5959141..f8a287a 100644
--- a/generalresearch/models/cint/question.py
+++ b/generalresearch/models/cint/question.py
@@ -1,8 +1,8 @@
from __future__ import annotations
import json
-from datetime import UTC, datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal, Self
from uuid import UUID
@@ -22,7 +22,7 @@ if TYPE_CHECKING:
)
-class CintQuestionType(str, Enum):
+class CintQuestionType(StrEnum):
SINGLE_SELECT = "s"
MULTI_SELECT = "m"
# Dummy means they're calculated
diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py
index 01615f6..cfc91ef 100644
--- a/generalresearch/models/cint/survey.py
+++ b/generalresearch/models/cint/survey.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
-from typing import Annotated, Any, Literal, Self, Type
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
from pydantic import (
diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py
index efb6e95..4ae8de4 100644
--- a/generalresearch/models/cint/task_collection.py
+++ b/generalresearch/models/cint/task_collection.py
@@ -1,7 +1,5 @@
from __future__ import annotations
-from typing import List
-
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
diff --git a/generalresearch/models/dynata/__init__.py b/generalresearch/models/dynata/__init__.py
index c6d3a67..0f3dfe7 100644
--- a/generalresearch/models/dynata/__init__.py
+++ b/generalresearch/models/dynata/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class DynataStatus(str, Enum):
+class DynataStatus(StrEnum):
OPEN = "OPEN"
PAUSED = "PAUSED"
CLOSED = "CLOSED"
diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py
index 7288588..b95f7c6 100644
--- a/generalresearch/models/dynata/question.py
+++ b/generalresearch/models/dynata/question.py
@@ -5,7 +5,7 @@ import json
import logging
import re
from datetime import timedelta
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from typing import Any, Literal
@@ -51,7 +51,7 @@ class DynataQuestionOption(BaseModel):
return clean_text(s)
-class DynataQuestionType(str, Enum):
+class DynataQuestionType(StrEnum):
"""
From the API: {'geo', 'multi_select', 'multi_select_searchable', 'none',
'single_select', 'single_select_grid', 'single_select_searchable', 'zip'}
diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py
index 097eea2..9adac42 100644
--- a/generalresearch/models/dynata/survey.py
+++ b/generalresearch/models/dynata/survey.py
@@ -2,10 +2,10 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal, Self, Type
+from typing import Any, Literal, Self
from more_itertools import flatten
from pydantic import (
diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py
index 5efd4f6..70f699b 100644
--- a/generalresearch/models/events.py
+++ b/generalresearch/models/events.py
@@ -1,6 +1,6 @@
-from datetime import UTC, datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from enum import StrEnum
-from typing import Annotated, Dict, Literal, Optional, Union
+from typing import Annotated, Literal
from uuid import uuid4
from pydantic import (
@@ -263,17 +263,14 @@ class SubscribeMessage(BaseModel):
product_id: UUIDStr = Field(examples=["4fe381fb7186416cb443a38fa66c6557"])
-ServerToClientMessage = Union[EventMessage, StatsMessage, PingMessage]
+ServerToClientMessage = EventMessage | StatsMessage | PingMessage
ServerToClientMessageField = Annotated[
ServerToClientMessage,
Field(discriminator="kind"),
]
ServerToClientMessageAdapter = TypeAdapter(ServerToClientMessageField)
-ClientToServerMessage = Union[
- SubscribeMessage,
- PongMessage,
-]
+ClientToServerMessage = SubscribeMessage | PongMessage
ClientToServerMessageField = Annotated[
ClientToServerMessage,
Field(discriminator="kind"),
diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py
index a67cf48..534b23f 100644
--- a/generalresearch/models/gr/business.py
+++ b/generalresearch/models/gr/business.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import json
import logging
import os
-from datetime import UTC, datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import Enum, StrEnum
from pathlib import Path
from typing import TYPE_CHECKING, Self
from uuid import uuid4
@@ -63,7 +63,7 @@ class TransferMethod(Enum, metaclass=ReprEnumMeta):
WIRE = 1
-class BusinessType(str, Enum, metaclass=ReprEnumMeta):
+class BusinessType(StrEnum, metaclass=ReprEnumMeta):
INDIVIDUAL = "i"
COMPANY = "c"
diff --git a/generalresearch/models/innovate/__init__.py b/generalresearch/models/innovate/__init__.py
index 26946e9..9d7ea75 100644
--- a/generalresearch/models/innovate/__init__.py
+++ b/generalresearch/models/innovate/__init__.py
@@ -1,4 +1,4 @@
-from enum import Enum
+from enum import StrEnum
from typing import Annotated
from pydantic import StringConstraints
@@ -9,17 +9,17 @@ InnovateQuestionID = Annotated[
]
-class InnovateStatus(str, Enum):
+class InnovateStatus(StrEnum):
LIVE = "LIVE"
NOT_LIVE = "NOT_LIVE"
-class InnovateQuotaStatus(str, Enum):
+class InnovateQuotaStatus(StrEnum):
OPEN = "OPEN"
CLOSED = "CLOSED"
-class InnovateDuplicateCheckLevel(str, Enum):
+class InnovateDuplicateCheckLevel(StrEnum):
# How we should check for de-dupes / survey exclusions.
# https://innovatemr.stoplight.io/docs/supplier-api/ZG9jOjEzNzYxMTg2-statuses-term-reasons-and-categories
# #duplicatedtoken
diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py
index 306274f..89310a2 100644
--- a/generalresearch/models/innovate/question.py
+++ b/generalresearch/models/innovate/question.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import json
import logging
-from enum import Enum
+from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -50,7 +50,7 @@ class InnovateQuestionOption(BaseModel):
order: int = Field()
-class InnovateQuestionType(str, Enum):
+class InnovateQuestionType(StrEnum):
# API response: {'Multipunch', 'Numeric Open Ended', 'Single Punch'}
# "Numeric Open Ended" must be wrong... It can't be numeric, as UK's
# postcode question is marked as this, but it wants alphanumeric
diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py
index 0359bd6..3c37fe3 100644
--- a/generalresearch/models/innovate/survey.py
+++ b/generalresearch/models/innovate/survey.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, date, timezone
+from datetime import UTC, date
from decimal import Decimal
from functools import cached_property
from typing import (
@@ -10,7 +10,6 @@ from typing import (
Any,
Literal,
Self,
- Type,
)
from more_itertools import flatten
diff --git a/generalresearch/models/legacy/definitions.py b/generalresearch/models/legacy/definitions.py
index 1755d2a..b3a6687 100644
--- a/generalresearch/models/legacy/definitions.py
+++ b/generalresearch/models/legacy/definitions.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class OfferwallReason(str, Enum):
+class OfferwallReason(StrEnum):
USER_BLOCKED = "USER_BLOCKED"
HIGH_RECON_RATE = "HIGH_RECON_RATE"
UNCOMMON_DEMOGRAPHICS = "UNCOMMON_DEMOGRAPHICS"
diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py
index 6cbcb73..af3420e 100644
--- a/generalresearch/models/lucid/question.py
+++ b/generalresearch/models/lucid/question.py
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
-from enum import Enum
+from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal, Self
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -40,7 +40,7 @@ class LucidQuestionOption(BaseModel):
order: int = Field()
-class LucidQuestionType(str, Enum):
+class LucidQuestionType(StrEnum):
SINGLE_SELECT = "s"
MULTI_SELECT = "m"
TEXT_ENTRY = "t"
diff --git a/generalresearch/models/morning/__init__.py b/generalresearch/models/morning/__init__.py
index 1bc15a7..e747586 100644
--- a/generalresearch/models/morning/__init__.py
+++ b/generalresearch/models/morning/__init__.py
@@ -1,4 +1,4 @@
-from enum import Enum
+from enum import StrEnum
from typing import Annotated
from pydantic import StringConstraints
@@ -9,7 +9,7 @@ MorningQuestionID = Annotated[
]
-class MorningStatus(str, Enum):
+class MorningStatus(StrEnum):
DRAFT = "draft"
ACTIVE = "active" # aka LIVE
PAUSED = "paused"
diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py
index 8a1f729..b64a44a 100644
--- a/generalresearch/models/morning/question.py
+++ b/generalresearch/models/morning/question.py
@@ -1,6 +1,6 @@
import json
-from enum import Enum
-from typing import Any, Dict, List, Literal, Optional, Self
+from enum import StrEnum
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -36,7 +36,7 @@ class MorningQuestionOption(BaseModel, frozen=True):
order: int = Field()
-class MorningQuestionType(str, Enum):
+class MorningQuestionType(StrEnum):
# The db stores these as a single letter
# Geographic questions represent geographic areas within a country.
diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py
index 3c5a0e0..5011698 100644
--- a/generalresearch/models/morning/survey.py
+++ b/generalresearch/models/morning/survey.py
@@ -2,20 +2,14 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
from typing import (
Annotated,
Any,
- Dict,
- List,
Literal,
- Optional,
Self,
- Set,
- Tuple,
- Type,
)
from pydantic import (
diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py
index e6a0fd3..55c2109 100644
--- a/generalresearch/models/network/nmap/result.py
+++ b/generalresearch/models/network/nmap/result.py
@@ -4,7 +4,7 @@ import json
from datetime import timedelta
from enum import StrEnum
from functools import cached_property
-from typing import Any, Literal, Set
+from typing import Any, Literal
from pydantic import BaseModel, Field, computed_field
diff --git a/generalresearch/models/pollfish/question.py b/generalresearch/models/pollfish/question.py
index ae0dc5a..3b658fd 100644
--- a/generalresearch/models/pollfish/question.py
+++ b/generalresearch/models/pollfish/question.py
@@ -3,7 +3,7 @@ from __future__ import annotations
# https://wss.pollfish.com/mediation/documentation
import json
import logging
-from enum import Enum
+from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal, Self
from pydantic import BaseModel, Field, model_validator
@@ -39,7 +39,7 @@ class PollfishQuestionOption(BaseModel):
order: int = Field()
-class PollfishQuestionType(str, Enum):
+class PollfishQuestionType(StrEnum):
"""
From the API: {'single_punch', 'multi_punch', 'open_ended'}
"""
diff --git a/generalresearch/models/precision/__init__.py b/generalresearch/models/precision/__init__.py
index 089c34a..07deb8e 100644
--- a/generalresearch/models/precision/__init__.py
+++ b/generalresearch/models/precision/__init__.py
@@ -1,10 +1,10 @@
-from enum import Enum
+from enum import StrEnum
from typing import Annotated
from pydantic import StringConstraints
-class PrecisionStatus(str, Enum):
+class PrecisionStatus(StrEnum):
# I made this up. They use isactive: "Yes" or "no", which I think is stupid
OPEN = "open"
CLOSED = "closed"
diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py
index 4030509..a2189d5 100644
--- a/generalresearch/models/precision/question.py
+++ b/generalresearch/models/precision/question.py
@@ -3,7 +3,7 @@ from __future__ import annotations
# https://integrations.precisionsample.com/api.html#Get%20Questions
import json
import logging
-from enum import Enum
+from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -43,7 +43,7 @@ class PrecisionQuestionOption(BaseModel):
order: int = Field()
-class PrecisionQuestionType(str, Enum):
+class PrecisionQuestionType(StrEnum):
"""
From the API: {'Drop Down', 'Multi Select', 'Single Select', 'Single Select Matrix', 'Vertical Question'}
Of course undocumented. And there doesn't seem to be a text entry option?
diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py
index be98a79..f515552 100644
--- a/generalresearch/models/precision/survey.py
+++ b/generalresearch/models/precision/survey.py
@@ -1,9 +1,9 @@
from __future__ import annotations
import json
-from datetime import UTC, timezone
+from datetime import UTC
from functools import cached_property
-from typing import Annotated, Any, Dict, List, Literal, Optional, Self, Set, Tuple, Type
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
from pydantic import (
diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py
index daea448..c8db2af 100644
--- a/generalresearch/models/precision/task_collection.py
+++ b/generalresearch/models/precision/task_collection.py
@@ -1,4 +1,4 @@
-from typing import Any, Dict, List
+from typing import Any
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
diff --git a/generalresearch/models/prodege/__init__.py b/generalresearch/models/prodege/__init__.py
index 5c6659a..64b148a 100644
--- a/generalresearch/models/prodege/__init__.py
+++ b/generalresearch/models/prodege/__init__.py
@@ -1,4 +1,4 @@
-from enum import Enum
+from enum import StrEnum
from typing import Annotated, Literal
from pydantic import Field
@@ -8,7 +8,7 @@ ProdegeQuestionIdType = Annotated[
]
-class ProdegeStatus(str, Enum):
+class ProdegeStatus(StrEnum):
LIVE = "LIVE"
# We need another status to mark if a survey we thought was live does not come back
# from the API, we'll mark it as NOT_FOUND
@@ -18,7 +18,7 @@ class ProdegeStatus(str, Enum):
INELIGIBLE = "INELIGIBLE"
-class ProdegePastParticipationType(str, Enum):
+class ProdegePastParticipationType(StrEnum):
# These come from the "participation_types" key in the survey API response
# which is how we filter by users' past_participation.
CLICK = "click"
diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py
index 3ef4772..574c4fd 100644
--- a/generalresearch/models/prodege/question.py
+++ b/generalresearch/models/prodege/question.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import StrEnum
from functools import cached_property
from typing import TYPE_CHECKING, Any, Literal
@@ -90,7 +90,7 @@ class ProdegeQuestionOption(BaseModel):
is_exclusive: bool = Field(default=False)
-class ProdegeQuestionType(str, Enum):
+class ProdegeQuestionType(StrEnum):
"""
{'Derived', 'Multi Punch', 'Numeric - Open End', 'Single Punch', 'Zip Code'}
"""
diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py
index 1601fa3..3f4c88f 100644
--- a/generalresearch/models/prodege/survey.py
+++ b/generalresearch/models/prodege/survey.py
@@ -4,10 +4,10 @@ from __future__ import annotations
import json
import logging
from collections import defaultdict
-from datetime import UTC, datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal, Type
+from typing import Any, Literal
from pydantic import (
BaseModel,
diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py
index d3e4a20..4544050 100644
--- a/generalresearch/models/prodege/task_collection.py
+++ b/generalresearch/models/prodege/task_collection.py
@@ -1,4 +1,4 @@
-from typing import Any, Dict, List
+from typing import Any
import pandas as pd
from pandera.pandas import Check, Column, DataFrameSchema, Index
diff --git a/generalresearch/models/repdata/__init__.py b/generalresearch/models/repdata/__init__.py
index 9706b28..4a3d8d0 100644
--- a/generalresearch/models/repdata/__init__.py
+++ b/generalresearch/models/repdata/__init__.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class RepDataStatus(str, Enum):
+class RepDataStatus(StrEnum):
LIVE = "LIVE"
DRAFT = "DRAFT"
PAUSED = "PAUSED"
diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py
index 8d0da13..9dda97f 100644
--- a/generalresearch/models/repdata/question.py
+++ b/generalresearch/models/repdata/question.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import json
import logging
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from typing import TYPE_CHECKING, Any, Literal
from uuid import UUID
@@ -86,7 +86,7 @@ class RepDataQuestionOption(BaseModel):
order: int = Field()
-class RepDataQuestionType(str, Enum):
+class RepDataQuestionType(StrEnum):
"""
{'Derived', 'Multi Punch', 'Numeric - Open End', 'Single Punch', 'Zip Code'}
"""
diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py
index fa71c04..43a592c 100644
--- a/generalresearch/models/repdata/survey.py
+++ b/generalresearch/models/repdata/survey.py
@@ -3,10 +3,10 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
from functools import cached_property
-from typing import Any, Literal, Self, Type
+from typing import Any, Literal, Self
from uuid import UUID
from pydantic import (
diff --git a/generalresearch/models/sago/__init__.py b/generalresearch/models/sago/__init__.py
index 19e7b6d..1059a6f 100644
--- a/generalresearch/models/sago/__init__.py
+++ b/generalresearch/models/sago/__init__.py
@@ -1,4 +1,4 @@
-from enum import Enum
+from enum import StrEnum
from typing import Annotated
from pydantic import Field
@@ -8,6 +8,6 @@ SagoQuestionIdType = Annotated[
]
-class SagoStatus(str, Enum):
+class SagoStatus(StrEnum):
LIVE = "LIVE"
NOT_LIVE = "NOT_LIVE"
diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py
index f911854..474543d 100644
--- a/generalresearch/models/sago/question.py
+++ b/generalresearch/models/sago/question.py
@@ -4,7 +4,7 @@ from __future__ import annotations
# -answers-lanaguge-languageid
import json
import logging
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from typing import TYPE_CHECKING, Any, Literal
@@ -57,7 +57,7 @@ class SagoQuestionOption(BaseModel):
return string_utils.remove_nbsp(s)
-class SagoQuestionType(str, Enum):
+class SagoQuestionType(StrEnum):
"""
From the API:
{1: 'Single Punch', 2: 'Multi Punch', 3: 'Open Ended', 4: 'Dummy',
diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py
index 0b5830b..83aad8c 100644
--- a/generalresearch/models/sago/survey.py
+++ b/generalresearch/models/sago/survey.py
@@ -2,10 +2,10 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, timezone
+from datetime import UTC
from decimal import Decimal
from functools import cached_property
-from typing import Annotated, Any, Literal, Self, Type
+from typing import Annotated, Any, Literal, Self
from more_itertools import flatten
from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator
diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py
index 81f8655..c8eea4a 100644
--- a/generalresearch/models/spectrum/question.py
+++ b/generalresearch/models/spectrum/question.py
@@ -3,8 +3,8 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, datetime, timezone
-from enum import Enum
+from datetime import UTC, datetime
+from enum import IntEnum, StrEnum
from functools import cached_property
from typing import TYPE_CHECKING, Any, Literal, Self
from uuid import UUID
@@ -93,7 +93,7 @@ class SpectrumQuestionOption(BaseModel):
return string_utils.remove_nbsp(s)
-class SpectrumQuestionType(str, Enum):
+class SpectrumQuestionType(StrEnum):
# The documentation defines 4 types (1,2,3,4), however 2 is the same as 1
# and never comes back in the api, and we also get back 5, 6, and 7,
# which are all undocumented.
@@ -135,7 +135,7 @@ class SpectrumQuestionType(str, Enum):
return api_type_map[a] if a in api_type_map else None
-class SpectrumQuestionClass(int, Enum):
+class SpectrumQuestionClass(IntEnum):
CORE = 1
EXTENDED = 2
CUSTOM = 3
diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py
index f842f46..aaf182e 100644
--- a/generalresearch/models/spectrum/survey.py
+++ b/generalresearch/models/spectrum/survey.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import json
import logging
-from datetime import UTC, timezone
+from datetime import UTC
from decimal import Decimal
-from typing import Any, Literal, Self, Type
+from typing import Any, Literal, Self
from more_itertools import flatten
from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator
diff --git a/generalresearch/models/thl/contest/definitions.py b/generalresearch/models/thl/contest/definitions.py
index 1a71408..cc0ee05 100644
--- a/generalresearch/models/thl/contest/definitions.py
+++ b/generalresearch/models/thl/contest/definitions.py
@@ -1,17 +1,17 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class ContestStatus(str, Enum):
+class ContestStatus(StrEnum):
ACTIVE = "active"
COMPLETED = "completed"
CANCELLED = "cancelled"
-class ContestType(str, Enum, metaclass=ReprEnumMeta):
+class ContestType(StrEnum, metaclass=ReprEnumMeta):
"""There are 3 contest types. They have a common base, with some unique
configurations and behaviors for each.
"""
@@ -26,7 +26,7 @@ class ContestType(str, Enum, metaclass=ReprEnumMeta):
MILESTONE = "milestone"
-class ContestEndReason(str, Enum):
+class ContestEndReason(StrEnum):
"""
Defines why a contest ended
"""
@@ -44,7 +44,7 @@ class ContestEndReason(str, Enum):
MAX_WINNERS = "max_winners"
-class ContestPrizeKind(str, Enum, metaclass=ReprEnumMeta):
+class ContestPrizeKind(StrEnum, metaclass=ReprEnumMeta):
# A physical prize (e.g. a iPhone, cash in the mail, dinner with Max)
PHYSICAL = "physical"
@@ -56,7 +56,7 @@ class ContestPrizeKind(str, Enum, metaclass=ReprEnumMeta):
CASH = "cash"
-class ContestEntryTrigger(str, Enum):
+class ContestEntryTrigger(StrEnum):
"""
Defines what action/event triggers a (possible) entry into the contest (automatically).
This only is valid on milestone contests
@@ -67,7 +67,7 @@ class ContestEntryTrigger(str, Enum):
REFERRAL = "referral"
-class ContestEntryType(str, Enum, metaclass=ReprEnumMeta):
+class ContestEntryType(StrEnum, metaclass=ReprEnumMeta):
"""
All entries into a contest must be of the same type, and match
the entry_type of the Contest itself.
@@ -83,7 +83,7 @@ class ContestEntryType(str, Enum, metaclass=ReprEnumMeta):
CASH = "cash"
-class LeaderboardTieBreakStrategy(str, Enum):
+class LeaderboardTieBreakStrategy(StrEnum):
"""
Strategies for resolving ties in leaderboard-based contests.
"""
diff --git a/generalresearch/models/thl/definitions.py b/generalresearch/models/thl/definitions.py
index 22ca4b9..0217a80 100644
--- a/generalresearch/models/thl/definitions.py
+++ b/generalresearch/models/thl/definitions.py
@@ -1,10 +1,10 @@
import copy
-from enum import Enum
+from enum import IntEnum, StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class ReservedQueryParameters(str, Enum, metaclass=ReprEnumMeta):
+class ReservedQueryParameters(StrEnum, metaclass=ReprEnumMeta):
PRODUCT_ID = "product_id"
PRODUCT_USER_ID = "bp_user_id"
BPUID = "bpuid"
@@ -48,7 +48,7 @@ class ReservedQueryParameters(str, Enum, metaclass=ReprEnumMeta):
N_BINS = "n_bins"
-class THLPaths(str, Enum, metaclass=ReprEnumMeta):
+class THLPaths(StrEnum, metaclass=ReprEnumMeta):
# Endpoints on thl-fsb
TASK_ADJUSTMENT = "f4484dbdf144451ab60cda256ce14266"
@@ -65,7 +65,7 @@ class THLPaths(str, Enum, metaclass=ReprEnumMeta):
GET_GRLIQ_JS_ATTR = "4a2954b34cc24f93be3e8b218e323b88"
-class Status(str, Enum, metaclass=ReprEnumMeta):
+class Status(StrEnum, metaclass=ReprEnumMeta):
"""
The outcome of a session or wall event. If the session is still in
progress, the status will be NULL.
@@ -85,7 +85,7 @@ class Status(str, Enum, metaclass=ReprEnumMeta):
TIMEOUT = "t"
-class WallAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class WallAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
# Task was reconciled to complete
ADJUSTED_TO_COMPLETE = "ac"
# Task was reconciled to incomplete
@@ -100,7 +100,7 @@ class WallAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
CONFIRMED_COMPLETE = "cc"
-class SessionAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class SessionAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
"""An adjusted_status is set if a session is adjusted by the marketplace
after the original return. A session can be adjusted multiple times.
This is the most recent status. If a session was originally a complete,
@@ -117,7 +117,7 @@ class SessionAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
PAYOUT_ADJUSTMENT = "pa"
-class StatusCode1(int, Enum, metaclass=ReprEnumMeta):
+class StatusCode1(IntEnum, metaclass=ReprEnumMeta):
"""
__High level status code for outcome of the session.__
This should only be NULL if the Status is ABANDON or TIMEOUT
@@ -177,7 +177,7 @@ class StatusCode1(int, Enum, metaclass=ReprEnumMeta):
SESSION_CONTINUE_QUALITY_FAIL = 19
-class SessionStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class SessionStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
__Status Detail__
This should be set if the Session.status_code_1 is SESSION_XXX_FAIL
@@ -217,7 +217,7 @@ class SessionStatusCode2(int, Enum, metaclass=ReprEnumMeta):
GRLIQ_MISSING = 13
-class WallStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class WallStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
This should be set if the Wall.status_code_1 is MARKETPLACE_FAIL
"""
@@ -289,7 +289,7 @@ WALL_ALLOWED_STATUS_CODE_1_2 = {
}
-class ReportValue(int, Enum, metaclass=ReprEnumMeta):
+class ReportValue(IntEnum, metaclass=ReprEnumMeta):
"""
The reason a user reported a task.
"""
@@ -316,7 +316,7 @@ class ReportValue(int, Enum, metaclass=ReprEnumMeta):
DIDNT_LIKE = 7
-class PayoutStatus(str, Enum, metaclass=ReprEnumMeta):
+class PayoutStatus(StrEnum, metaclass=ReprEnumMeta):
"""The max size of the db field that holds this value is 20, so please
don't add new values longer than that!
"""
diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py
index 4d4c8c2..b6a8be1 100644
--- a/generalresearch/models/thl/demographics.py
+++ b/generalresearch/models/thl/demographics.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import copy
from collections import Counter, defaultdict
from dataclasses import dataclass
-from enum import Enum
+from enum import Enum, StrEnum
from typing import TYPE_CHECKING, Any, Literal
import numpy as np
@@ -35,7 +35,7 @@ class DemographicTarget:
}
-class Gender(str, Enum):
+class Gender(StrEnum):
"""
The respondent's gender
"""
diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py
index dce3280..4d116a5 100644
--- a/generalresearch/models/thl/leaderboard.py
+++ b/generalresearch/models/thl/leaderboard.py
@@ -2,8 +2,8 @@ from __future__ import annotations
import logging
import math
-from datetime import UTC, datetime, timedelta, timezone
-from enum import Enum
+from datetime import UTC, datetime, timedelta
+from enum import StrEnum
from typing import Literal
from uuid import UUID, uuid3
from zoneinfo import ZoneInfo
@@ -27,7 +27,7 @@ from generalresearch.utils.enum import ReprEnumMeta
logger = logging.getLogger()
-class LeaderboardCode(str, Enum, metaclass=ReprEnumMeta):
+class LeaderboardCode(StrEnum, metaclass=ReprEnumMeta):
"""
The type of leaderboard. What the "values" represent.
"""
@@ -40,7 +40,7 @@ class LeaderboardCode(str, Enum, metaclass=ReprEnumMeta):
SUM_PAYOUTS = "sum_user_payout"
-class LeaderboardFrequency(str, Enum, metaclass=ReprEnumMeta):
+class LeaderboardFrequency(StrEnum, metaclass=ReprEnumMeta):
"""
The time period range for the leaderboard.
"""
diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py
index dd37d98..19dde20 100644
--- a/generalresearch/models/thl/ledger.py
+++ b/generalresearch/models/thl/ledger.py
@@ -1,8 +1,8 @@
from __future__ import annotations
-from datetime import UTC, datetime, timezone
-from enum import Enum
-from typing import Annotated, Any, Literal, Self, Union
+from datetime import UTC, datetime
+from enum import StrEnum
+from typing import Annotated, Any, Literal, Self
from uuid import uuid4
from pydantic import (
@@ -50,13 +50,13 @@ class Direction(int, Enum, metaclass=ReprEnumMeta):
DEBIT = 1
-class OrderBy(str, Enum, metaclass=ReprEnumMeta):
+class OrderBy(StrEnum, metaclass=ReprEnumMeta):
ASC = "ASC"
DESC = "DESC"
-class AccountType(str, Enum, metaclass=ReprEnumMeta):
+class AccountType(StrEnum, metaclass=ReprEnumMeta):
# Revenue from BP payment commission
BP_COMMISSION = "bp_commission"
# BP wallets (owed balance)
@@ -83,7 +83,7 @@ class AccountType(str, Enum, metaclass=ReprEnumMeta):
WA_CREDIT_LINE = "wa_credit_line"
-class TransactionMetadataColumns(str, Enum):
+class TransactionMetadataColumns(StrEnum):
BONUS = "bonus_id"
# Note: EVENT & EVENT2 represent the same concept. I accidentally made
# this inconsistent.
@@ -102,7 +102,7 @@ class TransactionMetadataColumns(str, Enum):
CONTEST = "contest"
-class TransactionType(str, Enum):
+class TransactionType(StrEnum):
"""These are used in the Ledger to annotate the type of transaction (in
metadata: tx_type)
"""
diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py
index 3acc592..e7c8e03 100644
--- a/generalresearch/models/thl/offerwall/__init__.py
+++ b/generalresearch/models/thl/offerwall/__init__.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import hashlib
import json
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from typing import Any, Literal, Self
from pydantic import (
@@ -31,7 +31,7 @@ from generalresearch.models.thl.product import (
from generalresearch.models.thl.user import User
-class OfferWallType(str, Enum):
+class OfferWallType(StrEnum):
"""
The specific offerwall type
"""
@@ -56,7 +56,7 @@ class OfferWallType(str, Enum):
STARWALL = "b59a2d2b"
-class OfferWallTypeClass(str, Enum):
+class OfferWallTypeClass(StrEnum):
"""
A higher level "class" to organize similar offerwall types.
For e.g. STARWALL_PLUS_BLOCK, STARWALL_PLUS, STARWALL all use the same
diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py
index 8377e45..5b7ab9a 100644
--- a/generalresearch/models/thl/product.py
+++ b/generalresearch/models/thl/product.py
@@ -8,7 +8,7 @@ import warnings
from collections import defaultdict
from collections.abc import Callable
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from functools import cached_property, partial
from typing import (
TYPE_CHECKING,
@@ -628,13 +628,13 @@ class SourceConfig(BaseModel):
)
-class Scope(str, Enum):
+class Scope(StrEnum):
GLOBAL = "global"
TEAM = "team"
PRODUCT = "product"
-class IntegrationMode(str, Enum):
+class IntegrationMode(StrEnum):
# We handle integration, get paid
PLATFORM = "platform"
# "external" credentials, we do not get paid for this activity
diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py
index 9e78692..e46e00a 100644
--- a/generalresearch/models/thl/profiling/upk_property.py
+++ b/generalresearch/models/thl/profiling/upk_property.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
from uuid import uuid4
@@ -11,7 +11,7 @@ from generalresearch.models.thl.category import Category
from generalresearch.utils.enum import ReprEnumMeta
-class PropertyType(str, Enum, metaclass=ReprEnumMeta):
+class PropertyType(StrEnum, metaclass=ReprEnumMeta):
# UserProfileKnowledge Item
UPK_ITEM = "i"
# UserProfileKnowledge Numerical
@@ -25,7 +25,7 @@ class PropertyType(str, Enum, metaclass=ReprEnumMeta):
# UPK_DATE = "d"
-class Cardinality(str, Enum, metaclass=ReprEnumMeta):
+class Cardinality(StrEnum, metaclass=ReprEnumMeta):
# Zero or More
ZERO_OR_MORE = "*"
# Zero or One
diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py
index 78f9511..08ea350 100644
--- a/generalresearch/models/thl/profiling/upk_question.py
+++ b/generalresearch/models/thl/profiling/upk_question.py
@@ -3,9 +3,9 @@ from __future__ import annotations
import hashlib
import json
import re
-from enum import Enum
+from enum import StrEnum
from functools import cached_property
-from typing import Annotated, Any, List, Literal, Union
+from typing import Annotated, Any, Literal
from pydantic import (
BaseModel,
@@ -98,7 +98,7 @@ class UpkQuestionChoiceOut(UpkQuestionChoice):
# importance: Optional[UPKImportance] = Field(default=None, exclude=True)
-class UpkQuestionType(str, Enum):
+class UpkQuestionType(StrEnum):
# The question has options that the user must select from. A MC question
# can be e.g. Selector.SINGLE_ANSWER or Selector.MULTIPLE_ANSWER to
# indicate only 1 or more than 1 option can be selected respectively.
@@ -111,7 +111,7 @@ class UpkQuestionType(str, Enum):
HIDDEN = "HIDDEN"
-class UpkQuestionSelector(str, Enum):
+class UpkQuestionSelector(StrEnum):
pass
diff --git a/generalresearch/models/thl/supplier_tag.py b/generalresearch/models/thl/supplier_tag.py
index ad84c9b..739b895 100644
--- a/generalresearch/models/thl/supplier_tag.py
+++ b/generalresearch/models/thl/supplier_tag.py
@@ -1,7 +1,7 @@
-from enum import Enum
+from enum import StrEnum
-class SupplierTag(str, Enum):
+class SupplierTag(StrEnum):
"""Available tags which can be used to annotate supplier traffic
Note: should not include commas!
diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py
index d6ba895..bc3d699 100644
--- a/generalresearch/models/thl/survey/__init__.py
+++ b/generalresearch/models/thl/survey/__init__.py
@@ -3,7 +3,6 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from decimal import Decimal
from itertools import product
-from typing import Type
from more_itertools import flatten
from pydantic import BaseModel, Field
diff --git a/generalresearch/models/thl/user_quality_event.py b/generalresearch/models/thl/user_quality_event.py
index 52903e3..cfb4ff3 100644
--- a/generalresearch/models/thl/user_quality_event.py
+++ b/generalresearch/models/thl/user_quality_event.py
@@ -1,8 +1,8 @@
from __future__ import annotations
-from datetime import UTC, datetime, timezone
+from datetime import UTC, datetime
from decimal import Decimal
-from enum import Enum
+from enum import StrEnum
from typing import Literal
from pydantic import BaseModel, Field, PositiveInt
@@ -18,7 +18,7 @@ Typically used internally. These affect a user's quality standing.
"""
-class QualityEventType(str, Enum, metaclass=ReprEnumMeta):
+class QualityEventType(StrEnum, metaclass=ReprEnumMeta):
"""
Currently, the grpc call SendUserQualityEvents handles both the
recons/task adj, access control, and "security/hash failure" events.
diff --git a/generalresearch/models/thl/user_streak.py b/generalresearch/models/thl/user_streak.py
index 278809b..cc4d643 100644
--- a/generalresearch/models/thl/user_streak.py
+++ b/generalresearch/models/thl/user_streak.py
@@ -1,7 +1,8 @@
from __future__ import annotations
from datetime import date, datetime, timedelta
-from enum import Enum
+from enum import StrEnum
+from zoneinfo import ZoneInfo
import pandas as pd
from pydantic import (
@@ -15,14 +16,13 @@ from pydantic import (
model_validator,
)
from pydantic.json_schema import SkipJsonSchema
-from zoneinfo import ZoneInfo
from generalresearch.managers.leaderboard import country_timezone
from generalresearch.models import MAX_INT32
from generalresearch.models.thl.locales import CountryISO
-class StreakPeriod(str, Enum):
+class StreakPeriod(StrEnum):
# Midnight to midnight in the tz associated with the user's country
DAY = "day"
# Sunday midnight - sunday midnight
@@ -31,7 +31,7 @@ class StreakPeriod(str, Enum):
MONTH = "month"
-class StreakFulfillment(str, Enum):
+class StreakFulfillment(StrEnum):
"""
What has to happen for a user to fulfill a period for a streak
"""
@@ -42,7 +42,7 @@ class StreakFulfillment(str, Enum):
COMPLETE = "complete"
-class StreakState(str, Enum):
+class StreakState(StrEnum):
# The activity for today was completed!
ACTIVE = "active"
# They had activity yesterday, but not today, and can still continue today
diff --git a/generalresearch/models/thl/userhealth.py b/generalresearch/models/thl/userhealth.py
index fb15572..f0b7473 100644
--- a/generalresearch/models/thl/userhealth.py
+++ b/generalresearch/models/thl/userhealth.py
@@ -1,8 +1,8 @@
from __future__ import annotations
-from datetime import UTC, datetime, timezone
+from datetime import UTC, datetime
from enum import Enum
-from typing import Dict, Optional, Self
+from typing import Self
from pydantic import BaseModel, Field, NonNegativeFloat, PositiveInt
diff --git a/generalresearch/models/thl/wallet/__init__.py b/generalresearch/models/thl/wallet/__init__.py
index 1403ddf..2d1eb8d 100644
--- a/generalresearch/models/thl/wallet/__init__.py
+++ b/generalresearch/models/thl/wallet/__init__.py
@@ -1,9 +1,9 @@
-from enum import Enum
+from enum import StrEnum
from generalresearch.utils.enum import ReprEnumMeta
-class PayoutType(str, Enum, metaclass=ReprEnumMeta):
+class PayoutType(StrEnum, metaclass=ReprEnumMeta):
"""
The method in which the requested payout is delivered.
"""
@@ -35,7 +35,7 @@ class PayoutType(str, Enum, metaclass=ReprEnumMeta):
AMT_ASSIGNMENT = "AMT_ASSIGNMENT"
-class Currency(str, Enum):
+class Currency(StrEnum):
# United States Dollar
USD = "USD"
# Canadian Dollar
diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py
index 4757eba..40c4717 100644
--- a/generalresearch/models/thl/wallet/cashout_method.py
+++ b/generalresearch/models/thl/wallet/cashout_method.py
@@ -2,9 +2,9 @@ from __future__ import annotations
import hashlib
import logging
-from datetime import datetime, timezone, UTC
-from enum import Enum
-from typing import Any, Literal
+from datetime import UTC, datetime
+from enum import StrEnum
+from typing import Any, Literal, Self
from pydantic import (
BaseModel,
@@ -16,7 +16,6 @@ from pydantic import (
field_validator,
model_validator,
)
-from typing import Self
from generalresearch.currency import USDCent
from generalresearch.models.custom_types import (
@@ -144,9 +143,7 @@ class CashoutMethod(CashoutMethodBase):
"a user may have a paypal cashout method with their paypal"
"email associated.",
)
- last_updated: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=UTC)
- )
+ last_updated: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
is_live: bool = Field(default=True)
@model_validator(mode="after")
@@ -249,7 +246,7 @@ class CashoutMethodsResponse(StatusResponse):
cashout_methods: list[CashoutMethodOut] = Field()
-class DeliveryStatus(str, Enum):
+class DeliveryStatus(StrEnum):
PENDING = "Pending"
SHIPPED = "Shipped"
IN_TRANSIT = "In Transit"
@@ -261,14 +258,14 @@ class DeliveryStatus(str, Enum):
LOST = "Lost"
-class ShippingCarrier(str, Enum):
+class ShippingCarrier(StrEnum):
USPS = "USPS"
FEDEX = "FedEx"
UPS = "UPS"
DHL = "DHL"
-class ShippingMethod(str, Enum):
+class ShippingMethod(StrEnum):
STANDARD = "Standard"
EXPRESS = "Express"
TWO_DAY = "Two-Day"
@@ -394,7 +391,7 @@ example_foreign_value = {
}
-class RedemptionCurrency(str, Enum, metaclass=ReprEnumMeta):
+class RedemptionCurrency(StrEnum, metaclass=ReprEnumMeta):
"""
Supported Currencies for Foreign Redemptions
"""
diff --git a/generalresearch/utils/aggregation.py b/generalresearch/utils/aggregation.py
index bd962f1..ef3a52d 100644
--- a/generalresearch/utils/aggregation.py
+++ b/generalresearch/utils/aggregation.py
@@ -1,5 +1,5 @@
from collections import defaultdict
-from typing import Any, Dict, List
+from typing import Any
def group_by_year(records: list[dict], datetime_field: str) -> dict[int, list[Any]]:
diff --git a/generalresearch/utils/enum.py b/generalresearch/utils/enum.py
index 56706ba..e59a383 100644
--- a/generalresearch/utils/enum.py
+++ b/generalresearch/utils/enum.py
@@ -3,7 +3,6 @@ from __future__ import annotations
import inspect
import re
from enum import EnumMeta
-from typing import Dict
class ReprEnumMeta(EnumMeta):
diff --git a/generalresearch/wall_status_codes/__init__.py b/generalresearch/wall_status_codes/__init__.py
index 1d80924..cfccb35 100644
--- a/generalresearch/wall_status_codes/__init__.py
+++ b/generalresearch/wall_status_codes/__init__.py
@@ -1,5 +1,3 @@
-from typing import Optional, Tuple
-
from generalresearch.models import Source
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.session import Wall
diff --git a/generalresearch/wall_status_codes/cint.py b/generalresearch/wall_status_codes/cint.py
index ecb6219..c31c0ba 100644
--- a/generalresearch/wall_status_codes/cint.py
+++ b/generalresearch/wall_status_codes/cint.py
@@ -1,4 +1,4 @@
-from typing import Any, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.wall_status_codes import lucid
diff --git a/generalresearch/wall_status_codes/dynata.py b/generalresearch/wall_status_codes/dynata.py
index 69f3e74..53eebe7 100644
--- a/generalresearch/wall_status_codes/dynata.py
+++ b/generalresearch/wall_status_codes/dynata.py
@@ -4,7 +4,7 @@ checked by Greg 2023-10-10
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py
index f37ffda..aeaa4c7 100644
--- a/generalresearch/wall_status_codes/fullcircle.py
+++ b/generalresearch/wall_status_codes/fullcircle.py
@@ -7,7 +7,7 @@ we'll try to infer based on the time spent in survey.
from collections import defaultdict
from datetime import timedelta
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py
index fba3f54..936ee6c 100644
--- a/generalresearch/wall_status_codes/innovate.py
+++ b/generalresearch/wall_status_codes/innovate.py
@@ -9,7 +9,7 @@ can map directly, and some we have to look at the category.
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py
index 3ce89b0..05b75fb 100644
--- a/generalresearch/wall_status_codes/lucid.py
+++ b/generalresearch/wall_status_codes/lucid.py
@@ -5,7 +5,7 @@ https://support.lucidhq.com/s/article/Collecting-Data-From-Redirects
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py
index 42e34b3..bac9669 100644
--- a/generalresearch/wall_status_codes/morning.py
+++ b/generalresearch/wall_status_codes/morning.py
@@ -1,5 +1,5 @@
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py
index d3c4471..b033cf2 100644
--- a/generalresearch/wall_status_codes/precision.py
+++ b/generalresearch/wall_status_codes/precision.py
@@ -7,7 +7,7 @@ f - client approved the Preliminary complete as Final Complete
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/prodege.py b/generalresearch/wall_status_codes/prodege.py
index aaac376..2876ba9 100644
--- a/generalresearch/wall_status_codes/prodege.py
+++ b/generalresearch/wall_status_codes/prodege.py
@@ -3,7 +3,7 @@ https://developer.prodege.com/surveys-feed/term-reasons
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py
index 8b690b5..ad64338 100644
--- a/generalresearch/wall_status_codes/repdata.py
+++ b/generalresearch/wall_status_codes/repdata.py
@@ -3,7 +3,7 @@ Status codes are in a xlsx file. See thl-repdata readme
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/sago.py b/generalresearch/wall_status_codes/sago.py
index 028192b..b66190c 100644
--- a/generalresearch/wall_status_codes/sago.py
+++ b/generalresearch/wall_status_codes/sago.py
@@ -3,7 +3,7 @@ https://developer-beta.market-cube.com/api-details#api=definition-api&operation=
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/spectrum.py b/generalresearch/wall_status_codes/spectrum.py
index cf200a9..610e239 100644
--- a/generalresearch/wall_status_codes/spectrum.py
+++ b/generalresearch/wall_status_codes/spectrum.py
@@ -3,7 +3,7 @@ https://purespectrum.atlassian.net/wiki/spaces/PA/pages/33613201/Minimizing+Clic
"""
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Tuple
+from typing import Any
from generalresearch.models.thl.definitions import Status, StatusCode1
diff --git a/generalresearch/wall_status_codes/wxet.py b/generalresearch/wall_status_codes/wxet.py
index 1b7c514..6b3f67c 100644
--- a/generalresearch/wall_status_codes/wxet.py
+++ b/generalresearch/wall_status_codes/wxet.py
@@ -1,5 +1,4 @@
from collections import defaultdict
-from typing import Dict, List, Optional, Tuple
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.wxet.models.definitions import (
diff --git a/generalresearch/wxet/models/definitions.py b/generalresearch/wxet/models/definitions.py
index b08e178..ea0b51f 100644
--- a/generalresearch/wxet/models/definitions.py
+++ b/generalresearch/wxet/models/definitions.py
@@ -1,11 +1,10 @@
-from enum import Enum
-from typing import Optional, Tuple
+from enum import IntEnum, StrEnum
from generalresearch.currency import USDMill
from generalresearch.utils.enum import ReprEnumMeta
-class IncExcFilterType(str, Enum, metaclass=ReprEnumMeta):
+class IncExcFilterType(StrEnum, metaclass=ReprEnumMeta):
INCLUDE = "include"
EXCLUDE = "exclude"
@@ -13,7 +12,7 @@ class IncExcFilterType(str, Enum, metaclass=ReprEnumMeta):
# Note: This is exactly the same as the generalresearch:models/thl/definitions.py:Status.
# Keeping this because the comments (and as a result, the documentation)
# is slightly different, and specific to wxet.
-class WXETStatus(str, Enum, metaclass=ReprEnumMeta):
+class WXETStatus(StrEnum, metaclass=ReprEnumMeta):
"""
The outcome of a task attempt. If the attempt is still in progress, the status will be NULL.
"""
@@ -36,7 +35,7 @@ class WXETStatus(str, Enum, metaclass=ReprEnumMeta):
# Basically same note as for WxetStatus for WallAdjustedStatus
-class WXETAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
+class WXETAdjustedStatus(StrEnum, metaclass=ReprEnumMeta):
# Task was reconciled to complete
ADJUSTED_TO_COMPLETE = "ac"
@@ -52,7 +51,7 @@ class WXETAdjustedStatus(str, Enum, metaclass=ReprEnumMeta):
POSTBACK_COMPLETE = "pc"
-class WXETStatusCode1(int, Enum, metaclass=ReprEnumMeta):
+class WXETStatusCode1(IntEnum, metaclass=ReprEnumMeta):
"""
__High level status code for outcome of the attempt.__
This should only be NULL if the WXETStatus is ABANDON or TIMEOUT
@@ -103,10 +102,10 @@ class WXETStatusCode1(int, Enum, metaclass=ReprEnumMeta):
"""This property helper indicates if the WXET Attempt made it into
the WXET Account's (eg: the "buyer"'s) Task.
"""
- return False if self.value > 10 else True
+ return not self.value > 10
-class WXETStatusCode2(int, Enum, metaclass=ReprEnumMeta):
+class WXETStatusCode2(IntEnum, metaclass=ReprEnumMeta):
"""
__Status Detail__
These are generally only set if the StatusCode1 is WXET_FAIL,
diff --git a/generalresearch/wxet/models/finish_type.py b/generalresearch/wxet/models/finish_type.py
index a57dce8..17923d6 100644
--- a/generalresearch/wxet/models/finish_type.py
+++ b/generalresearch/wxet/models/finish_type.py
@@ -1,11 +1,10 @@
-from enum import Enum
-from typing import Optional, Set
+from enum import StrEnum
from generalresearch.utils.enum import ReprEnumMeta
from generalresearch.wxet.models.definitions import WXETStatus, WXETStatusCode1
-class FinishType(str, Enum, metaclass=ReprEnumMeta):
+class FinishType(StrEnum, metaclass=ReprEnumMeta):
"""A Task can be classified as "finished" based on different outcomes.
This controls how the `Task.required_finish_count` value
diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py
index 29d93fe..e98eecd 100644
--- a/tests/incite/schemas/test_admin_responses.py
+++ b/tests/incite/schemas/test_admin_responses.py
@@ -1,6 +1,5 @@
-from datetime import UTC, datetime, timedelta, timezone
+from datetime import UTC, datetime, timedelta
from random import sample
-from typing import List
import numpy as np
import pandas as pd
--
cgit v1.2.3
From aeeb7fef2594ccd34fbe96a77f6c5b392299fed7 Mon Sep 17 00:00:00 2001
From: Max Nanis
Date: Thu, 27 Aug 2026 17:46:11 -0700
Subject: Ruff afternoon
---
generalresearch/managers/thl/product.py | 1 -
generalresearch/models/network/nmap/parser.py | 8 +-
generalresearch/models/precision/question.py | 4 +-
generalresearch/models/precision/survey.py | 17 +-
generalresearch/models/prodege/survey.py | 3 +-
generalresearch/models/spectrum/question.py | 9 +-
generalresearch/models/spectrum/survey.py | 13 +-
generalresearch/models/spectrum/task_collection.py | 2 +-
generalresearch/models/thl/contest/contest.py | 10 +-
generalresearch/models/thl/contest/leaderboard.py | 13 +-
generalresearch/models/thl/contest/milestone.py | 20 +-
generalresearch/models/thl/contest/raffle.py | 46 ++--
generalresearch/models/thl/finance.py | 12 +-
generalresearch/models/thl/product.py | 36 +--
.../models/thl/profiling/other_option.py | 4 +-
.../models/thl/profiling/upk_question.py | 13 +-
.../models/thl/profiling/user_question_answer.py | 31 +--
pyproject.toml | 5 +-
.../collections/test_df_collection_item_thl_web.py | 273 ++++++++++-----------
.../mergers/foundations/test_enriched_session.py | 52 ++--
.../foundations/test_enriched_task_adjust.py | 38 ++-
.../mergers/foundations/test_enriched_wall.py | 73 +++---
tests/incite/mergers/test_merge_collection.py | 53 +++-
tests/incite/mergers/test_merge_collection_item.py | 25 +-
tests/incite/mergers/test_pop_ledger.py | 109 ++++----
tests/incite/mergers/test_ym_survey_merge.py | 55 +++--
tests/incite/schemas/test_admin_responses.py | 36 +--
tests/incite/schemas/test_thl_web.py | 8 +-
tests/incite/test_collection_base.py | 62 ++---
tests/incite/test_collection_base_item.py | 74 +++---
tests/incite/test_interval_idx.py | 6 +-
tests/managers/gr/test_business.py | 60 +++--
tests/managers/gr/test_team.py | 58 +++--
tests/managers/network/test_label.py | 4 +-
.../managers/thl/test_contest/test_leaderboard.py | 5 +-
tests/managers/thl/test_contest/test_milestone.py | 5 +-
tests/managers/thl/test_contest/test_raffle.py | 12 +-
tests/managers/thl/test_ledger/test_lm_accounts.py | 2 +-
tests/managers/thl/test_ledger/test_lm_tx.py | 4 -
.../thl/test_ledger/test_thl_lm_accounts.py | 11 +-
.../thl/test_ledger/test_thl_lm_bp_payout.py | 54 ++--
tests/managers/thl/test_ledger/test_thl_lm_tx.py | 26 +-
tests/managers/thl/test_payout.py | 2 -
tests/managers/thl/test_survey.py | 6 +-
tests/managers/thl/test_user_manager/test_base.py | 2 +-
tests/models/network/test_nmap.py | 2 +-
tests/models/spectrum/test_survey.py | 44 ++++
tests/models/thl/test_product.py | 4 +-
48 files changed, 815 insertions(+), 597 deletions(-)
(limited to 'tests/incite/schemas/test_admin_responses.py')
diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py
index d924e17..54fa7c8 100644
--- a/generalresearch/managers/thl/product.py
+++ b/generalresearch/managers/thl/product.py
@@ -33,7 +33,6 @@ if TYPE_CHECKING:
ProfilingConfig,
SessionConfig,
SourcesConfig,
- SupplyConfigs,
UserCreateConfig,
UserHealthConfig,
UserWalletConfig,
diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py
index ecaf2d1..866b4bd 100644
--- a/generalresearch/models/network/nmap/parser.py
+++ b/generalresearch/models/network/nmap/parser.py
@@ -48,7 +48,7 @@ class NmapXmlParser:
try:
root = ET.fromstring(nmap_data)
- except Exception as e:
+ except ET.ParseError as e:
emsg = f"Wrong XML structure: cannot parse data: {e}"
raise NmapParserException(emsg)
@@ -103,7 +103,7 @@ class NmapXmlParser:
@classmethod
def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo:
- data = dict()
+ data = {}
data["type"] = NmapScanType(scaninfo_el.attrib["type"])
data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"])
data["num_services"] = scaninfo_el.attrib["numservices"]
@@ -132,7 +132,7 @@ class NmapXmlParser:
@classmethod
def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict:
- nmap_data = dict()
+ nmap_data = {}
nmaprun = dict(nmaprun_el.attrib)
nmap_data["command_line"] = nmaprun["args"]
nmap_data["started_at"] = datetime.fromtimestamp(
@@ -148,7 +148,7 @@ class NmapXmlParser:
Receives a XML tag representing a scanned host with
its services.
"""
- data = dict()
+ data = {}
#
status_el = host_el.find("status")
diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py
index a2189d5..cc90aa9 100644
--- a/generalresearch/models/precision/question.py
+++ b/generalresearch/models/precision/question.py
@@ -6,7 +6,7 @@ import logging
from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal
-from pydantic import BaseModel, Field, field_validator, model_validator
+from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator
from generalresearch.models import Source, string_utils
from generalresearch.models.precision import PrecisionQuestionID
@@ -112,7 +112,7 @@ class PrecisionQuestion(MarketplaceQuestion):
"""
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py
index f515552..b27b8c4 100644
--- a/generalresearch/models/precision/survey.py
+++ b/generalresearch/models/precision/survey.py
@@ -99,11 +99,11 @@ class PrecisionQuota(BaseModel):
self, criteria_evaluation: dict[str, bool | None]
) -> tuple[bool | None, list[str]]:
# Passes back "matches" (T/F/none) and a list of unknown criterion hashes
- unknowns = list()
+ unknowns = []
for c in self.condition_hashes:
eval_value = criteria_evaluation.get(c)
if eval_value is False:
- return False, list()
+ return False, []
if eval_value is None:
unknowns.append(c)
if unknowns:
@@ -245,11 +245,10 @@ class PrecisionSurvey(MarketplaceTask):
# Fancy repr that abbreviates exclude_pids and excluded_surveys
repr_args = list(self.__repr_args__())
for n, (k, v) in enumerate(repr_args):
- if k in {"excluded_surveys"}:
- if v and len(v) > 6:
- v = sorted(v)
- v = v[:3] + ["…"] + v[-3:]
- repr_args[n] = (k, v)
+ if k in {"excluded_surveys"} and v and len(v) > 6:
+ v = sorted(v)
+ v = v[:3] + ["…"] + v[-3:]
+ repr_args[n] = (k, v)
join_str = ", "
repr_str = join_str.join(
repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args
@@ -369,6 +368,4 @@ class PrecisionSurvey(MarketplaceTask):
return False
if self.group_id in att_group_ids:
return False
- if self.excluded_surveys & att_survey_ids:
- return False
- return True
+ return not self.excluded_surveys & att_survey_ids
diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py
index 3f4c88f..7ab6df6 100644
--- a/generalresearch/models/prodege/survey.py
+++ b/generalresearch/models/prodege/survey.py
@@ -13,6 +13,7 @@ from pydantic import (
BaseModel,
ConfigDict,
Field,
+ ValidationError,
computed_field,
field_validator,
model_validator,
@@ -513,7 +514,7 @@ class ProdegeSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> ProdegeSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py
index c8eea4a..7add692 100644
--- a/generalresearch/models/spectrum/question.py
+++ b/generalresearch/models/spectrum/question.py
@@ -13,6 +13,7 @@ from pydantic import (
BaseModel,
Field,
PositiveInt,
+ ValidationError,
field_validator,
model_validator,
)
@@ -132,7 +133,7 @@ class SpectrumQuestionType(StrEnum):
@classmethod
def from_api(cls, a: int):
api_type_map = cls.get_api_map()
- return api_type_map[a] if a in api_type_map else None
+ return api_type_map.get(a, None)
class SpectrumQuestionClass(IntEnum):
@@ -260,7 +261,7 @@ class SpectrumQuestion(MarketplaceQuestion):
return None
try:
return cls._from_api(d, country_iso, language_iso)
- except Exception as e:
+ except ValidationError as e:
logger.warning(f"Unable to parse question: {d}. {e}")
return None
@@ -280,7 +281,9 @@ class SpectrumQuestion(MarketplaceQuestion):
]
created = (
- datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=UTC)
+ datetime.fromtimestamp(timestamp=d["crtd_on"] / 1000, tz=UTC).replace(
+ tzinfo=UTC
+ )
if d.get("crtd_on")
else None
)
diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py
index f9a5e27..424d206 100644
--- a/generalresearch/models/spectrum/survey.py
+++ b/generalresearch/models/spectrum/survey.py
@@ -75,8 +75,7 @@ class SpectrumCondition(MarketplaceCondition):
rs["from"] = round(rs["from"] / 12)
rs["to"] = round(rs["to"] / 12)
d["values"] = [
- f"{rs["from"] or "inf"}-{rs["to"] or "inf"}"
- for rs in d["range_sets"]
+ f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"]
]
d["value_type"] = ConditionValueType.RANGE
return cls.model_validate(d)
@@ -103,7 +102,7 @@ class SpectrumQuota(BaseModel):
# There is no explicit status. The quota is closed if the count is 0
def __hash__(self) -> int:
- return hash(tuple((tuple(self.condition_hashes), self.remaining_count)))
+ return hash((tuple(self.condition_hashes), self.remaining_count))
@property
def is_open(self) -> bool:
@@ -113,7 +112,7 @@ class SpectrumQuota(BaseModel):
return self.remaining_count >= min_open_spots
@classmethod
- def from_api(cls, d: dict) -> Self:
+ def from_api(cls, d: dict[str, Any]) -> Self:
d["remaining_count"] = d["quantities"]["currently_open"]
return cls.model_validate(d)
@@ -323,7 +322,7 @@ class SpectrumSurvey(MarketplaceTask):
def from_api(cls, d: dict[str, Any]) -> SpectrumSurvey | None:
try:
return cls._from_api(d)
- except Exception as e:
+ except (AssertionError, ValueError) as e:
logger.warning(f"Unable to parse survey: {d}. {e}")
return None
@@ -336,7 +335,7 @@ class SpectrumSurvey(MarketplaceTask):
else TaskCalculationType.COMPLETES
)
- d["conditions"] = dict()
+ d["conditions"] = {}
# If we haven't hit the "detail" endpoint, we won't get this
d.setdefault("qualifications", [])
@@ -454,7 +453,7 @@ class SpectrumSurvey(MarketplaceTask):
quota_eval = {
quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas
}
- evals = set(g[0] for g in quota_eval.values())
+ evals = {g[0] for g in quota_eval.values()}
if any(m[0] is True and not q.is_open for q, m in quota_eval.items()):
# matched a full quota
return False, set()
diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py
index 6715378..8ca5a93 100644
--- a/generalresearch/models/spectrum/task_collection.py
+++ b/generalresearch/models/spectrum/task_collection.py
@@ -91,7 +91,7 @@ class SpectrumTaskCollection(TaskCollection):
"survey_id",
]
rows = []
- d = dict()
+ d = {}
for k in fields:
d[k] = getattr(s, k) if hasattr(s, k) else None
d["used_question_ids"] = list(s.used_question_ids)
diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py
index 2a8853d..bd0fc04 100644
--- a/generalresearch/models/thl/contest/contest.py
+++ b/generalresearch/models/thl/contest/contest.py
@@ -136,10 +136,12 @@ class Contest(ContestBase):
# return False
def should_end(self) -> tuple[bool, ContestEndReason | None]:
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.ends_at:
- if datetime.now(tz=UTC) >= self.end_condition.ends_at:
- return True, ContestEndReason.ENDS_AT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.ends_at
+ and datetime.now(tz=UTC) >= self.end_condition.ends_at
+ ):
+ return True, ContestEndReason.ENDS_AT
return False, None
diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py
index 696cdea..e923383 100644
--- a/generalresearch/models/thl/contest/leaderboard.py
+++ b/generalresearch/models/thl/contest/leaderboard.py
@@ -151,7 +151,8 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
len(self.country_isos) == 1
), "Can only set 1 country_iso in a leaderboard contest"
assert (
- list(self.country_isos)[0] == self.leaderboard_key_parts["country_iso"]
+ next(iter(self.country_isos))
+ == self.leaderboard_key_parts["country_iso"]
), "leaderboard_key country_iso must match the country_isos"
else:
self.country_isos = {self.leaderboard_key_parts["country_iso"]}
@@ -192,10 +193,12 @@ class LeaderboardContest(LeaderboardContestCreate, Contest):
return lbm
def should_end(self) -> tuple[bool, ContestEndReason | None]:
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.ends_at:
- if datetime.now(tz=UTC) >= self.end_condition.ends_at:
- return True, ContestEndReason.ENDS_AT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.ends_at
+ and datetime.now(tz=UTC) >= self.end_condition.ends_at
+ ):
+ return True, ContestEndReason.ENDS_AT
return False, None
diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py
index 8d96fcb..5fc27fa 100644
--- a/generalresearch/models/thl/contest/milestone.py
+++ b/generalresearch/models/thl/contest/milestone.py
@@ -132,10 +132,12 @@ class MilestoneContest(MilestoneContestCreate, Contest):
if res:
return res, msg
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.max_winners:
- if self.win_count >= self.end_condition.max_winners:
- return True, ContestEndReason.MAX_WINNERS
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.max_winners
+ and self.win_count >= self.end_condition.max_winners
+ ):
+ return True, ContestEndReason.MAX_WINNERS
return False, None
@@ -189,16 +191,10 @@ class MilestoneUserView(MilestoneContest, ContestUserView):
)
def should_award(self):
- if self.status == ContestStatus.ACTIVE:
- if self.should_have_awarded():
- return True
- return False
+ return bool(self.status == ContestStatus.ACTIVE and self.should_have_awarded())
def should_have_awarded(self):
- if self.target_amount:
- if self.user_amount >= self.target_amount:
- return True
- return False
+ return bool(self.target_amount and self.user_amount >= self.target_amount)
def is_user_eligible(self, country_iso: str) -> tuple[bool, str]:
passes, msg = super().is_user_eligible(country_iso=country_iso)
diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py
index 08243f4..16a0a47 100644
--- a/generalresearch/models/thl/contest/raffle.py
+++ b/generalresearch/models/thl/contest/raffle.py
@@ -127,7 +127,7 @@ class RaffleContest(RaffleContestCreate, Contest):
# If there is more than 1 prize, the winning entry is subtracted
# from the user's entry count
user_amount = defaultdict(int)
- user_id_user = dict()
+ user_id_user = {}
for entry in self.entries:
user_amount[entry.user.user_id] += entry.amount
user_id_user[entry.user.user_id] = entry.user
@@ -149,10 +149,12 @@ class RaffleContest(RaffleContestCreate, Contest):
res, msg = super().should_end()
if res:
return res, msg
- if self.status == ContestStatus.ACTIVE:
- if self.end_condition.target_entry_amount:
- if self.current_amount >= self.end_condition.target_entry_amount:
- return True, ContestEndReason.TARGET_ENTRY_AMOUNT
+ if (
+ self.status == ContestStatus.ACTIVE
+ and self.end_condition.target_entry_amount
+ and self.current_amount >= self.end_condition.target_entry_amount
+ ):
+ return True, ContestEndReason.TARGET_ENTRY_AMOUNT
return False, None
@staticmethod
@@ -278,17 +280,19 @@ class RaffleUserView(RaffleContest, ContestUserView):
return probs
def is_entry_eligible(self, entry: ContestEntry) -> tuple[bool, str]:
- if self.entry_rule.max_entry_amount_per_user:
- if (
- self.user_amount + entry.amount
- ) > self.entry_rule.max_entry_amount_per_user:
- return False, "Entry would exceed max amount per user."
-
- if self.entry_rule.max_daily_entries_per_user:
- if (
- self.user_amount_today + entry.amount
- ) > self.entry_rule.max_daily_entries_per_user:
- return False, "Entry would exceed max amount per user per day."
+ if (
+ self.entry_rule.max_entry_amount_per_user
+ and (self.user_amount + entry.amount)
+ > self.entry_rule.max_entry_amount_per_user
+ ):
+ return False, "Entry would exceed max amount per user."
+
+ if (
+ self.entry_rule.max_daily_entries_per_user
+ and (self.user_amount_today + entry.amount)
+ > self.entry_rule.max_daily_entries_per_user
+ ):
+ return False, "Entry would exceed max amount per user per day."
return True, ""
def is_user_eligible(self, country_iso: str) -> tuple[bool, str]:
@@ -296,16 +300,18 @@ class RaffleUserView(RaffleContest, ContestUserView):
if not passes:
return False, msg
- if self.entry_rule.max_entry_amount_per_user:
+ if self.entry_rule.max_entry_amount_per_user: # noqa: SIM102
# Greater or equal b/c we're asking if the user is eligible to
# enter MORE, now! If it equals, nothing is wrong, just that they
# are not eligible anymore.
if self.user_amount >= self.entry_rule.max_entry_amount_per_user:
return False, "Reached max amount per user."
- if self.entry_rule.max_daily_entries_per_user:
- if self.user_amount_today >= self.entry_rule.max_daily_entries_per_user:
- return False, "Reached max amount today."
+ if (
+ self.entry_rule.max_daily_entries_per_user
+ and self.user_amount_today >= self.entry_rule.max_daily_entries_per_user
+ ):
+ return False, "Reached max amount today."
# This would indicate something is wrong, as something else should have done this
e, _ = self.should_end()
diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py
index 0856825..8c94390 100644
--- a/generalresearch/models/thl/finance.py
+++ b/generalresearch/models/thl/finance.py
@@ -124,10 +124,10 @@ class POPFinancial(BaseModel):
Direction,
)
- assert all([a.account_type == AccountType.BP_WALLET for a in accounts])
- assert all([a.normal_balance == Direction.CREDIT for a in accounts])
+ assert all(a.account_type == AccountType.BP_WALLET for a in accounts)
+ assert all(a.normal_balance == Direction.CREDIT for a in accounts)
if not is_debug():
- assert all([a.currency == "USD" for a in accounts])
+ assert all(a.currency == "USD" for a in accounts)
if input_data.empty:
return []
@@ -850,11 +850,11 @@ class BusinessBalances(BaseModel):
# Validate the input accounts
assert len(accounts) > 0, "Must provide accounts"
- assert all([a.account_type == AccountType.BP_WALLET for a in accounts])
- assert all([a.normal_balance == Direction.CREDIT for a in accounts])
+ assert all(a.account_type == AccountType.BP_WALLET for a in accounts)
+ assert all(a.normal_balance == Direction.CREDIT for a in accounts)
if not is_debug():
- assert all([a.currency == "USD" for a in accounts])
+ assert all(a.currency == "USD" for a in accounts)
# Validate the input dataframe
assert input_data.index.name == "account_id"
diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py
index 76a8e83..a7ecd55 100644
--- a/generalresearch/models/thl/product.py
+++ b/generalresearch/models/thl/product.py
@@ -430,8 +430,8 @@ class UserWalletConfig(BaseModel):
@field_serializer("supported_payout_types", when_used="json")
def serialize_supported_payout_types_in_order(
self, supported_payout_types: set[PayoutType]
- ) -> set[PayoutType]:
- return set(sorted(supported_payout_types))
+ ) -> list[PayoutType]:
+ return sorted(supported_payout_types)
@field_validator("min_cashout", mode="after")
@classmethod
@@ -552,14 +552,14 @@ class PayoutTransformation(BaseModel):
min_payout = Decimal(0)
pct = Decimal(pct)
- payout = Decimal(payout)
+ _payout = Decimal(payout)
min_payout = Decimal(min_payout)
max_payout = Decimal(max_payout) if max_payout else None
- payout: Decimal = payout * pct
- payout: Decimal = max([payout, min_payout])
- payout: Decimal = min([payout, max_payout]) if max_payout else payout
- return payout
+ _payout: Decimal = _payout * pct
+ _payout: Decimal = max([_payout, min_payout])
+ _payout: Decimal = min([_payout, max_payout]) if max_payout else payout
+ return _payout
def payout_transformation_amt(
self, payout: Decimal, user_wallet_balance: Decimal | None = None
@@ -569,22 +569,22 @@ class PayoutTransformation(BaseModel):
# (display, adjustment) so ignore the 7-cent rounding.
if user_wallet_balance is None:
return self.payout_transformation_percent(payout=payout, pct=Decimal(".95"))
- payout = Decimal(payout)
+ _payout = Decimal(payout)
- payout: Decimal = payout * Decimal("0.95")
- new_balance = payout + user_wallet_balance
+ _payout: Decimal = _payout * Decimal("0.95")
+ new_balance = _payout + user_wallet_balance
# If the new_balance is <0, we aren't paying anything, so use the
# full amount
if new_balance < 0:
- return payout
+ return _payout
amt = (5 * math.floor((int(new_balance * 100) - 2) / 5)) + 2
rounded_new_balance = Decimal(amt / 100).quantize(Decimal("0.00"))
- payout = rounded_new_balance - user_wallet_balance
- if payout < Decimal(0):
+ _payout = rounded_new_balance - user_wallet_balance
+ if _payout < Decimal(0):
return Decimal(0)
- return payout
+ return _payout
class SourceConfig(BaseModel):
@@ -731,8 +731,8 @@ class SupplyConfig(BaseModel):
Use global config.
"""
d = self.global_scoped_policies_dict.copy()
- d.update(self.team_scoped_policies_dict.get(team_id, dict()))
- d.update(self.product_scoped_policies_dict.get(product_id, dict()))
+ d.update(self.team_scoped_policies_dict.get(team_id, {}))
+ d.update(self.product_scoped_policies_dict.get(product_id, {}))
return d
def get_config_for_product(self, product: Product) -> MergedSupplyConfig:
@@ -751,7 +751,7 @@ class SupplyConfig(BaseModel):
supply_policy=policy_dict[source],
source_config=sources_dict[source],
)
- for source in policy_dict.keys()
+ for source in policy_dict
]
)
@@ -1000,10 +1000,12 @@ class Product(BaseModel, validate_assignment=True):
@property
def business_uuid(self) -> UUIDStr:
+ assert self.business_id
return self.business_id
@property
def team_uuid(self) -> UUIDStr:
+ assert self.team_id
return self.team_id
@property
diff --git a/generalresearch/models/thl/profiling/other_option.py b/generalresearch/models/thl/profiling/other_option.py
index 6d789e5..2f3cac9 100644
--- a/generalresearch/models/thl/profiling/other_option.py
+++ b/generalresearch/models/thl/profiling/other_option.py
@@ -51,6 +51,4 @@ def option_is_catch_all(c: UpkQuestionChoice) -> bool:
return True
if c.text.lower() in texts_exact:
return True
- if any(t in c.text.lower() for t in texts_in):
- return True
- return False
+ return bool(any(t in c.text.lower() for t in texts_in))
diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py
index 77bba6f..3bb0733 100644
--- a/generalresearch/models/thl/profiling/upk_question.py
+++ b/generalresearch/models/thl/profiling/upk_question.py
@@ -475,10 +475,9 @@ class UpkQuestion(BaseModel):
# Almost nothing has >1k options, besides location stuff (cities,
# etc.) which should get harmonized. When presenting them, we'll
# filter down options to at most 50.
- if self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000):
- return False
-
- return True
+ return not (
+ self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000)
+ )
@property
def md5sum(self):
@@ -534,7 +533,7 @@ class UpkQuestion(BaseModel):
), "Multiple of the same answer submitted"
if self.type == UpkQuestionType.MULTIPLE_CHOICE:
assert len(answer) >= 1, "MC question with no selected answers"
- choice_codes = set(x.id for x in self.choices)
+ choice_codes = {x.id for x in self.choices}
if self.selector == UpkQuestionSelectorMC.SINGLE_ANSWER:
assert (
len(answer) == 1
@@ -563,9 +562,7 @@ class UpkQuestion(BaseModel):
assert len(answer) == 1, "Only one answer allowed"
answer = answer[0]
assert len(answer) > 0, "Must provide answer"
- max_length = (
- self.configuration.max_length if self.configuration else 0 or 100000
- )
+ max_length = self.configuration.max_length if self.configuration else 100000
assert len(answer) <= max_length, "Answer longer than allowed"
if self.validation and self.validation.patterns:
for pattern in self.validation.patterns:
diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py
index 2db07b7..378345e 100644
--- a/generalresearch/models/thl/profiling/user_question_answer.py
+++ b/generalresearch/models/thl/profiling/user_question_answer.py
@@ -3,7 +3,7 @@ from __future__ import annotations
import json
from collections.abc import Iterator
from datetime import UTC, datetime, timedelta
-from typing import Any, Literal, Self
+from typing import Any, Literal
from pydantic import (
BaseModel,
@@ -14,7 +14,6 @@ from pydantic import (
model_validator,
)
-from generalresearch.grpc import timestamp_to_datetime
from generalresearch.models import MAX_INT32, Source
from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
from generalresearch.models.thl.locales import CountryISO, LanguageISO
@@ -39,19 +38,23 @@ class UserQuestionAnswer(BaseModel):
calc_answers: dict[str, tuple[str, ...]] | None = Field(default=None)
@field_validator("calc_answers")
- def sorted_calc_answers(cls, calc_answers) -> dict[str, tuple[str, ...]] | None:
+ def sorted_calc_answers(
+ cls, calc_answers: dict[str, tuple[str, ...]] | None
+ ) -> dict[str, tuple[str, ...]] | None:
if calc_answers is None:
return None
return {k: tuple(sorted(v)) for k, v in calc_answers.items()}
@field_validator("calc_answers")
- def validate_keys(cls, calc_answers) -> dict[str, tuple[str, ...]] | None:
+ def validate_keys(
+ cls, calc_answers: dict[str, tuple[str, ...]] | None
+ ) -> dict[str, tuple[str, ...]] | None:
if calc_answers is None:
return None
assert all(
- ":" in k for k in calc_answers.keys()
+ ":" in k for k in calc_answers
), "calc_answers expects the keys to be in format source:question_code"
return calc_answers
@@ -66,6 +69,7 @@ class UserQuestionAnswer(BaseModel):
return d
def get_mrpqs(self) -> Iterator[MarketplaceResearchProfileQuestion]:
+ assert self.calc_answers
for k, v in self.calc_answers.items():
source, question_code = k.split(":", 1)
yield MarketplaceResearchProfileQuestion(
@@ -105,21 +109,6 @@ class UserQuestionAnswer(BaseModel):
def is_stale(self) -> bool:
return self.timestamp < datetime.now(tz=UTC) - timedelta(days=30)
- @classmethod
- def from_grpc(cls, msg, default_timestamp: datetime) -> Self:
- """
- Handles correctly issues with grpc timestamps
- :param msg: "thl.protos.generalresearch_pb2.ProfilingQuestionAnswer"
- """
- assert default_timestamp.tzinfo is not None, "must use tz-aware timestamps"
- timestamp = timestamp_to_datetime(msg.timestamp)
- timestamp = default_timestamp if timestamp < datetime(2000, 1, 1) else timestamp
- return cls(
- question_id=msg.question_id,
- answer=tuple(msg.answer),
- timestamp=timestamp,
- )
-
# We can't set a redis list to [] vs None. We'll push this dummy answer into
# the cache to signify the user has no answered questions. It'll get removed
@@ -131,7 +120,7 @@ DUMMY_UQA = UserQuestionAnswer(
country_iso="xx",
language_iso="xxx",
property_code="dummy",
- calc_answers=dict(),
+ calc_answers={},
)
diff --git a/pyproject.toml b/pyproject.toml
index dbdf3b9..03a1a1f 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -56,4 +56,7 @@ testpaths = ["tests"]
addopts = "-v --tb=short"
[tool.ruff]
-target-version = "py314"
\ No newline at end of file
+target-version = "py314"
+exclude = [
+ "generalresearch/thl_django",
+]
\ No newline at end of file
diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py
index 8038d3b..edf90f7 100644
--- a/tests/incite/collections/test_df_collection_item_thl_web.py
+++ b/tests/incite/collections/test_df_collection_item_thl_web.py
@@ -5,12 +5,12 @@ from datetime import UTC, datetime, timedelta
from itertools import product as iter_product
from os.path import join as pjoin
from pathlib import Path, PurePath
-from typing import TYPE_CHECKING
from uuid import uuid4
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
from distributed import Client, Scheduler, Worker
# noinspection PyUnresolvedReferences
@@ -21,20 +21,19 @@ from faker import Faker
from pandera.pandas import DataFrameSchema
from pydantic import FilePath
-from generalresearch.incite.base import CollectionItemBase
+from generalresearch.incite.base import CollectionItemBase, GRLDatasets
from generalresearch.incite.collections import (
+ DFCollection,
DFCollectionItem,
DFCollectionType,
)
from generalresearch.incite.schemas import ARCHIVE_AFTER
+from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.models.thl.product import Product
from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
from generalresearch.sql_helper import PostgresDsn
-if TYPE_CHECKING:
- from generalresearch.incite.base import GRLDatasets
-
fake = Faker()
df_collections = [
@@ -72,7 +71,12 @@ class TestDFCollectionItemBase:
)
class TestDFCollectionItemProperties:
- def test_filename(self, df_collection_data_type, df_collection, offset: str):
+ def test_filename(
+ self,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
+ offset: str,
+ ):
for i in df_collection.items:
assert isinstance(i.filename, str)
@@ -89,37 +93,59 @@ class TestDFCollectionItemProperties:
)
class TestDFCollectionItemPropertiesBase:
- def test_name(self, df_collection_data_type, offset: str, df_collection):
+ def test_name(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.name, str)
- def test_finish(self, df_collection_data_type, offset: str, df_collection):
+ def test_finish(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.finish, datetime)
- def test_interval(self, df_collection_data_type, offset: str, df_collection):
+ def test_interval(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.interval, pd.Interval)
def test_partial_filename(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection: DFCollection,
):
for i in df_collection.items:
assert isinstance(i.partial_filename, str)
- def test_empty_filename(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_filename(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_filename, str)
- def test_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.path, FilePath)
- def test_partial_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_partial_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.partial_path, FilePath)
- def test_empty_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_path, FilePath)
@@ -138,11 +164,8 @@ class TestDFCollectionItemMethod:
def test_has_mysql(
self,
- df_collection,
+ df_collection: DFCollection,
thl_web_rr: PostgresConfig,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
@@ -168,12 +191,6 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_update_partial_archive(
self,
- df_collection,
- offset: str,
- duration: timedelta,
- thl_web_rw: PostgresConfig,
- df_collection_data_type,
- delete_df_collection: Callable[..., None],
):
# for i in collection.items:
# assert i.update_partial_archive()
@@ -183,28 +200,12 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_create_partial_archive(
self,
- df_collection,
- offset: str,
- duration: str,
- create_main_accounts: Callable[..., None],
- thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
- user_factory: Callable[..., User],
- product: product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection: Callable[..., None],
- mnt_filepath: GRLDatasets,
):
assert 1 + 1 == 2
def test_dict(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- df_collection,
+ df_collection: DFCollection,
delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
@@ -225,15 +226,15 @@ class TestDFCollectionItemMethod:
def test_from_mysql(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
create_main_accounts: Callable[..., None],
thl_web_rw: PostgresConfig,
user_factory: Callable[..., User],
- product: product: Product,
- incite_item_factory,
+ product: Product,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
):
@@ -253,12 +254,14 @@ class TestDFCollectionItemMethod:
if df_collection.data_type == DFCollectionType.LEDGER:
assert df is None
else:
+ assert isinstance(df, pd.DataFrame)
assert df.empty
assert set(df.columns) == set(df_collection._schema.columns.keys())
incite_item_factory(user=u1, item=item)
df = item.from_mysql()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
assert set(df.columns) == set(df_collection._schema.columns.keys())
if df_collection.data_type == DFCollectionType.LEDGER:
@@ -270,13 +273,13 @@ class TestDFCollectionItemMethod:
def test_from_mysql_standard(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
user_factory: Callable[..., User],
- product: product: Product,
- incite_item_factory,
+ product: Product,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
):
@@ -293,7 +296,7 @@ class TestDFCollectionItemMethod:
# We're using parametrize, so this If statement is just to
# confirm other Item Types will always raise an assertion
with pytest.raises(expected_exception=AssertionError) as cm:
- res = item.from_mysql_standard()
+ _ = item.from_mysql_standard()
assert (
"Can't call from_mysql_standard for Ledger DFCollectionItem"
in str(cm.value)
@@ -304,32 +307,34 @@ class TestDFCollectionItemMethod:
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
df = item.from_mysql_standard()
+ assert isinstance(df, pd.DataFrame)
assert df.empty
assert set(df.columns) == set(df_collection._schema.columns.keys())
incite_item_factory(user=u1, item=item)
df = item.from_mysql_standard()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
assert set(df.columns) == set(df_collection._schema.columns.keys())
assert df.shape[0] > 0
def test_from_mysql_ledger(
self,
- df_collection,
+ df_collection: DFCollection,
user: User,
create_main_accounts: Callable[..., None],
offset: str,
duration: timedelta,
thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
+ thl_ledger_manager: ThlLedgerManager,
+ df_collection_data_type: DFCollectionType,
user_factory: Callable[..., User],
- product: product: Product,
- client_no_amm,
- incite_item_factory,
+ product: Product,
+ client_no_amm: DaskClient,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath,
+ mnt_filepath: GRLDatasets,
):
if df_collection.data_type != DFCollectionType.LEDGER:
@@ -370,17 +375,17 @@ class TestDFCollectionItemMethod:
def test_to_archive(
self,
- df_collection,
+ df_collection: DFCollection,
user: User,
offset: str,
duration: timedelta,
- df_collection_data_type,
+ df_collection_data_type: DFCollectionType,
user_factory: Callable[..., User],
- product: product: Product,
- client_no_amm,
- incite_item_factory,
+ product: Product,
+ client_no_amm: DaskClient,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath,
+ mnt_filepath: GRLDatasets,
):
if df_collection.data_type in unsupported_mock_types:
@@ -407,17 +412,17 @@ class TestDFCollectionItemMethod:
def test__to_archive(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
- product: product: Product,
+ product: Product,
offset: str,
duration: timedelta,
- client_no_amm,
+ client_no_amm: DaskClient,
user: User,
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath,
+ mnt_filepath: GRLDatasets,
):
"""We already have a test for the "non-private" version of this,
which primarily just uses the respective Client to determine if
@@ -480,19 +485,19 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_to_archive_numbered_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_initial_load(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_clear_corrupt_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@@ -505,34 +510,40 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_path_exists(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_next_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_search_highest_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_tmp_filename(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
- def test_tmp_path(self, df_collection_data_type, offset: str, duration: timedelta):
+ def test_tmp_path(
+ self,
+ ):
pass
def test_is_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
"""
test_has_empty was merged into this because item.has_empty is
@@ -549,7 +560,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_empty()
def test_has_partial_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
assert not item.has_partial_archive()
@@ -557,7 +569,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_partial_archive()
def test_has_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
# (1) Originally, nothing exists... so let's just make a file and
@@ -594,7 +607,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_archive(include_empty=True)
def test_delete_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
item: DFCollectionItem
@@ -617,7 +631,8 @@ class TestDFCollectionItemMethodBase:
assert not item.partial_path.exists()
def test_should_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
schema: DataFrameSchema = df_collection._schema
aa = schema.metadata[ARCHIVE_AFTER]
@@ -635,12 +650,13 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_set_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
def test_valid_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
# Originally, nothing has been saved or anything.. so confirm it
# always comes back as None
@@ -664,18 +680,19 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_validate_df(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_from_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
def test__to_dict(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
@@ -694,19 +711,19 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_delete_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_cleanup_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@pytest.mark.skip
def test_delete_dangling_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
):
pass
@@ -726,7 +743,9 @@ async def test_client(client, s, worker):
)
@gen_cluster(client=True, nthreads=[("127.0.0.1", 1)])
@pytest.mark.anyio
-async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str):
+async def test_client_parametrize(
+ c, s, w, df_collection_data_type: DFCollectionType, offset: str
+):
"""c,s,a are all required - the secondary Worker (b) is not required"""
assert isinstance(c, Client), f"c is not Client, it's {type(c)}"
@@ -750,17 +769,12 @@ class TestDFCollectionItemFunctionalTest:
def test_to_archive_and_ddf(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- client_no_amm,
- df_collection,
- user: User,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
- product: product: Product,
- incite_item_factory,
+ product: Product,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath: GRLDatasets,
):
if df_collection.data_type in unsupported_mock_types:
@@ -799,17 +813,11 @@ class TestDFCollectionItemFunctionalTest:
def test_filesize_estimate(
self,
- df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- client_no_amm,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
- product: product: Product,
- df_collection_data_type,
- incite_item_factory,
+ product: Product,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath: GRLDatasets,
):
"""A functional test to write some Parquet files for the
DFCollection and then confirm that the files get written
@@ -846,16 +854,12 @@ class TestDFCollectionItemFunctionalTest:
def test_to_archive_client(
self,
- client_no_amm,
- df_collection,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
- product: product: Product,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
- incite_item_factory,
+ product: Product,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
- mnt_filepath: GRLDatasets,
):
delete_df_collection(coll=df_collection)
@@ -885,7 +889,8 @@ class TestDFCollectionItemFunctionalTest:
@pytest.mark.skip
def test_get_items(
- self, df_collection, product: product: Product, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
with pytest.warns(expected_warning=ResourceWarning) as cm:
df_collection.get_items_last365()
@@ -898,16 +903,11 @@ class TestDFCollectionItemFunctionalTest:
def test_saving_protections(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
- product: product: Product,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
+ product: Product,
):
"""Don't allow creating an archive for data that will likely be
overwritten or updated
@@ -939,15 +939,8 @@ class TestDFCollectionItemFunctionalTest:
def test_empty_item(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
+ df_collection: DFCollection,
delete_df_collection: Callable[..., None],
- user: User,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
):
delete_df_collection(coll=df_collection)
@@ -967,16 +960,12 @@ class TestDFCollectionItemFunctionalTest:
def test_file_touching(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
- product: product: Product,
- offset: str,
- duration: timedelta,
- mnt_filepath,
+ product: Product,
):
delete_df_collection(coll=df_collection)
diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py
index 8254d81..2a161e4 100644
--- a/tests/incite/mergers/foundations/test_enriched_session.py
+++ b/tests/incite/mergers/foundations/test_enriched_session.py
@@ -1,3 +1,6 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from itertools import product
@@ -5,10 +8,24 @@ from itertools import product
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+)
+from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+)
from generalresearch.incite.schemas.admin_responses import (
AdminPOPSessionSchema,
)
+from generalresearch.models.admin.request import (
+ ReportRequest,
+)
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
@@ -25,21 +42,20 @@ class TestEnrichedSession:
def test_base(
self,
- client_no_amm,
+ client_no_amm: DaskClient,
product: Product,
user_factory: Callable[..., User],
- wall_collection,
- session_collection,
- enriched_session_merge,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
thl_web_rr: PostgresConfig,
delete_df_collection: Callable[..., None],
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
- u1: User = user_factory(product=product: Product, created=session_collection.start)
+ u1: User = user_factory(product=product, created=session_collection.start)
for item in session_collection.items:
incite_item_factory(item=item, user=u1)
@@ -52,7 +68,7 @@ class TestEnrichedSession:
client=client_no_amm,
wall_coll=wall_collection,
session_coll=session_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
# --
@@ -85,16 +101,16 @@ class TestEnrichedSessionAdmin:
def test_to_admin_response(
self,
- event_report_request,
- enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
+ event_report_request: ReportRequest,
+ enriched_session_merge: EnrichedSessionMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
- session_report_request,
+ session_report_request: ReportRequest,
user_factory: Callable[..., User],
- start,
- session_factory,
+ start: datetime,
+ session_factory: Callable[..., Session],
product_factory: Callable[..., Product],
delete_df_collection: Callable[..., None],
):
@@ -107,7 +123,7 @@ class TestEnrichedSessionAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -120,7 +136,7 @@ class TestEnrichedSessionAdmin:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
df = enriched_session_merge.to_admin_response(
diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
index a33a55a..0606b6f 100644
--- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py
+++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
@@ -1,9 +1,28 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import timedelta
from itertools import product as iter_product
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+
+from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ TaskAdjustmentDFCollection,
+ WallDFCollection,
+)
+from generalresearch.incite.mergers.foundations.enriched_task_adjust import (
+ EnrichedTaskAdjustMerge,
+)
+from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
+)
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
+from generalresearch.pg_helper import PostgresConfig
@pytest.mark.parametrize(
@@ -20,19 +39,18 @@ class TestEnrichedTaskAdjust:
@pytest.mark.skip
def test_base(
self,
- client_no_amm,
+ client_no_amm: DaskClient,
user_factory: Callable[..., User],
product: Product,
- task_adj_collection,
- wall_collection,
- session_collection,
- enriched_wall_merge,
- enriched_task_adjust_merge,
- incite_item_factory,
+ task_adj_collection: TaskAdjustmentDFCollection,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ enriched_task_adjust_merge: EnrichedTaskAdjustMerge,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
@@ -48,14 +66,14 @@ class TestEnrichedTaskAdjust:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
enriched_task_adjust_merge.build(
client=client_no_amm,
task_adjust_coll=task_adj_collection,
enriched_wall=enriched_wall_merge,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
# --
diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py
index a0ca4dd..0cb8f60 100644
--- a/tests/incite/mergers/foundations/test_enriched_wall.py
+++ b/tests/incite/mergers/foundations/test_enriched_wall.py
@@ -1,3 +1,4 @@
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from itertools import product as iter_product
@@ -5,11 +6,23 @@ from itertools import product as iter_product
import dask.dataframe as dd
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+
+from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+)
# noinspection PyUnresolvedReferences
from generalresearch.incite.mergers.foundations.enriched_wall import (
+ EnrichedWallMerge,
EnrichedWallMergeItem,
)
+from generalresearch.models.admin.request import ReportRequest
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+from generalresearch.pg_helper import PostgresConfig
@pytest.mark.parametrize(
@@ -20,22 +33,21 @@ class TestEnrichedWall:
def test_base(
self,
- client_no_amm,
+ client_no_amm: DaskClient,
product: Product,
user_factory: Callable[..., User],
- wall_collection,
+ wall_collection: WallDFCollection,
thl_web_rr: PostgresConfig,
- session_collection,
- enriched_wall_merge,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
delete_df_collection: Callable[..., None],
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
delete_df_collection(coll=wall_collection)
- u1: User = user_factory(product=product: Product, created=session_collection.start)
+ u1: User = user_factory(product=product, created=session_collection.start)
for item in session_collection.items:
incite_item_factory(item=item, user=u1)
@@ -48,7 +60,7 @@ class TestEnrichedWall:
client=client_no_amm,
wall_coll=wall_collection,
session_coll=session_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
# --
@@ -63,19 +75,19 @@ class TestEnrichedWall:
def test_base_item(
self,
- client_no_amm,
+ client_no_amm: DaskClient,
product: Product,
user_factory: Callable[..., User],
- wall_collection,
- session_collection,
- enriched_wall_merge,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
delete_df_collection: Callable[..., None],
thl_web_rr: PostgresConfig,
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
):
# -- Build & Setup
delete_df_collection(coll=session_collection)
- u = user_factory(product=product: Product, created=session_collection.start)
+ u = user_factory(product=product, created=session_collection.start)
for item in session_collection.items:
incite_item_factory(item=item, user=u)
@@ -87,7 +99,7 @@ class TestEnrichedWall:
client=client_no_amm,
wall_coll=wall_collection,
session_coll=session_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
# --
@@ -99,14 +111,14 @@ class TestEnrichedWall:
try:
modified_time1 = path.stat().st_mtime
- except Exception:
+ except OSError:
modified_time1 = 0
item.build(
client=client_no_amm,
wall_coll=wall_collection,
session_coll=session_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
modified_time2 = path.stat().st_mtime
@@ -150,7 +162,12 @@ class TestEnrichedWallToAdmin:
def duration(self) -> timedelta | None:
return timedelta(days=5)
- def test_empty(self, enriched_wall_merge, client_no_amm, start):
+ def test_empty(
+ self,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ start: datetime,
+ ):
from generalresearch.models.admin.request import ReportRequest
rr = ReportRequest.model_validate({"interval": "5min", "start": start})
@@ -167,18 +184,18 @@ class TestEnrichedWallToAdmin:
def test_to_admin_response(
self,
- event_report_request,
- enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
+ event_report_request: ReportRequest,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
- user,
- session_factory,
+ user: User,
+ session_factory: Callable[..., Session],
delete_df_collection: Callable[..., None],
product_factory: Callable[..., Product],
user_factory: Callable[..., User],
- start,
+ start: datetime,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
@@ -189,7 +206,7 @@ class TestEnrichedWallToAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=2,
wall_req_cpi=Decimal("1.00"),
@@ -203,7 +220,7 @@ class TestEnrichedWallToAdmin:
client=client_no_amm,
wall_coll=wall_collection,
session_coll=session_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
df = enriched_wall_merge.to_admin_response(
diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py
index 15fa4db..cf8315f 100644
--- a/tests/incite/mergers/test_merge_collection.py
+++ b/tests/incite/mergers/test_merge_collection.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from datetime import UTC, datetime, timedelta
from itertools import product
@@ -5,12 +7,13 @@ import pandas as pd
import pytest
from pandera.pandas import DataFrameSchema
+from generalresearch.incite.base import GRLDatasets
from generalresearch.incite.mergers import (
MergeCollection,
MergeType,
)
-merge_types = list(e for e in MergeType if e != MergeType.TEST)
+merge_types = [e for e in MergeType if e != MergeType.TEST]
@pytest.mark.parametrize(
@@ -26,7 +29,11 @@ merge_types = list(e for e in MergeType if e != MergeType.TEST)
)
class TestMergeCollection:
- def test_init(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_init(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ ):
with pytest.raises(expected_exception=ValueError) as cm:
MergeCollection(archive_path=mnt_filepath.data_src)
assert "Must explicitly provide a merge_type" in str(cm.value)
@@ -37,7 +44,14 @@ class TestMergeCollection:
)
assert instance.merge_type == merge_type
- def test_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -48,7 +62,14 @@ class TestMergeCollection:
assert len(instance.interval_range) == len(instance.items)
- def test_progress(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_progress(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -62,7 +83,11 @@ class TestMergeCollection:
assert instance.progress.shape[1] == 7
assert instance.progress["group_by"].isnull().all()
- def test_schema(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_schema(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ ):
instance = MergeCollection(
merge_type=merge_type,
archive_path=mnt_filepath.archive_path(enum_type=merge_type),
@@ -70,7 +95,14 @@ class TestMergeCollection:
assert isinstance(instance._schema, DataFrameSchema)
- def test_load(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_load(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
start=start,
@@ -82,7 +114,14 @@ class TestMergeCollection:
# Confirm that there are no archives available yet
assert instance.progress.has_archive.eq(False).all()
- def test_get_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_get_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
start=start,
finished=start + duration,
diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py
index 3d0b644..5ca2f6b 100644
--- a/tests/incite/mergers/test_merge_collection_item.py
+++ b/tests/incite/mergers/test_merge_collection_item.py
@@ -1,10 +1,16 @@
+from __future__ import annotations
+
from datetime import timedelta
from itertools import product
from pathlib import PurePath
import pytest
-from generalresearch.incite.mergers import MergeCollectionItem, MergeType
+from generalresearch.incite.mergers import (
+ MergeCollection,
+ MergeCollectionItem,
+ MergeType,
+)
@pytest.mark.parametrize(
@@ -19,7 +25,10 @@ from generalresearch.incite.mergers import MergeCollectionItem, MergeType
)
class TestMergeCollectionItem:
- def test_file_naming(self, merge_collection, offset, duration, start):
+ def test_file_naming(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
items: list[MergeCollectionItem] = merge_collection.items
@@ -34,7 +43,10 @@ class TestMergeCollectionItem:
assert i._collection.offset in i.filename
assert i.start.strftime("%Y-%m-%d-%H-%M-%S") in i.filename
- def test_archives(self, merge_collection, offset, duration, start):
+ def test_archives(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
for i in merge_collection.items:
@@ -44,10 +56,13 @@ class TestMergeCollectionItem:
assert not i.has_partial_archive()
assert i.has_archive() == i.path_exists(generic_path=i.path)
- res = set([i.should_archive() for i in merge_collection.items])
+ res = {i.should_archive() for i in merge_collection.items}
assert len(res) == 1
- def test_item_to_archive(self, merge_collection, offset, duration, start):
+ def test_item_to_archive(
+ self,
+ merge_collection: MergeCollection,
+ ):
for item in merge_collection.items:
item: MergeCollectionItem
assert not item.has_archive()
diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py
index d054eb6..529a641 100644
--- a/tests/incite/mergers/test_pop_ledger.py
+++ b/tests/incite/mergers/test_pop_ledger.py
@@ -1,12 +1,25 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from itertools import product as iter_product
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+from generalresearch.incite.base import GRLDatasets
+from generalresearch.incite.collections.thl_web import (
+ LedgerDFCollection,
+ SessionDFCollection,
+)
+from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.incite.schemas.mergers.pop_ledger import (
numerical_col_names,
)
+from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
@pytest.mark.parametrize(
@@ -30,20 +43,20 @@ class TestMergePOPLedger:
def test_base(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
product: Product,
user_factory: Callable[..., User],
create_main_accounts: Callable[..., None],
- thl_lm,
+ thl_ledger_manager: ThlLedgerManager,
delete_df_collection: Callable[..., None],
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
delete_ledger_db: Callable[..., None],
):
from generalresearch.models.thl.ledger import LedgerAccount
- u = user_factory(product=product: Product, created=ledger_collection.start)
+ u = user_factory(product=product, created=ledger_collection.start)
# -- Build & Setup
delete_ledger_db()
@@ -73,19 +86,21 @@ class TestMergePOPLedger:
# --
- user_wallet_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet(
- user=u
+ user_wallet_account: LedgerAccount = (
+ thl_ledger_manager.get_account_or_create_user_wallet(user=u)
+ )
+ cash_account: LedgerAccount = thl_ledger_manager.get_account_cash()
+ rev_account: LedgerAccount = (
+ thl_ledger_manager.get_account_task_complete_revenue()
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
item_finishes = [i.finish for i in ledger_collection.items]
item_finishes.sort(reverse=True)
last_item_finish = item_finishes[0]
# Pure SQL based lookups
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert cash_balance > rev_balance
# (1) Test Cash Account
@@ -123,39 +138,42 @@ class TestMergePOPLedger:
def test_pydantic_init(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
product: Product,
user_factory: Callable[..., User],
create_main_accounts: Callable[..., None],
- offset,
- duration,
- start,
- thl_lm,
- incite_item_factory,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
+ incite_item_factory: Callable[..., None],
delete_df_collection: Callable[..., None],
delete_ledger_db: Callable[..., None],
- session_collection,
+ session_collection: SessionDFCollection,
):
from generalresearch.models.thl.finance import ProductBalances
from generalresearch.models.thl.ledger import LedgerAccount
from generalresearch.models.thl.product import Product
- u = user_factory(product=product: Product, created=session_collection.start)
+ u = user_factory(product=product, created=session_collection.start)
assert ledger_collection.finished is not None
- assert isinstance(u.product: Product, Product)
+ assert isinstance(u.product, Product)
delete_ledger_db()
- create_main_accounts(),
+ create_main_accounts()
+
delete_df_collection(coll=ledger_collection)
- bp_account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(
+ bp_account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet(
product=u.product
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
+ cash_account: LedgerAccount = thl_ledger_manager.get_account_cash()
+ rev_account: LedgerAccount = (
+ thl_ledger_manager.get_account_task_complete_revenue()
+ )
for item in ledger_collection.items:
incite_item_factory(item=item, user=u)
@@ -185,8 +203,10 @@ class TestMergePOPLedger:
assert instance.payout == instance.net == instance.bp_payment_credit
assert instance.available_balance < instance.net
assert instance.available_balance + instance.retainer == instance.net
- assert instance.balance == thl_lm.get_account_balance(bp_account)
- assert df["bp_payment.CREDIT"].sum() == thl_lm.get_account_balance(bp_account)
+ assert instance.balance == thl_ledger_manager.get_account_balance(bp_account)
+ assert df["bp_payment.CREDIT"].sum() == thl_ledger_manager.get_account_balance(
+ bp_account
+ )
# (2) Filter by the Cash Account
ddf = pop_ledger_merge.ddf(
@@ -199,7 +219,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
assert df["bp_payment.CREDIT"].sum() == 0
assert cash_balance > 0
assert df["mp_payment.CREDIT"].sum() == 0
@@ -216,7 +236,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert rev_balance == 0
assert df["bp_payment.CREDIT"].sum() == 0
assert df["mp_payment.DEBIT"].sum() == 0
@@ -224,27 +244,28 @@ class TestMergePOPLedger:
def test_resample(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
user_factory: Callable[..., User],
product: Product,
create_main_accounts: Callable[..., None],
- offset,
- duration,
- start,
- thl_lm,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
delete_df_collection: Callable[..., None],
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
assert ledger_collection.finished is not None
delete_df_collection(coll=ledger_collection)
u1: User = user_factory(product=product)
- bp_account = thl_lm.get_account_or_create_bp_wallet(product=u1.product)
+ bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=u1.product
+ )
for item in ledger_collection.items:
incite_item_factory(user=u1, item=item)
@@ -274,7 +295,7 @@ class TestMergePOPLedger:
assert isinstance(df.index, pd.Index)
assert isinstance(df.index, pd.DatetimeIndex)
- bp_account_balance = thl_lm.get_account_balance(account=bp_account)
+ bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account)
# Initial sum
initial_sum = df.sum().sum()
diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py
index a0b8b87..8a4897b 100644
--- a/tests/incite/mergers/test_ym_survey_merge.py
+++ b/tests/incite/mergers/test_ym_survey_merge.py
@@ -1,8 +1,24 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from datetime import UTC, datetime, timedelta
from itertools import product
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
+
+from generalresearch.incite.collections.thl_web import (
+ SessionDFCollection,
+ WallDFCollection,
+)
+from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+)
+from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge
+from generalresearch.models.thl.product import Product
+from generalresearch.models.thl.user import User
+from generalresearch.pg_helper import PostgresConfig
# noinspection PyUnresolvedReferences
@@ -27,21 +43,20 @@ class TestYMSurveyMerge:
def test_base(
self,
- client_no_amm,
+ client_no_amm: DaskClient,
user_factory: Callable[..., User],
product: Product,
- ym_survey_wall_merge,
- wall_collection,
- session_collection,
- enriched_session_merge,
+ ym_survey_wall_merge: YMSurveyWallMerge,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
delete_df_collection: Callable[..., None],
- incite_item_factory,
+ incite_item_factory: Callable[..., None],
thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
- user: User = user_factory(product=product: Product, created=session_collection.start)
+ user: User = user_factory(product=product, created=session_collection.start)
# -- Build & Setup
assert ym_survey_wall_merge.start is None
@@ -61,15 +76,15 @@ class TestYMSurveyMerge:
client=client_no_amm,
session_coll=session_collection,
wall_coll=wall_collection,
- pg_config=thl_web_rr: PostgresConfig,
+ pg_config=thl_web_rr,
)
assert enriched_session_merge.progress.has_archive.eq(True).all()
ddf = enriched_session_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df1: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df1, pd.DataFrame)
+ assert not df1.empty
# --
@@ -83,18 +98,18 @@ class TestYMSurveyMerge:
# --
ddf = ym_survey_wall_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df2: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df2, pd.DataFrame)
+ assert not df2.empty
# --
- assert df.product_id.nunique() == 1
- assert df.team_id.nunique() == 1
- assert df.source.nunique() > 1
+ assert df2.product_id.nunique() == 1
+ assert df2.team_id.nunique() == 1
+ assert df2.source.nunique() > 1
- started_min_ts = df.started.min()
- started_max_ts = df.started.max()
+ started_min_ts = df2.started.min()
+ started_max_ts = df2.started.max()
assert type(started_min_ts) is pd.Timestamp
assert type(started_max_ts) is pd.Timestamp
diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py
index e98eecd..d2658ea 100644
--- a/tests/incite/schemas/test_admin_responses.py
+++ b/tests/incite/schemas/test_admin_responses.py
@@ -1,8 +1,11 @@
+from __future__ import annotations
+
from datetime import UTC, datetime, timedelta
from random import sample
import numpy as np
import pandas as pd
+import pandera as pa
import pytest
from generalresearch.incite.schemas import empty_dataframe_from_schema
@@ -16,12 +19,14 @@ from generalresearch.locales import Localelator
class TestAdminPOPSchema:
schema_df = empty_dataframe_from_schema(AdminPOPSchema)
countries = list(Localelator().get_all_countries())[:5]
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10) # noqa
+ ]
@classmethod
def assign_valid_vals(cls, df: pd.DataFrame) -> pd.DataFrame:
for c in df.columns:
- check_attrs: dict = AdminPOPSchema.columns[c].checks[0].statistics
+ check_attrs = AdminPOPSchema.columns[c].checks[0].statistics
df[c] = np.random.randint(
check_attrs["min_value"], check_attrs["max_value"], df.shape[0]
)
@@ -29,7 +34,7 @@ class TestAdminPOPSchema:
return df
def test_empty(self):
- with pytest.raises(Exception):
+ with pytest.raises(pa.errors.SchemaError):
AdminPOPSchema.validate(pd.DataFrame())
def test_new_empty_df(self):
@@ -42,7 +47,7 @@ class TestAdminPOPSchema:
def test_valid(self):
# (1) Works with raw naive datetime
dates = [
- datetime(year=2024, month=1, day=i, tzinfo=None).isoformat()
+ datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() # noqa
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -57,7 +62,10 @@ class TestAdminPOPSchema:
assert isinstance(df, pd.DataFrame)
# (2) Works with isoformat naive datetime
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) # noqa
+ for i in range(1, 10)
+ ]
df = pd.DataFrame(
index=pd.MultiIndex.from_product(
iterables=[dates, self.countries], names=["index0", "index1"]
@@ -84,12 +92,12 @@ class TestAdminPOPSchema:
# Initially, they're all set with a timezone
timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz == UTC for ts in timestmaps])
+ assert all(ts.tz == UTC for ts in timestmaps)
# After validation, the timezone is removed
df = AdminPOPSchema.validate(df)
timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ assert all(ts.tz is None for ts in timestmaps)
def test_index_tz_no_future_beyond_one_year(self):
now = datetime.now(tz=UTC)
@@ -123,12 +131,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, float) for v in vals])
+ assert all(isinstance(v, float) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# --- int to str ---
@@ -142,12 +150,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, int) for v in vals])
+ assert all(isinstance(v, int) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# a = 1
assert isinstance(df, pd.DataFrame)
@@ -170,7 +178,7 @@ class TestAdminPOPSchema:
assert isinstance(df, pd.DataFrame)
timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ assert all(ts.tz is None for ts in timestmaps)
# (2) Timezones are removed
dates = [
@@ -187,12 +195,12 @@ class TestAdminPOPSchema:
# Has tz before validation, and none after
timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is UTC for ts in timestmaps])
+ assert all(ts.tz is UTC for ts in timestmaps)
df = AdminPOPSchema.validate(df, lazy=True)
timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)]
- assert all([ts.tz is None for ts in timestmaps])
+ assert all(ts.tz is None for ts in timestmaps)
def test_clipping(self):
df = pd.DataFrame(
diff --git a/tests/incite/schemas/test_thl_web.py b/tests/incite/schemas/test_thl_web.py
index 7f4434b..9b34ce0 100644
--- a/tests/incite/schemas/test_thl_web.py
+++ b/tests/incite/schemas/test_thl_web.py
@@ -16,7 +16,7 @@ class TestWallSchema:
df = pd.DataFrame(columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_no_rows(self):
@@ -24,7 +24,7 @@ class TestWallSchema:
df = pd.DataFrame(index=["uuid"], columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_new_empty_df(self):
@@ -50,7 +50,7 @@ class TestSessionSchema:
df = pd.DataFrame(columns=THLSessionSchema.columns.keys())
df.set_index("uuid", inplace=True)
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_no_rows(self):
@@ -58,7 +58,7 @@ class TestSessionSchema:
df = pd.DataFrame(index=["id"], columns=THLSessionSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_new_empty_df(self):
diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py
index 7e1577a..d6ce2b1 100644
--- a/tests/incite/test_collection_base.py
+++ b/tests/incite/test_collection_base.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from datetime import UTC, datetime, timedelta, timezone
from os.path import exists as pexists
from os.path import join as pjoin
@@ -9,7 +11,7 @@ import pandas as pd
import pytest
from _pytest._code.code import ExceptionInfo
-from generalresearch.incite.base import CollectionBase
+from generalresearch.incite.base import CollectionBase, GRLDatasets
AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0)
AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0)
@@ -17,11 +19,11 @@ AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0)
class TestCollectionBase:
- def test_init(self, mnt_filepath):
+ def test_init(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.df.empty is True
- def test_init_df(self, mnt_filepath):
+ def test_init_df(self, mnt_filepath: GRLDatasets):
# Only an empty pd.DataFrame can ever be provided
instance = CollectionBase(
df=pd.DataFrame({}), archive_path=mnt_filepath.data_src
@@ -43,7 +45,7 @@ class TestCollectionBase:
)
assert "Do not provide a pd.DataFrame" in str(cm.value)
- def test_init_start(self, mnt_filepath):
+ def test_init_start(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(
@@ -74,7 +76,7 @@ class TestCollectionBase:
cm.value
)
- def test_init_archive_path(self, mnt_filepath):
+ def test_init_archive_path(self, mnt_filepath: GRLDatasets):
"""DirectoryPath is apparently smart enough to confirm that the
directory path exists.
"""
@@ -99,7 +101,7 @@ class TestCollectionBase:
CollectionBase(archive_path=new_path)
assert "Path does not point to a directory" in str(cm.value)
- def test_init_offset(self, mnt_filepath):
+ def test_init_offset(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(offset="1:X", archive_path=mnt_filepath.data_src)
@@ -118,14 +120,14 @@ class TestCollectionBase:
class TestCollectionBaseProperties:
- def test_items(self, mnt_filepath):
+ def test_items(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- x = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
- def test_interval_range(self, mnt_filepath):
+ def test_interval_range(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
# Private method requires the end parameter
with pytest.raises(expected_exception=AssertionError) as cm:
@@ -147,7 +149,7 @@ class TestCollectionBaseProperties:
assert res.is_monotonic_increasing
assert res.is_unique
- def test_interval_range2(self, mnt_filepath):
+ def test_interval_range2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert isinstance(instance.interval_range, list)
@@ -166,16 +168,16 @@ class TestCollectionBaseProperties:
)
assert len(instance.interval_range) == 2
- def test_progress(self, mnt_filepath):
+ def test_progress(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(
start=AGO_15min, offset="3min", archive_path=mnt_filepath.data_src
)
- x = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_progress2(self, mnt_filepath):
+ def test_progress2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
@@ -184,10 +186,10 @@ class TestCollectionBaseProperties:
assert instance.df.empty
with pytest.raises(expected_exception=NotImplementedError) as cm:
- df = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_items2(self, mnt_filepath):
+ def test_items2(self, mnt_filepath: GRLDatasets):
"""There can't be a test for this because the Items need a path whic
isn't possible in the generic form
"""
@@ -197,7 +199,7 @@ class TestCollectionBaseProperties:
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
- items = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
# item = items[-3]
@@ -208,19 +210,19 @@ class TestCollectionBaseProperties:
# assert str(df.product_id.dtype) == "object"
# assert str(ddf.product_id.dtype) == "string"
- def test_items3(self, mnt_filepath):
+ def test_items3(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
archive_path=mnt_filepath.data_src,
)
with pytest.raises(expected_exception=NotImplementedError) as cm:
- item = instance.items[0]
+ _ = instance.items[0]
assert "Must override" in str(cm.value)
class TestCollectionBaseMethodsCleanup:
- def test_fetch_force_rr_latest(self, mnt_filepath):
+ def test_fetch_force_rr_latest(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=Exception) as cm:
@@ -228,7 +230,7 @@ class TestCollectionBaseMethodsCleanup:
coll.fetch_force_rr_latest(sources=[])
assert "Must override" in str(cm.value)
- def test_fetch_all_paths(self, mnt_filepath):
+ def test_fetch_all_paths(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
@@ -242,16 +244,16 @@ class TestCollectionBaseMethodsCleanup:
class TestCollectionBaseMethodsCleanup:
@pytest.mark.skip
- def test_cleanup_partials(self, mnt_filepath):
+ def test_cleanup_partials(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.cleanup_partials() is None # it doesn't return anything
- def test_clear_tmp_archives(self, mnt_filepath):
+ def test_clear_tmp_archives(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.clear_tmp_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_clear_corrupt_archives(self, mnt_filepath):
+ def test_clear_corrupt_archives(self, mnt_filepath: GRLDatasets):
"""TODO: expand this so it actually has corrupt archives that we
check to see if they're removed
"""
@@ -259,14 +261,14 @@ class TestCollectionBaseMethodsCleanup:
assert instance.clear_corrupt_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_rebuild_symlinks(self, mnt_filepath):
+ def test_rebuild_symlinks(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.rebuild_symlinks() is None
class TestCollectionBaseMethodsSourceTiming:
- def test_get_item(self, mnt_filepath):
+ def test_get_item(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
i = pd.Interval(left=1, right=2, closed="left")
@@ -274,7 +276,7 @@ class TestCollectionBaseMethodsSourceTiming:
instance.get_item(interval=i)
assert "Must override" in str(cm.value)
- def test_get_item_start(self, mnt_filepath):
+ def test_get_item_start(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
dt = datetime.now(tz=UTC)
@@ -284,7 +286,7 @@ class TestCollectionBaseMethodsSourceTiming:
instance.get_item_start(start=start)
assert "Must override" in str(cm.value)
- def test_get_items(self, mnt_filepath):
+ def test_get_items(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
dt = datetime.now(tz=UTC)
@@ -293,21 +295,21 @@ class TestCollectionBaseMethodsSourceTiming:
instance.get_items(since=dt)
assert "Must override" in str(cm.value)
- def test_get_items_from_year(self, mnt_filepath):
+ def test_get_items_from_year(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_from_year(year=2020)
assert "Must override" in str(cm.value)
- def test_get_items_last90(self, mnt_filepath):
+ def test_get_items_last90(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_last90()
assert "Must override" in str(cm.value)
- def test_get_items_last365(self, mnt_filepath):
+ def test_get_items_last365(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py
index 7a0a581..e09f54a 100644
--- a/tests/incite/test_collection_base_item.py
+++ b/tests/incite/test_collection_base_item.py
@@ -1,3 +1,5 @@
+from __future__ import annotations
+
from datetime import UTC, datetime
from os.path import join as pjoin
from pathlib import Path
@@ -8,7 +10,7 @@ import pandas as pd
import pytest
from pydantic import ValidationError
-from generalresearch.incite.base import CollectionItemBase
+from generalresearch.incite.base import CollectionItemBase, GRLDatasets
class TestCollectionItemBase:
@@ -40,20 +42,20 @@ class TestCollectionItemBaseProperties:
def test_finish(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.finish
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.finish
def test_interval(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.interval
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.interval
def test_filename(self):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -61,7 +63,7 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -69,27 +71,27 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.path
def test_partial_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.partial_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.partial_path
def test_empty_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.empty_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.empty_path
class TestCollectionItemBaseMethods:
@@ -106,41 +108,41 @@ class TestCollectionItemBaseMethods:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.tmp_filename()
+ _ = instance.tmp_filename()
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_tmp_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.tmp_path()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.tmp_path()
def test_is_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.is_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.is_empty()
def test_has_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_empty()
def test_has_partial_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_partial_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_partial_archive()
@pytest.mark.parametrize("include_empty", [True, False])
- def test_has_archive(self, include_empty):
+ def test_has_archive(self, include_empty: bool):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_archive(include_empty=include_empty)
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_archive(include_empty=include_empty)
- def test_delete_archive_file(self, mnt_filepath):
+ def test_delete_archive_file(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}.zip"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -155,7 +157,7 @@ class TestCollectionItemBaseMethods:
CollectionItemBase.delete_archive(generic_path=path1)
assert not path1.exists()
- def test_delete_archive_dir(self, mnt_filepath):
+ def test_delete_archive_dir(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -174,20 +176,20 @@ class TestCollectionItemBaseMethods:
def test_should_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.should_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.should_archive()
def test_set_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.set_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.set_empty()
def test_valid_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.valid_archive(generic_path=None, sample=None)
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.valid_archive(generic_path=None, sample=None)
class TestCollectionItemBaseMethodsORM:
@@ -197,11 +199,11 @@ class TestCollectionItemBaseMethodsORM:
pass
@pytest.mark.parametrize("is_partial", [True, False])
- def test_to_archive(self, is_partial):
+ def test_to_archive(self, is_partial: bool):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.to_archive(
+ _ = instance.to_archive(
ddf=dd.from_pandas(data=pd.DataFrame()), is_partial=is_partial
)
assert "Must override" in str(cm.value)
diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py
index 3034c21..03d29ea 100644
--- a/tests/incite/test_interval_idx.py
+++ b/tests/incite/test_interval_idx.py
@@ -1,4 +1,4 @@
-from datetime import datetime
+from datetime import UTC, datetime
import pandas as pd
@@ -6,8 +6,8 @@ import pandas as pd
class TestIntervalIndex:
def test_init(self):
- start = datetime(year=2000, month=1, day=1)
- end = datetime(year=2000, month=1, day=10)
+ start = datetime(year=2000, month=1, day=1, tzinfo=UTC)
+ end = datetime(year=2000, month=1, day=10, tzinfo=UTC)
iv_r: pd.IntervalIndex = pd.interval_range(
start=start, end=end, freq="1d", closed="left"
diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py
index 3490403..ed141b1 100644
--- a/tests/managers/gr/test_business.py
+++ b/tests/managers/gr/test_business.py
@@ -2,17 +2,36 @@ from uuid import uuid4
import pytest
+from generalresearch.managers.gr.business import (
+ BusinessAddressManager,
+ BusinessBankAccountManager,
+ BusinessManager,
+)
+from generalresearch.managers.gr.team import MembershipManager, TeamManager
+from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+ BusinessBankAccount,
+ TransferMethod,
+)
+from generalresearch.pg_helper import PostgresConfig
+
class TestBusinessBankAccountManager:
- def test_init(self, business_bank_account_manager, gr_db):
+ def test_init(
+ self,
+ business_bank_account_manager: BusinessBankAccountManager,
+ gr_db: PostgresConfig,
+ ):
assert business_bank_account_manager.pg_config == gr_db
- def test_create(self, business: Business, business_bank_account_manager):
- from generalresearch.models.gr.business import (
- BusinessBankAccount,
- TransferMethod,
- )
+ def test_create(
+ self,
+ business: Business,
+ business_bank_account_manager: BusinessBankAccountManager,
+ ):
instance = business_bank_account_manager.create(
business_id=business.id,
@@ -33,8 +52,9 @@ class TestBusinessBankAccountManager:
class TestBusinessAddressManager:
- def test_create(self, business: Business, business_address_manager):
- from generalresearch.models.gr.business import BusinessAddress
+ def test_create(
+ self, business: Business, business_address_manager: BusinessAddressManager
+ ):
res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id)
assert isinstance(res, BusinessAddress)
@@ -43,14 +63,13 @@ class TestBusinessAddressManager:
class TestBusinessManager:
- def test_create(self, business_manager):
- from generalresearch.models.gr.business import Business
+ def test_create(self, business_manager: BusinessManager):
instance = business_manager.create_dummy()
assert isinstance(instance, Business)
assert isinstance(instance.id, int)
- def test_get_or_create(self, business_manager):
+ def test_get_or_create(self, business_manager: BusinessManager):
uuid_key = uuid4().hex
assert business_manager.get_by_uuid(business_uuid=uuid_key) is None
@@ -61,9 +80,10 @@ class TestBusinessManager:
)
res = business_manager.get_by_uuid(business_uuid=uuid_key)
+ assert isinstance(res, Business)
assert res.id == instance.id
- def test_get_all(self, business_manager):
+ def test_get_all(self, business_manager: BusinessManager):
res1 = business_manager.get_all()
assert isinstance(res1, list)
@@ -76,7 +96,11 @@ class TestBusinessManager:
pass
def test_get_by_user_id(
- self, business_manager, gr_user, team_manager, membership_manager
+ self,
+ business_manager: BusinessManager,
+ gr_user: GRUser,
+ team_manager: TeamManager,
+ membership_manager: MembershipManager,
):
res = business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
@@ -93,7 +117,7 @@ class TestBusinessManager:
# Create a Membership for the gr_user to the Team... but it doesn't
# matter because the Team doesn't have any Business yet
- m1 = membership_manager.create(team=t1, gr_user=gr_user)
+ _ = membership_manager.create(team=t1, gr_user=gr_user)
res = business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 0
@@ -113,15 +137,17 @@ class TestBusinessManager:
def test_get_uuids_by_user_id(self):
pass
- def test_get_by_uuid(self, business: Business, business_manager):
+ def test_get_by_uuid(self, business: Business, business_manager: BusinessManager):
instance = business_manager.get_by_uuid(business_uuid=business.uuid)
+ assert isinstance(instance, Business)
assert business.id == instance.id
- def test_get_by_id(self, business: Business, business_manager):
+ def test_get_by_id(self, business: Business, business_manager: BusinessManager):
instance = business_manager.get_by_id(business_id=business.id)
+ assert isinstance(instance, Business)
assert business.uuid == instance.uuid
- def test_cache_key(self, business):
+ def test_cache_key(self, business: Business):
assert "business:" in business.cache_key
# def test_create_raise_on_duplicate(self):
diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py
index 5e5c565..ae3e1bb 100644
--- a/tests/managers/gr/test_team.py
+++ b/tests/managers/gr/test_team.py
@@ -1,18 +1,29 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from uuid import uuid4
+from generalresearch.managers.gr.authentication import GRUserManager
+from generalresearch.managers.gr.team import MembershipManager, TeamManager
+from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.team import Membership, Team
+from generalresearch.models.thl.product import Product
+from generalresearch.pg_helper import PostgresConfig
+from generalresearch.redis_helper import RedisConfig
+
class TestMembershipManager:
- def test_init(self, membership_manager, gr_db):
+ def test_init(self, membership_manager: MembershipManager, gr_db: PostgresConfig):
assert membership_manager.pg_config == gr_db
class TestTeamManager:
- def test_init(self, team_manager, gr_db):
+ def test_init(self, team_manager: TeamManager, gr_db: PostgresConfig):
assert team_manager.pg_config == gr_db
- def test_get_or_create(self, team_manager):
+ def test_get_or_create(self, team_manager: TeamManager):
from generalresearch.models.gr.team import Team
new_uuid = uuid4().hex
@@ -24,7 +35,7 @@ class TestTeamManager:
assert team.uuid == new_uuid
assert team.name == "< Unknown >"
- def test_get_all(self, team_manager):
+ def test_get_all(self, team_manager: TeamManager):
res1 = team_manager.get_all()
assert isinstance(res1, list)
@@ -32,16 +43,20 @@ class TestTeamManager:
res2 = team_manager.get_all()
assert len(res1) == len(res2) - 1
- def test_create(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_create(self, team_manager: TeamManager):
team: Team = team_manager.create_dummy()
assert isinstance(team, Team)
assert isinstance(team.id, int)
- def test_add_user(self, team, team_manager, gr_um, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Membership
+ def test_add_user(
+ self,
+ team: Team,
+ team_manager: TeamManager,
+ gr_um: GRUserManager,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
user: GRUser = gr_um.create_dummy()
@@ -54,25 +69,23 @@ class TestTeamManager:
assert len(team.gr_users)
assert team.gr_users == [user]
- def test_get_by_uuid(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_get_by_uuid(self, team_manager: TeamManager):
team: Team = team_manager.create_dummy()
instance = team_manager.get_by_uuid(team_uuid=team.uuid)
assert team.id == instance.id
- def test_get_by_id(self, team_manager):
- from generalresearch.models.gr.team import Team
+ def test_get_by_id(self, team_manager: TeamManager):
team: Team = team_manager.create_dummy()
instance = team_manager.get_by_id(team_id=team.id)
assert team.uuid == instance.uuid
- def test_get_by_user(self, team, team_manager, gr_um):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Team
+ def test_get_by_user(
+ self, team: Team, team_manager: TeamManager, gr_um: GRUserManager
+ ):
user: GRUser = gr_um.create_dummy()
team_manager.add_user(team=team, gr_user=user)
@@ -86,15 +99,12 @@ class TestTeamManager:
def test_get_by_user_duplicates(
self,
- gr_user_token,
- gr_user,
- membership,
+ gr_user: GRUser,
product_factory: Callable[..., Product],
- membership_factory,
- team,
- thl_web_rr: PostgresConfig,
- gr_redis_config,
- gr_db,
+ membership_factory: Callable[..., Membership],
+ team: Team,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
):
product_factory(team=team)
membership_factory(team=team, gr_user=gr_user)
diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py
index bfc7518..71efa95 100644
--- a/tests/managers/network/test_label.py
+++ b/tests/managers/network/test_label.py
@@ -27,7 +27,7 @@ def ip_label(utc_now) -> IPLabel:
provider="GeoNodE",
created_at=utc_now,
ip=ip,
- metadata=IPLabelMetadata(services=["RDP"])
+ metadata=IPLabelMetadata(services=["RDP"]),
)
@@ -181,7 +181,7 @@ def test_label_cidr_and_ipinfo(
ip = fake.ipv6()
ip_information_factory(ip=ip, geoname=ip_geoname)
# We normalize for storage into ipinfo table
- ip_norm, prefix = normalize_ip(ip)
+ ip_norm, _ = normalize_ip(ip)
# Test with a larger network
ip_48 = ipaddress.IPv6Network((ip, 48), strict=False)
diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py
index 07d8d74..3a63075 100644
--- a/tests/managers/thl/test_contest/test_leaderboard.py
+++ b/tests/managers/thl/test_contest/test_leaderboard.py
@@ -5,7 +5,6 @@ from zoneinfo import ZoneInfo
from generalresearch.currency import USDCent
from generalresearch.managers.thl.contest_manager import ContestManager
-from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.managers.thl.user_manager.user_manager import UserManager
from generalresearch.models.thl.contest.definitions import (
@@ -116,7 +115,9 @@ class TestLeaderboardContestCRUD:
assert decision
assert reason == ContestEndReason.ENDS_AT
- contest_manager.end_contest_if_over(contest=contest, ledger_manager=thl_lm)
+ contest_manager.end_contest_if_over(
+ contest=contest, ledger_manager=thl_ledger_manager
+ )
c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid)
assert c.status == ContestStatus.COMPLETED
diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py
index a2d575b..e3889bc 100644
--- a/tests/managers/thl/test_contest/test_milestone.py
+++ b/tests/managers/thl/test_contest/test_milestone.py
@@ -292,7 +292,10 @@ class TestMilestoneContestUserViews:
assert len(cs) == 1
contest_manager.enter_milestone_contest(
- contest_uuid=c.uuid, user=user, country_iso="us", ledger_manager=thl_lm
+ contest_uuid=c.uuid,
+ user=user,
+ country_iso="us",
+ ledger_manager=thl_ledger_manager,
)
# User isn't eligible anymore
diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py
index b435576..06d4676 100644
--- a/tests/managers/thl/test_contest/test_raffle.py
+++ b/tests/managers/thl/test_contest/test_raffle.py
@@ -14,6 +14,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
)
from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
from generalresearch.models.thl.contest import (
+ Contest,
ContestEndCondition,
ContestEntryRule,
ContestPrize,
@@ -40,8 +41,6 @@ class TestRaffleContest:
def test_should_end(
self,
contest: RaffleContest,
- thl_ledger_manager: ThlLedgerManager,
- contest_manager: ContestManager,
):
# contest is active and has no entries
should, msg = contest.should_end()
@@ -67,7 +66,6 @@ class TestRaffleContestCRUD:
self,
contest_create: RaffleContestCreate,
product_user_wallet_yes: Product,
- thl_ledger_manager: ThlLedgerManager,
contest_manager: ContestManager,
):
c = contest_manager.create(
@@ -329,7 +327,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user." in str(e.value)
@@ -342,7 +340,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
@@ -354,7 +352,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# Then can't anymore
@@ -366,7 +364,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py
index 540bea8..7b65b2d 100644
--- a/tests/managers/thl/test_ledger/test_lm_accounts.py
+++ b/tests/managers/thl/test_ledger/test_lm_accounts.py
@@ -63,7 +63,7 @@ class TestLedgerAccountManagerNoResults:
acct_id: UUIDStr,
lm: LedgerManager,
):
- qn = ":".join([currency, kind, acct_id])
+ qn = f"{currency}:{kind}:{acct_id}"
# (1) .get_many_
assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == []
diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py
index 13495a7..ce609d6 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx.py
@@ -24,8 +24,6 @@ class TestLedgerManagerCreateTx:
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
"""
- acct_uuid = uuid4().hex
-
# (1) With no Permissions defined
test_lm = LedgerManager(
pg_config=ledger_manager.pg_config,
@@ -44,8 +42,6 @@ class TestLedgerManagerCreateTx:
def test_create_assertions(
self,
- ledger_account_debit: LedgerAccount,
- ledger_account_credit: LedgerAccount,
ledger_manager: LedgerManager,
):
with pytest.raises(expected_exception=ValueError) as excinfo:
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
index dce9116..60eb71c 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py
@@ -229,6 +229,7 @@ class TestThlLedgerManagerAccounts:
# (1) known account and confirm it comes back
res = ledger_manager.get_account(qualified_name=account1.qualified_name)
+ assert isinstance(res, LedgerAccount)
assert account1.model_dump_json() == res.model_dump_json()
# (2) known accounts and confirm they both come back
@@ -291,6 +292,7 @@ class TestThlLedgerManagerAccounts:
assert len(res) == 2
# Confirm an empty array comes back for all unknown qualified names
+ assert isinstance(ledger_manager.currency, LedgerCurrency)
res = ledger_manager.get_accounts_if_exists(
qualified_names=[
f"{ledger_manager.currency.value}:bp_wall:{uuid4().hex}"
@@ -328,7 +330,7 @@ class TestThlLedgerManagerAccounts:
product_uuids=product_uuids
)
assert len(res) == len(product_uuids)
- assert all([isinstance(i, LedgerAccount) for i in res])
+ assert all(isinstance(i, LedgerAccount) for i in res)
class TestLedgerAccountManager:
@@ -351,10 +353,10 @@ class TestLedgerAccountManager:
# First we want to validate that using the get_account method raises
# an error for a random LedgerAccount which we know does not exist.
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_account(qualified_name=account.qualified_name)
+ ledger_account_manager.get_account(qualified_name=account.qualified_name)
# Now that we know it doesn't exist, get_or_create for it
- instance = lam.get_account_or_create(account=account)
+ instance = ledger_account_manager.get_account_or_create(account=account)
# It should always return
assert isinstance(instance, LedgerAccount)
@@ -364,10 +366,11 @@ class TestLedgerAccountManager:
self,
user: User,
thl_ledger_manager: ThlLedgerManager,
- ledger_manager: LedgerManager,
ledger_account_manager: LedgerAccountManager,
):
+ assert isinstance(user.product, Product)
+
with pytest.raises(LedgerAccountDoesntExistError):
ledger_account_manager.get_account(
qualified_name=f"test:bp_wallet:{user.product.id}"
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
index e4a25a3..b518453 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py
@@ -133,16 +133,17 @@ class TestThlLedgerManagerBPPayout:
)
payoutevent_uuid = uuid4().hex
- with caplog.at_level(logging.INFO):
- with pytest.raises(LedgerTransactionConditionFailedError):
- thl_ledger_manager.create_tx_bp_payout(
- user.product,
- amount=USDCent(10_000),
- created=now + timedelta(minutes=2),
- skip_one_per_day_check=True,
- skip_wallet_balance_check=False,
- payoutevent_uuid=payoutevent_uuid,
- )
+ with caplog.at_level(logging.INFO), pytest.raises(
+ LedgerTransactionConditionFailedError
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ user.product,
+ amount=USDCent(10_000),
+ created=now + timedelta(minutes=2),
+ skip_one_per_day_check=True,
+ skip_wallet_balance_check=False,
+ payoutevent_uuid=payoutevent_uuid,
+ )
assert "failed condition check balance:" in caplog.text
thl_ledger_manager.create_tx_bp_payout(
@@ -197,17 +198,18 @@ class TestThlLedgerManagerBPPayout:
assert balance == int(rand_amount) * -1
# Test some basic assertions
- with caplog.at_level(logging.INFO):
- with pytest.raises(expected_exception=Exception):
- thl_ledger_manager.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=UTC),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with caplog.at_level(logging.INFO), pytest.raises(
+ expected_exception=ValueError
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
def test_create_tx_redis_failure(
@@ -291,7 +293,7 @@ class TestThlLedgerManagerBPPayout:
# Will fail due to multiple per day
payoutevent_uuid2 = uuid4().hex
with pytest.raises(expected_exception=Exception) as e:
- tx = thl_ledger_manager.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid2,
@@ -348,7 +350,7 @@ class TestThlLedgerManagerBPPayout:
# Create TX will fail on lock exit, after the tx was created!
with pytest.raises(expected_exception=Exception) as e:
- tx = thl_ledger_manager.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=product,
amount=rand_amount,
payoutevent_uuid=payoutevent_uuid,
@@ -384,7 +386,9 @@ class TestPayoutEventManagerBPPayout:
product, rand_amount, now, direction=Direction.CREDIT
)
assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ brokerage_product_payout_event_manager.set_account_lookup_table(
+ thl_lm=thl_ledger_manager
+ )
pe = brokerage_product_payout_event_manager.create_bp_payout_event(
thl_ledger_manager=thl_ledger_manager,
@@ -557,7 +561,7 @@ class TestPayoutEventManagerBPPayout:
# Will fail on lock exit, after the tx was created!
# But it'll see that the tx was created and so everything will be fine
Lock.release = broken_release
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
+ brokerage_product_payout_event_manager.create_bp_payout_event(
thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py
index 89adb0b..1860d6d 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py
@@ -23,7 +23,6 @@ from generalresearch.models.thl.definitions import (
WALL_ALLOWED_STATUS_STATUS_CODE,
)
from generalresearch.models.thl.ledger import (
- AccountType,
Direction,
LedgerAccount,
TransactionType,
@@ -287,17 +286,18 @@ class TestThlLedgerTxManager:
assert balance == int(rand_amount) * -1
# Test some basic assertions
- with caplog.at_level(logging.INFO):
- with pytest.raises(expected_exception=Exception):
- thl_ledger_manager.create_tx_bp_payout(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=UTC),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with caplog.at_level(logging.INFO), pytest.raises(
+ expected_exception=ValueError
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
def test_create_tx_bp_payout_(
@@ -1794,7 +1794,7 @@ class TestThlLedgerManagerAdj:
)
thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started)
- revenue = ththl_ledger_managerl_lm.get_account_task_complete_revenue()
+ revenue = thl_ledger_manager.get_account_task_complete_revenue()
bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet(
user.product
)
diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py
index e6c597b..0f3f103 100644
--- a/tests/managers/thl/test_payout.py
+++ b/tests/managers/thl/test_payout.py
@@ -727,8 +727,6 @@ class TestBusinessPayoutEventManager:
# {"uuid": bp_pe.uuid, "status": PayoutStatus.FAILED},
# )
- assert 1 == 0
-
def test_ach_payment(
self,
mnt_filepath: GRLDatasets,
diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py
index 2c2bf9d..c3ab162 100644
--- a/tests/managers/thl/test_survey.py
+++ b/tests/managers/thl/test_survey.py
@@ -11,9 +11,6 @@ from generalresearch.managers.thl.buyer import BuyerManager
from generalresearch.managers.thl.profiling.question import (
QuestionManager,
)
-from generalresearch.managers.thl.profiling.schema import (
- UpkSchemaManager,
-)
from generalresearch.managers.thl.profiling.uqa import UQAManager
from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager
from generalresearch.models import Source
@@ -183,7 +180,8 @@ class TestSurvey:
]
uqad = {}
for uqa in uqas:
- for k, _ in uqa.calc_answers.items():
+ assert uqa.calc_answers
+ for k in uqa.calc_answers:
if k in qualifying_questions:
uqad[k] = uqa
uqad[uqa.property_code] = uqa
diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py
index 5822207..8cd83ad 100644
--- a/tests/managers/thl/test_user_manager/test_base.py
+++ b/tests/managers/thl/test_user_manager/test_base.py
@@ -21,7 +21,7 @@ from generalresearch.managers.thl.user_manager.user_manager import (
UserManager,
)
from generalresearch.managers.thl.userhealth import AuditLogManager
-from generalresearch.models.thl.product import Product, UserCreateConfig, product
+from generalresearch.models.thl.product import Product, UserCreateConfig
from generalresearch.models.thl.user import User
from generalresearch.pg_helper import PostgresConfig
diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py
index 5e9f4d0..db39997 100644
--- a/tests/models/network/test_nmap.py
+++ b/tests/models/network/test_nmap.py
@@ -8,7 +8,7 @@ from generalresearch.managers.network.tool_run import ToolRunManager
from generalresearch.models.network.definitions import IPProtocol
from generalresearch.models.network.nmap.execute import execute_nmap
from generalresearch.models.network.nmap.result import NmapResult, PortState
-from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName
+from generalresearch.models.network.tool_run import NmapRun, ToolClass, ToolName
fake = faker.Faker()
diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py
index 7ddd407..f97860f 100644
--- a/tests/models/spectrum/test_survey.py
+++ b/tests/models/spectrum/test_survey.py
@@ -405,3 +405,47 @@ class TestSpectrumSurvey:
assert (None, {"c", "d"}) == s.determine_eligibility_soft(
{"a": True, "b": True, "c": None, "d": None}
)
+
+
+def test_spectrum_something(spectrum_api_surveys_json: list[str]):
+ # make sure hashes for 111111 are in db
+ c1 = SpectrumCondition(
+ question_id="1001",
+ value_type=ConditionValueType.LIST,
+ values=["a", "b", "c"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c2 = SpectrumCondition(
+ question_id="1001",
+ value_type=ConditionValueType.LIST,
+ values=["a"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c3 = SpectrumCondition(
+ question_id="1002",
+ value_type=ConditionValueType.RANGE,
+ values=["18-24", "30-32"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c4 = SpectrumCondition(
+ question_id="212",
+ value_type=ConditionValueType.LIST,
+ values=["23", "24"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ c5 = SpectrumCondition(
+ question_id="1031",
+ value_type=ConditionValueType.LIST,
+ values=["113", "114", "121"],
+ negate=False,
+ logical_operator=LogicalOperator.OR,
+ )
+ _conditions = [c1, c2, c3, c4, c5]
+
+ survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0])
+ assert c1.criterion_hash in survey.qualifications
+ assert c3.criterion_hash in survey.qualifications
diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py
index adf276d..880799a 100644
--- a/tests/models/thl/test_product.py
+++ b/tests/models/thl/test_product.py
@@ -1013,7 +1013,7 @@ class TestProductCache:
assert res is None
with pytest.raises(expected_exception=AssertionError):
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
@@ -1035,7 +1035,7 @@ class TestProductCache:
# Now try again with everything in place
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
--
cgit v1.2.3