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') 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') 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') 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