diff options
31 files changed, 447 insertions, 559 deletions
diff --git a/generalresearch/incite/collections/thl_web.py b/generalresearch/incite/collections/thl_web.py index 951406c..d60c77f 100644 --- a/generalresearch/incite/collections/thl_web.py +++ b/generalresearch/incite/collections/thl_web.py @@ -1,6 +1,6 @@ from typing import Literal -from generalresearch.incite.collections import DFCollection, DFCollectionType +from generalresearch.incite.collections.base import DFCollection, DFCollectionType class UserDFCollection(DFCollection): diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 773555c..368b74a 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from generalresearch.incite.base import GRLDatasets diff --git a/generalresearch/incite/mergers/base.py b/generalresearch/incite/mergers/base.py index 810bedc..69c6c46 100644 --- a/generalresearch/incite/mergers/base.py +++ b/generalresearch/incite/mergers/base.py @@ -1,4 +1,3 @@ -import logging import os.path import subprocess from datetime import UTC, datetime @@ -12,6 +11,7 @@ from dask.distributed import Client from pandera.pandas import DataFrameSchema from pydantic import Field, ValidationInfo, field_validator, model_validator +from generalresearch.incite import LOG from generalresearch.incite.base import CollectionBase, CollectionItemBase from generalresearch.incite.schemas import PARTITION_ON from generalresearch.incite.schemas.mergers.foundations.enriched_session import ( @@ -37,8 +37,6 @@ from generalresearch.incite.schemas.mergers.ym_wall_summary import ( ) from generalresearch.models.custom_types import AwareDatetimeISO -LOG = logging.getLogger("incite") - class MergeType(StrEnum): TEST = "test" diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py index 4a300fc..e006768 100644 --- a/generalresearch/incite/mergers/foundations/enriched_session.py +++ b/generalresearch/incite/mergers/foundations/enriched_session.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging from datetime import timedelta from typing import TYPE_CHECKING, Any, Literal @@ -10,6 +9,7 @@ from dask.distributed import Client as DaskClient from dask.distributed import as_completed from more_itertools import chunked, flatten +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import ( SessionDFCollection, WallDFCollection, @@ -30,13 +30,11 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_session import EnrichedSessionSchema, ) from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.models.admin.request import ReportRequest - -LOG = logging.getLogger("incite") + from generalresearch.models.thl.user import User class EnrichedSessionMergeItem(MergeCollectionItem): diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py index f7c679f..be491b7 100644 --- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging from typing import Any, Literal import dask.dataframe as dd @@ -8,6 +7,7 @@ import pandas as pd from distributed import Client from sentry_sdk import capture_exception +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import ( TaskAdjustmentDFCollection, ) @@ -28,8 +28,6 @@ from generalresearch.incite.schemas.mergers.foundations.enriched_task_adjust imp ) from generalresearch.pg_helper import PostgresConfig -LOG = logging.getLogger("incite") - class EnrichedTaskAdjustMergeItem(MergeCollectionItem): """Because a single wall event can have multiple "alerted" times, diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py index 9db5328..396c2be 100644 --- a/generalresearch/incite/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/mergers/foundations/enriched_wall.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging from datetime import timedelta from typing import TYPE_CHECKING, Any, Literal @@ -8,6 +7,7 @@ import dask.dataframe as dd import pandas as pd from distributed import Client +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import ( SessionDFCollection, WallDFCollection, @@ -31,8 +31,6 @@ if TYPE_CHECKING: from generalresearch.models.admin.request import ReportRequest from generalresearch.models.thl.user import User -LOG = logging.getLogger("incite") - class EnrichedWallMergeItem(MergeCollectionItem): diff --git a/generalresearch/incite/mergers/foundations/user_id_product.py b/generalresearch/incite/mergers/foundations/user_id_product.py index 3fb7b36..e682b7f 100644 --- a/generalresearch/incite/mergers/foundations/user_id_product.py +++ b/generalresearch/incite/mergers/foundations/user_id_product.py @@ -1,10 +1,10 @@ from __future__ import annotations -import logging from typing import Any, Literal from distributed import Client +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import UserDFCollection from generalresearch.incite.mergers.base import ( MergeCollection, @@ -12,8 +12,6 @@ from generalresearch.incite.mergers.base import ( MergeType, ) -LOG = logging.getLogger("incite") - class UserIdProductMergeItem(MergeCollectionItem): diff --git a/generalresearch/incite/mergers/pop_ledger.py b/generalresearch/incite/mergers/pop_ledger.py index cf63cad..54f1b7e 100644 --- a/generalresearch/incite/mergers/pop_ledger.py +++ b/generalresearch/incite/mergers/pop_ledger.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging from typing import Any, Literal import dask.dataframe as dd @@ -8,6 +7,7 @@ import pandas as pd from distributed import Client from more_itertools import flatten +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import LedgerDFCollection from generalresearch.incite.mergers.base import ( MergeCollection, @@ -17,8 +17,6 @@ from generalresearch.incite.mergers.base import ( from generalresearch.incite.schemas.mergers.pop_ledger import PopLedgerSchema from generalresearch.models.thl.ledger import Direction, TransactionType -LOG = logging.getLogger("incite") - class PopLedgerMergeItem(MergeCollectionItem): diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py index 3cbf543..8b66eb5 100644 --- a/generalresearch/incite/mergers/ym_survey_wall.py +++ b/generalresearch/incite/mergers/ym_survey_wall.py @@ -1,6 +1,5 @@ from __future__ import annotations -import logging from datetime import timedelta from typing import Any, Literal @@ -9,6 +8,7 @@ import pandas as pd from distributed import Client from sentry_sdk import capture_exception +from generalresearch.incite import LOG from generalresearch.incite.collections.thl_web import WallDFCollection from generalresearch.incite.exceptions import BuildError from generalresearch.incite.mergers.base import ( @@ -24,8 +24,6 @@ from generalresearch.incite.schemas.mergers.ym_survey_wall import ( ) from generalresearch.models.custom_types import AwareDatetimeISO -LOG = logging.getLogger("incite") - class YMSurveyWallMergeCollectionItem(MergeCollectionItem): @@ -39,7 +37,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): LOG.info(f"YMSurveyWallMerge.build({self.start=}, {self.finish=})") ir: pd.Interval = self.interval start, _ = self.start, self.finish - ddf = wall_coll.ddf( + ddf: dd.DataFrame | None = wall_coll.ddf( items=wall_coll.get_items(start), force_rr_latest=False, include_partial=True, @@ -61,6 +59,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): ], filters=[("started", ">=", start)], ) + assert isinstance(ddf, dd.DataFrame) ddf = ddf[ddf["started"] > start] LOG.warning("YMSurveyWallMerge: merge session_collection") @@ -92,6 +91,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): ) df = client.compute(ddf, resources=client_resources, sync=True) + assert isinstance(df, pd.DataFrame) df["elapsed"] = (df["finished"] - df["started"]).dt.total_seconds() df["elapsed"] = df["elapsed"].round().astype("Int64") df = df.drop(columns={"finished", "payout"}, errors="ignore") diff --git a/generalresearch/incite/schemas/mergers/foundations/user_id_product.py b/generalresearch/incite/schemas/mergers/foundations/user_id_product.py index a07ca42..9e73aed 100644 --- a/generalresearch/incite/schemas/mergers/foundations/user_id_product.py +++ b/generalresearch/incite/schemas/mergers/foundations/user_id_product.py @@ -4,7 +4,7 @@ from pandera.pandas import Category, Check, Column, DataFrameSchema, Index from generalresearch.incite.schemas import ARCHIVE_AFTER -BIGINT = 9223372036854775807 +BIGINT = 9_223_372_036_854_775_807 UserIdIndex = Index( name="id", diff --git a/generalresearch/incite/schemas/mergers/pop_ledger.py b/generalresearch/incite/schemas/mergers/pop_ledger.py index fdb1b05..8452eb9 100644 --- a/generalresearch/incite/schemas/mergers/pop_ledger.py +++ b/generalresearch/incite/schemas/mergers/pop_ledger.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import timedelta import pandas as pd @@ -22,6 +25,11 @@ from generalresearch.models.thl.ledger import Direction, TransactionType # If an amount is "very" large, something is def wrong. Defining "very" somewhat arbitrarily here. SUSPICIOUSLY_LARGE_NUMBER = (2**32 / 2) - 1 # 2147483647 +_tz_min_freq: Callable[[pd.Series], pd.Series] = lambda i: (i.dt.second == 0) & ( + i.dt.microsecond == 0 +) + + NonNegativeAmount = Column( dtype="Int32", nullable=True, @@ -49,7 +57,7 @@ PopLedgerSchema = DataFrameSchema( | { "time_idx": Column( dtype=pd.DatetimeTZDtype(tz="UTC"), - checks=Check(lambda x: (x.dt.second == 0) & (x.dt.microsecond == 0)), + checks=Check(_tz_min_freq), nullable=False, ), "account_id": TxSchema.columns["account_id"], diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index 9595cab..98a9a32 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -156,17 +156,15 @@ class IPGeonameManager(PostgresManager): if len(filter_ids) == 0: return [] - with self.pg_config.make_connection() as sql_connection: - sql_connection: pymysql.Connection - with sql_connection.cursor() as c: - res = [] - for chunk in chunked(filter_ids, 500): - res.extend( - self.fetch_geoname_ids_( - c=c, - filter_ids=chunk, - ) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + res = [] + for chunk in chunked(filter_ids, 500): + res.extend( + self.fetch_geoname_ids_( + c=c, + filter_ids=chunk, ) + ) return res def fetch_geoname_ids_( @@ -635,16 +633,14 @@ class GeoIpInfoManager(PostgresManagerWithRedis): if len(ips) == 0: return {} - with self.pg_config.make_connection() as sql_connection: - sql_connection: pymysql.Connection - with sql_connection.cursor() as c: - res = {} - for chunk in chunked(ips, 500): - inner = self.get_mysql_multi_chunk( - c=c, - ips=chunk, - ) - res.update(inner) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + res = {} + for chunk in chunked(ips, 500): + inner = self.get_mysql_multi_chunk( + c=c, + ips=chunk, + ) + res.update(inner) return res def get_mysql_multi_chunk( diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py index b3fa8f6..449dc2a 100644 --- a/generalresearch/managers/thl/user_manager/__init__.py +++ b/generalresearch/managers/thl/user_manager/__init__.py @@ -2,9 +2,6 @@ from __future__ import annotations import csv import logging -import os -import threading -import time from pathlib import Path from threading import RLock from typing import Any @@ -15,17 +12,7 @@ from generalresearch.models.thl.product import Product logger = logging.getLogger() - -class UserDoesntExistError(Exception): - pass - - -class UserCreateNotAllowedError(Exception): - pass - - -def download_bp_trust(): - raise DeprecationWarning("No more S3") +convert_int = lambda x: int(float(x)) @cached(TTLCache(maxsize=1, ttl=5 * 60), lock=RLock()) @@ -39,22 +26,11 @@ def get_bp_trust_df(): # 'product_name', 'bp_trust', 'team_trust', 'entrance_limit_expire_sec', # 'entrance_limit_value'] - if not os.path.exists(fp): - Path(fp).touch() - threading.Thread(target=download_bp_trust).start() - # raise exception so its not cached - raise FileNotFoundError() - if time.time() - os.path.getmtime(fp) > 3600: - Path(fp).touch() - threading.Thread(target=download_bp_trust).start() bptrust = parse_bp_trust_df(fp) return bptrust -convert_int = lambda x: int(float(x)) - - def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]: dtype = { "bp_trust": float, diff --git a/generalresearch/managers/thl/user_manager/exceptions.py b/generalresearch/managers/thl/user_manager/exceptions.py new file mode 100644 index 0000000..6d8c87c --- /dev/null +++ b/generalresearch/managers/thl/user_manager/exceptions.py @@ -0,0 +1,6 @@ +class UserDoesntExistError(Exception): + pass + + +class UserCreateNotAllowedError(Exception): + pass diff --git a/generalresearch/managers/thl/user_manager/rate_limit.py b/generalresearch/managers/thl/user_manager/rate_limit.py index f938664..5d239c6 100644 --- a/generalresearch/managers/thl/user_manager/rate_limit.py +++ b/generalresearch/managers/thl/user_manager/rate_limit.py @@ -5,9 +5,11 @@ from limits.limits import TIME_TYPES, safe_string from pydantic import RedisDsn from generalresearch.managers.thl.user_manager import ( - UserCreateNotAllowedError, get_bp_user_create_limit_hourly, ) +from generalresearch.managers.thl.user_manager.exceptions import ( + UserCreateNotAllowedError, +) from generalresearch.models.thl.product import Product logger = logging.getLogger() diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 7486869..26b1fd6 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -4,12 +4,13 @@ import logging from collections.abc import Collection from datetime import datetime from functools import lru_cache +from typing import TYPE_CHECKING from pydantic import RedisDsn from generalresearch.managers.base import Permission from generalresearch.managers.thl.product import ProductManager -from generalresearch.managers.thl.user_manager import UserDoesntExistError +from generalresearch.managers.thl.user_manager.exceptions import UserDoesntExistError from generalresearch.managers.thl.user_manager.mysql_user_manager import ( MysqlUserManager, ) @@ -19,12 +20,16 @@ from generalresearch.managers.thl.user_manager.rate_limit import ( from generalresearch.managers.thl.user_manager.redis_user_manager import ( RedisUserManager, ) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.utils.copying_cache import deepcopy_return +if TYPE_CHECKING: + from generalresearch.managers.thl.userhealth import AuditLogManager + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.user import User + from generalresearch.models.thl.userhealth import AuditLog + logging.basicConfig() logger = logging.getLogger() auditlog = logging.getLogger("auditlog") @@ -85,17 +90,16 @@ class UserManager: def audit_log( self, + alm: AuditLogManager, user: User, level: int, event_type: str, event_msg: str | None = None, event_value: float | None = None, - ) -> None: - from generalresearch.managers.thl.userhealth import AuditLogManager + ) -> AuditLog: from generalresearch.models.thl.userhealth import AuditLogLevel - alm = AuditLogManager(pg_config=self.mysql_user_manager.pg_config) - alm.create( + return alm.create( user_id=user.user_id, level=AuditLogLevel(level), event_type=event_type, @@ -103,7 +107,6 @@ class UserManager: event_value=event_value, ) - def cache_clear(self): # Generally this is used in testing. This clears the .get_user's lru_cache. # There is no way of clearing only a specific key from the cache. @@ -149,6 +152,7 @@ class UserManager: # We can use the read-replica here b/c when we create a user we'll # put it in the in-memory cache mysql_user_manager = self.mysql_user_manager_rr or self.mysql_user_manager + assert mysql_user_manager user = mysql_user_manager.get_user_from_mysql( product_id=product_id, product_user_id=product_user_id, @@ -245,6 +249,9 @@ class UserManager: self.user_manager_limiter is not None ), "Need user_manager_limiter to get_or_create_user" # Attempt to create common_struct solely for validation purposes + + from generalresearch.models.thl.user import User + if not User.is_valid_ubp( product_id=product_id, product_user_id=product_user_id ): @@ -282,6 +289,7 @@ class UserManager: # if product.id not in {}: # self.user_manager_limiter.raise_allow_user_create(product=product) + assert self.mysql_user_manager user = self.mysql_user_manager.create_user( product_user_id=product_user_id, product_id=product.id, @@ -294,6 +302,7 @@ class UserManager: def product_id_exists(self, product_id: str) -> bool: mysql_user_manager = self.mysql_user_manager_rr or self.mysql_user_manager + assert mysql_user_manager return mysql_user_manager.product_id_exists(product_id) def block_user(self, user: User) -> bool: @@ -327,6 +336,7 @@ class UserManager: Currently, this sets a key in the userprofile_userstat table. TODO: this should be a property of the user? """ + assert self.mysql_user_manager return self.mysql_user_manager.is_whitelisted(user=user) def fetch_by_bpuids( @@ -337,6 +347,7 @@ class UserManager: ) -> list[User]: assert product_id, "must pass product_id" assert len(product_user_ids) > 0, "must pass 1 or more product_user_ids" + assert self.mysql_user_manager_rr return self.mysql_user_manager_rr.fetch_by_bpuids( product_id=product_id, product_user_ids=product_user_ids ) @@ -350,6 +361,7 @@ class UserManager: assert (user_ids or user_uuids) and not ( user_ids and user_uuids ), "Must pass ONE of user_ids, user_uuids" + assert self.mysql_user_manager_rr return self.mysql_user_manager_rr.fetch( user_ids=user_ids, user_uuids=user_uuids ) diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index f28fe0c..2bbdfab 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -4,7 +4,7 @@ import ipaddress from collections.abc import Collection from datetime import UTC, datetime, timedelta from itertools import zip_longest -from typing import Any +from typing import TYPE_CHECKING, Any import faker from pydantic import NonNegativeInt, PositiveInt @@ -18,7 +18,6 @@ from generalresearch.managers.base import ( from generalresearch.managers.thl.ipinfo import GeoIpInfoManager from generalresearch.models.custom_types import IPvAnyAddressStr from generalresearch.models.thl.product import Product -from generalresearch.models.thl.user import User from generalresearch.models.thl.user_iphistory import ( IPRecord, UserIPHistory, @@ -28,6 +27,9 @@ from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.models.thl.user import User + fake = faker.Faker() diff --git a/generalresearch/models/network/mtr/command.py b/generalresearch/models/network/mtr/command.py index fd5a8d0..220bfc7 100644 --- a/generalresearch/models/network/mtr/command.py +++ b/generalresearch/models/network/mtr/command.py @@ -1,11 +1,14 @@ from __future__ import annotations import subprocess +from typing import TYPE_CHECKING from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.mtr.parser import parse_mtr_output from generalresearch.models.network.mtr.result import MTRResult -from generalresearch.models.network.tool_run_command import MTRRunCommand + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run_command import MTRRunCommand SUPPORTED_PROTOCOLS = { IPProtocol.TCP, diff --git a/generalresearch/models/network/nmap/command.py b/generalresearch/models/network/nmap/command.py index 6a524a8..42b7178 100644 --- a/generalresearch/models/network/nmap/command.py +++ b/generalresearch/models/network/nmap/command.py @@ -1,10 +1,13 @@ from __future__ import annotations import subprocess +from typing import TYPE_CHECKING from generalresearch.models.network.nmap.parser import parse_nmap_xml from generalresearch.models.network.nmap.result import NmapResult -from generalresearch.models.network.tool_run_command import NmapRunCommand + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run_command import NmapRunCommand def build_nmap_command( diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py index bccead0..c63d6d2 100644 --- a/generalresearch/models/network/rdns/command.py +++ b/generalresearch/models/network/rdns/command.py @@ -1,8 +1,11 @@ import subprocess +from typing import TYPE_CHECKING from generalresearch.models.network.rdns.parser import parse_rdns_output from generalresearch.models.network.rdns.result import RDNSResult -from generalresearch.models.network.tool_run_command import RDNSRunCommand + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run_command import RDNSRunCommand def run_rdns(config: RDNSRunCommand) -> RDNSResult: diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py index abc129b..cb04b29 100644 --- a/generalresearch/models/thl/__init__.py +++ b/generalresearch/models/thl/__init__.py @@ -1,26 +1,26 @@ from decimal import Decimal -from generalresearch.models.thl.finance import ( - POPFinancial, - ProductBalances, -) -from generalresearch.models.thl.payout import ( - BrokerageProductPayoutEvent, - PayoutEvent, -) -from generalresearch.models.thl.product import Product +# from generalresearch.models.thl.finance import ( +# POPFinancial, +# ProductBalances, +# ) +# from generalresearch.models.thl.payout import ( +# BrokerageProductPayoutEvent, +# PayoutEvent, +# ) +# from generalresearch.models.thl.product import Product -_ = ( - Product, - PayoutEvent, - BrokerageProductPayoutEvent, - ProductBalances, - POPFinancial, -) +# _ = ( +# Product, +# PayoutEvent, +# BrokerageProductPayoutEvent, +# ProductBalances, +# POPFinancial, +# ) -Product.model_rebuild() -PayoutEvent.model_rebuild() -BrokerageProductPayoutEvent.model_rebuild() +# Product.model_rebuild() +# PayoutEvent.model_rebuild() +# BrokerageProductPayoutEvent.model_rebuild() def decimal_to_int_cents(usd: Decimal | None) -> int | None: diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 5814bef..6fc60f6 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -77,6 +77,10 @@ class ContestBase(BaseModel, ABC): self.model_config["validate_assignment"] = True self.__class__.model_validate(self) + @classmethod + def example_json_schema_extra(cls, schema: dict[str, Any]) -> None: + schema["examples"] = [cls.example().model_dump(mode="json")] + class Contest(ContestBase): id: int | None = Field( diff --git a/generalresearch/models/thl/contest/examples.py b/generalresearch/models/thl/contest/examples.py deleted file mode 100644 index 018f090..0000000 --- a/generalresearch/models/thl/contest/examples.py +++ /dev/null @@ -1,409 +0,0 @@ -from __future__ import annotations - -from typing import Any - -from pydantic import HttpUrl - -from generalresearch.config import EXAMPLE_PRODUCT_ID -from generalresearch.currency import USDCent - - -def _example_raffle_create(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestEndCondition, - ContestEntryRule, - ContestPrize, - ) - from generalresearch.models.thl.contest.contest_entry import ( - ContestEntryType, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.raffle import ( - RaffleContestCreate, - ) - - schema["example"] = RaffleContestCreate( - name="Win an iPhone", - description="iPhone winner will be drawn in proportion to entry " - "amount. Contest ends once $800 has been entered.", - contest_type=ContestType.RAFFLE, - end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PHYSICAL, - name="iPhone 16", - estimated_cash_value=USDCent(800_00), - ) - ], - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=None, - entry_rule=ContestEntryRule( - max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 - ), - country_isos={"us", "ca"}, - entry_type=ContestEntryType.CASH, - ).model_dump(mode="json") - - -def _example_raffle(schema: dict) -> None: - from generalresearch.models.thl.contest import ( - ContestEndCondition, - ContestEntryRule, - ContestPrize, - ) - from generalresearch.models.thl.contest.contest_entry import ( - ContestEntryType, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestStatus, - ContestType, - ) - from generalresearch.models.thl.contest.raffle import RaffleContest - - schema["example"] = RaffleContest( - name="Win an iPhone", - description="iPhone winner will be drawn in proportion to entry " - "amount. Contest ends once $800 has been entered.", - contest_type=ContestType.RAFFLE, - end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PHYSICAL, - name="iPhone 16", - estimated_cash_value=USDCent(800_00), - ) - ], - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=None, - entry_rule=ContestEntryRule( - max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 - ), - country_isos={"us", "ca"}, - entry_type=ContestEntryType.CASH, - status=ContestStatus.ACTIVE, - uuid="ce3968b8e18a4b96af62007f262ed7f7", - created_at="2025-06-12T21:12:58.061205Z", - updated_at="2025-06-12T21:12:58.061205Z", - current_amount=4723, - current_participants=12, - product_id=EXAMPLE_PRODUCT_ID, - ).model_dump(mode="json") - - - -def _example_raffle_user_view(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestEndCondition, - ContestEntryRule, - ContestPrize, - ) - from generalresearch.models.thl.contest.contest_entry import ( - ContestEntryType, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestStatus, - ContestType, - ) - from generalresearch.models.thl.contest.raffle import RaffleUserView - - schema["example"] = RaffleUserView( - name="Win an iPhone", - description="iPhone winner will be drawn in proportion to entry " - "amount. Contest ends once $800 has been entered.", - contest_type=ContestType.RAFFLE, - end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PHYSICAL, - name="iPhone 16", - estimated_cash_value=USDCent(800_00), - ) - ], - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=None, - entry_rule=ContestEntryRule( - max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 - ), - country_isos={"us", "ca"}, - entry_type=ContestEntryType.CASH, - status=ContestStatus.ACTIVE, - uuid="ce3968b8e18a4b96af62007f262ed7f7", - created_at="2025-06-12T21:12:58.061205Z", - updated_at="2025-06-12T21:12:58.061205Z", - current_amount=4723, - current_participants=12, - product_id=EXAMPLE_PRODUCT_ID, - user_amount=420, - user_amount_today=0, - product_user_id="test-user", - ).model_dump(mode="json") - - - -def _example_milestone_create(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, - MilestoneContestCreate, - MilestoneContestEndCondition, - ) - - schema["example"] = MilestoneContestCreate( - name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", - description="Only valid for the first 50 users", - contest_type=ContestType.MILESTONE, - end_condition=MilestoneContestEndCondition(max_winners=50), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PROMOTION, - name="50% bonus on completes for 7 days", - estimated_cash_value=USDCent(0), - ), - ContestPrize( - kind=ContestPrizeKind.CASH, - name="$5.00 Bonus", - cash_amount=USDCent(5_00), - estimated_cash_value=USDCent(5_00), - ), - ], - entry_trigger=ContestEntryTrigger.TASK_COMPLETE, - target_amount=10, - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=HttpUrl("https://www.example.com"), - ).model_dump(mode="json") - - - -def _example_milestone(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, - MilestoneContest, - MilestoneContestEndCondition, - ) - - schema["example"] = MilestoneContest( - name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", - description="Only valid for the first 50 users", - contest_type=ContestType.MILESTONE, - end_condition=MilestoneContestEndCondition(max_winners=50), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PROMOTION, - name="50% bonus on completes for 7 days", - estimated_cash_value=USDCent(0), - ), - ContestPrize( - kind=ContestPrizeKind.CASH, - name="$5.00 Bonus", - cash_amount=USDCent(5_00), - estimated_cash_value=USDCent(5_00), - ), - ], - entry_trigger=ContestEntryTrigger.TASK_COMPLETE, - target_amount=10, - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=HttpUrl("https://www.example.com"), - product_id=EXAMPLE_PRODUCT_ID, - uuid="747fe3b709ae460e816821dcb81aebb9", - created_at="2025-06-12T21:12:58.061205Z", - updated_at="2025-06-12T21:12:58.061205Z", - win_count=12, - ).model_dump(mode="json") - - - -def _example_milestone_user_view(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ContestPrize - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, - MilestoneContestEndCondition, - MilestoneUserView, - ) - - schema["example"] = MilestoneUserView( - name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", - description="Only valid for the first 50 users", - contest_type=ContestType.MILESTONE, - end_condition=MilestoneContestEndCondition(max_winners=50), - prizes=[ - ContestPrize( - kind=ContestPrizeKind.PROMOTION, - name="50% bonus on completes for 7 days", - estimated_cash_value=USDCent(0), - ), - ContestPrize( - kind=ContestPrizeKind.CASH, - name="$5.00 Bonus", - cash_amount=USDCent(5_00), - estimated_cash_value=USDCent(5_00), - ), - ], - entry_trigger=ContestEntryTrigger.TASK_COMPLETE, - target_amount=10, - starts_at="2025-06-12T21:12:58.061170Z", - terms_and_conditions=HttpUrl("https://www.example.com"), - product_id=EXAMPLE_PRODUCT_ID, - uuid="747fe3b709ae460e816821dcb81aebb9", - created_at="2025-06-12T21:12:58.061205Z", - updated_at="2025-06-12T21:12:58.061205Z", - win_count=12, - user_amount=8, - product_user_id="test-user", - ).model_dump(mode="json") - - - -def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContestCreate, - ) - - schema["example"] = LeaderboardContestCreate( - name="Prizes for top survey takers this week", - description="$15 1st place, $10 2nd, $5 3rd place US weekly", - contest_type=ContestType.LEADERBOARD, - prizes=[ - ContestPrize( - name="$15 Cash", - estimated_cash_value=USDCent(15_00), - cash_amount=USDCent(15_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=1, - ), - ContestPrize( - name="$10 Cash", - estimated_cash_value=USDCent(10_00), - cash_amount=USDCent(10_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=2, - ), - ContestPrize( - name="$5 Cash", - estimated_cash_value=USDCent(5_00), - cash_amount=USDCent(5_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=3, - ), - ], - leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count", - ).model_dump(mode="json") - - - -def _example_leaderboard_contest(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContest, - ) - - schema["example"] = LeaderboardContest( - name="Prizes for top survey takers this week", - description="$15 1st place, $10 2nd, $5 3rd place US weekly", - contest_type=ContestType.LEADERBOARD, - prizes=[ - ContestPrize( - name="$15 Cash", - estimated_cash_value=USDCent(15_00), - cash_amount=USDCent(15_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=1, - ), - ContestPrize( - name="$10 Cash", - estimated_cash_value=USDCent(10_00), - cash_amount=USDCent(10_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=2, - ), - ContestPrize( - name="$5 Cash", - estimated_cash_value=USDCent(5_00), - cash_amount=USDCent(5_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=3, - ), - ], - leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count", - product_id=EXAMPLE_PRODUCT_ID, - ).model_dump(mode="json") - - - -def _example_leaderboard_contest_user_view(schema: dict[str, Any]) -> None: - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContestUserView, - ) - - schema["example"] = LeaderboardContestUserView( - name="Prizes for top survey takers this week", - description="$15 1st place, $10 2nd, $5 3rd place US weekly", - contest_type=ContestType.LEADERBOARD, - prizes=[ - ContestPrize( - name="$15 Cash", - estimated_cash_value=USDCent(15_00), - cash_amount=USDCent(15_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=1, - ), - ContestPrize( - name="$10 Cash", - estimated_cash_value=USDCent(10_00), - cash_amount=USDCent(10_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=2, - ), - ContestPrize( - name="$5 Cash", - estimated_cash_value=USDCent(5_00), - cash_amount=USDCent(5_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=3, - ), - ], - leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count", - product_id=EXAMPLE_PRODUCT_ID, - product_user_id="test-user", - ).model_dump(mode="json") diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index e923383..080f356 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -12,6 +12,7 @@ from pydantic import ( ) from redis import Redis +from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.managers.leaderboard import country_timezone from generalresearch.managers.leaderboard.manager import LeaderboardManager @@ -20,6 +21,7 @@ from generalresearch.managers.thl.user_manager.user_manager import ( ) from generalresearch.models.thl.contest import ( ContestEndCondition, + ContestPrize, ContestWinner, ) from generalresearch.models.thl.contest.contest import ( @@ -34,11 +36,6 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, LeaderboardTieBreakStrategy, ) -from generalresearch.models.thl.contest.examples import ( - _example_leaderboard_contest, - _example_leaderboard_contest_create, - _example_leaderboard_contest_user_view, -) from generalresearch.models.thl.leaderboard import ( Leaderboard, LeaderboardCode, @@ -50,7 +47,6 @@ class LeaderboardContestCreate(ContestBase): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_leaderboard_contest_create, ) contest_type: Literal[ContestType.LEADERBOARD] = Field( @@ -124,12 +120,45 @@ class LeaderboardContestCreate(ContestBase): parts | {"row_count": 0, "bpid": parts["product_id"]} ) + @classmethod + def example(cls) -> LeaderboardContestCreate: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + + return cls( + name="Prizes for top survey takers this week", + description="$15 1st place, $10 2nd, $5 3rd place US weekly", + contest_type=ContestType.LEADERBOARD, + prizes=[ + ContestPrize( + name="$15 Cash", + estimated_cash_value=USDCent(15_00), + cash_amount=USDCent(15_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=1, + ), + ContestPrize( + name="$10 Cash", + estimated_cash_value=USDCent(10_00), + cash_amount=USDCent(10_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=2, + ), + ContestPrize( + name="$5 Cash", + estimated_cash_value=USDCent(5_00), + cash_amount=USDCent(5_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=3, + ), + ], + leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count", + ) + class LeaderboardContest(LeaderboardContestCreate, Contest): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_leaderboard_contest, arbitrary_types_allowed=True, ) @@ -246,12 +275,46 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): ) return d + @classmethod + def example(cls) -> LeaderboardContest: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + + return cls( + name="Prizes for top survey takers this week", + description="$15 1st place, $10 2nd, $5 3rd place US weekly", + contest_type=ContestType.LEADERBOARD, + prizes=[ + ContestPrize( + name="$15 Cash", + estimated_cash_value=USDCent(15_00), + cash_amount=USDCent(15_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=1, + ), + ContestPrize( + name="$10 Cash", + estimated_cash_value=USDCent(10_00), + cash_amount=USDCent(10_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=2, + ), + ContestPrize( + name="$5 Cash", + estimated_cash_value=USDCent(5_00), + cash_amount=USDCent(5_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=3, + ), + ], + leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count", + product_id=product_id, + ) + class LeaderboardContestUserView(LeaderboardContest, ContestUserView): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_leaderboard_contest_user_view, ) @computed_field(description="The current rank of this user in this contest") @@ -291,3 +354,39 @@ class LeaderboardContestUserView(LeaderboardContest, ContestUserView): return False, "contest is over" return True, "" + + @classmethod + def example(cls) -> LeaderboardContestUserView: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + + return cls( + name="Prizes for top survey takers this week", + description="$15 1st place, $10 2nd, $5 3rd place US weekly", + contest_type=ContestType.LEADERBOARD, + prizes=[ + ContestPrize( + name="$15 Cash", + estimated_cash_value=USDCent(15_00), + cash_amount=USDCent(15_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=1, + ), + ContestPrize( + name="$10 Cash", + estimated_cash_value=USDCent(10_00), + cash_amount=USDCent(10_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=2, + ), + ContestPrize( + name="$5 Cash", + estimated_cash_value=USDCent(5_00), + cash_amount=USDCent(5_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=3, + ), + ], + leaderboard_key=f"leaderboard:{product_id}:us:weekly:2025-05-26:complete_count", + product_id=product_id, + product_user_id="test-user", + ) diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index 5fc27fa..db5ba2f 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -8,10 +8,15 @@ from pydantic import ( BaseModel, ConfigDict, Field, + HttpUrl, PositiveInt, ) +from generalresearch.currency import USDCent from generalresearch.models.custom_types import AwareDatetimeISO +from generalresearch.models.thl.contest import ( + ContestPrize, +) from generalresearch.models.thl.contest.contest import ( Contest, ContestBase, @@ -22,14 +27,10 @@ from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestEntryTrigger, ContestEntryType, + ContestPrizeKind, ContestStatus, ContestType, ) -from generalresearch.models.thl.contest.examples import ( - _example_milestone, - _example_milestone_create, - _example_milestone_user_view, -) logging.basicConfig() LOG = logging.getLogger() @@ -103,19 +104,46 @@ class MilestoneContestCreate(ContestBase, MilestoneContestConfig): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_milestone_create, + # json_schema_extra=json_example_milestone_create, ) contest_type: Literal[ContestType.MILESTONE] = Field(default=ContestType.MILESTONE) end_condition: MilestoneContestEndCondition = Field() + @classmethod + def example(cls) -> MilestoneContestCreate: + + return cls( + name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", + description="Only valid for the first 50 users", + contest_type=ContestType.MILESTONE, + end_condition=MilestoneContestEndCondition(max_winners=50), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PROMOTION, + name="50% bonus on completes for 7 days", + estimated_cash_value=USDCent(0), + ), + ContestPrize( + kind=ContestPrizeKind.CASH, + name="$5.00 Bonus", + cash_amount=USDCent(5_00), + estimated_cash_value=USDCent(5_00), + ), + ], + entry_trigger=ContestEntryTrigger.TASK_COMPLETE, + target_amount=10, + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=HttpUrl("https://www.example.com"), + ) + class MilestoneContest(MilestoneContestCreate, Contest): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_milestone, + # json_schema_extra=json_example_milestone, ) entry_type: Literal[ContestEntryType.COUNT] = Field(default=ContestEntryType.COUNT) @@ -173,12 +201,43 @@ class MilestoneContest(MilestoneContestCreate, Contest): ) return super().model_validate_mysql(data) + @classmethod + def example(cls) -> MilestoneContest: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + return cls( + name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", + description="Only valid for the first 50 users", + contest_type=ContestType.MILESTONE, + end_condition=MilestoneContestEndCondition(max_winners=50), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PROMOTION, + name="50% bonus on completes for 7 days", + estimated_cash_value=USDCent(0), + ), + ContestPrize( + kind=ContestPrizeKind.CASH, + name="$5.00 Bonus", + cash_amount=USDCent(5_00), + estimated_cash_value=USDCent(5_00), + ), + ], + entry_trigger=ContestEntryTrigger.TASK_COMPLETE, + target_amount=10, + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=HttpUrl("https://www.example.com"), + product_id=product_id, + uuid="747fe3b709ae460e816821dcb81aebb9", + created_at="2025-06-12T21:12:58.061205Z", + updated_at="2025-06-12T21:12:58.061205Z", + win_count=12, + ) + class MilestoneUserView(MilestoneContest, ContestUserView): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_milestone_user_view, ) valid_until: AwareDatetimeISO | None = Field( @@ -218,3 +277,37 @@ class MilestoneUserView(MilestoneContest, ContestUserView): # TODO: others in self.entry_rule ... min_completes, id_verified, etc. return True, "" + + @classmethod + def example(cls) -> MilestoneUserView: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + return cls( + name="Win a 50% bonus for 7 days and a $5 bonus after your first 10 completes!", + description="Only valid for the first 50 users", + contest_type=ContestType.MILESTONE, + end_condition=MilestoneContestEndCondition(max_winners=50), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PROMOTION, + name="50% bonus on completes for 7 days", + estimated_cash_value=USDCent(0), + ), + ContestPrize( + kind=ContestPrizeKind.CASH, + name="$5.00 Bonus", + cash_amount=USDCent(5_00), + estimated_cash_value=USDCent(5_00), + ), + ], + entry_trigger=ContestEntryTrigger.TASK_COMPLETE, + target_amount=10, + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=HttpUrl("https://www.example.com"), + product_id=product_id, + uuid="747fe3b709ae460e816821dcb81aebb9", + created_at="2025-06-12T21:12:58.061205Z", + updated_at="2025-06-12T21:12:58.061205Z", + win_count=12, + user_amount=8, + product_user_id="test-user", + ) diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 072f011..14b3fb6 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -17,7 +17,9 @@ from scipy.stats import hypergeom from generalresearch.currency import USDCent from generalresearch.models.thl.contest import ( + ContestEndCondition, ContestEntryRule, + ContestPrize, ContestWinner, ) from generalresearch.models.thl.contest.contest import ( @@ -25,18 +27,16 @@ from generalresearch.models.thl.contest.contest import ( ContestBase, ContestUserView, ) -from generalresearch.models.thl.contest.contest_entry import ContestEntry +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, + ContestEntryType, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, - ContestEntryType, + ContestPrizeKind, ContestStatus, ContestType, ) -from generalresearch.models.thl.contest.examples import ( - _example_raffle, - _example_raffle_create, - _example_raffle_user_view, -) logging.basicConfig() LOG = logging.getLogger() @@ -47,7 +47,7 @@ class RaffleContestCreate(ContestBase): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_raffle_create, + # json_schema_extra=json_example_raffle_create, ) contest_type: Literal[ContestType.RAFFLE] = Field(default=ContestType.RAFFLE) @@ -63,12 +63,36 @@ class RaffleContestCreate(ContestBase): raise ValueError("At least one end condition must be specified") return self + @classmethod + def example(cls) -> RaffleContestCreate: + return cls( + name="Win an iPhone", + description="iPhone winner will be drawn in proportion to entry " + "amount. Contest ends once $800 has been entered.", + contest_type=ContestType.RAFFLE, + end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PHYSICAL, + name="iPhone 16", + estimated_cash_value=USDCent(800_00), + ) + ], + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=None, + entry_rule=ContestEntryRule( + max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 + ), + country_isos={"us", "ca"}, + entry_type=ContestEntryType.CASH, + ) + class RaffleContest(RaffleContestCreate, Contest): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_raffle, + # json_schema_extra=json_example_raffle, ) entries: list[ContestEntry] = Field(default_factory=list, exclude=True) @@ -92,13 +116,16 @@ class RaffleContest(RaffleContestCreate, Contest): return self @field_validator("current_amount", mode="before") - def coerce_current_amount(cls, v, info): + def coerce_current_amount(cls, v: int | USDCent, info): if v is None: return None + if info.data.get("entry_type") == ContestEntryType.CASH: return USDCent(v) + elif info.data.get("entry_type") == ContestEntryType.COUNT: return int(v) + return v @model_validator(mode="after") @@ -215,12 +242,44 @@ class RaffleContest(RaffleContestCreate, Contest): data["entry_rule"] = ContestEntryRule.model_validate(data["entry_rule"]) return super().model_validate_mysql(data) + @classmethod + def example(cls) -> RaffleContest: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + return cls( + name="Win an iPhone", + description="iPhone winner will be drawn in proportion to entry " + "amount. Contest ends once $800 has been entered.", + contest_type=ContestType.RAFFLE, + end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PHYSICAL, + name="iPhone 16", + estimated_cash_value=USDCent(800_00), + ) + ], + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=None, + entry_rule=ContestEntryRule( + max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 + ), + country_isos={"us", "ca"}, + entry_type=ContestEntryType.CASH, + status=ContestStatus.ACTIVE, + uuid="ce3968b8e18a4b96af62007f262ed7f7", + created_at="2025-06-12T21:12:58.061205Z", + updated_at="2025-06-12T21:12:58.061205Z", + current_amount=4723, + current_participants=12, + product_id=product_id, + ) + class RaffleUserView(RaffleContest, ContestUserView): model_config = ConfigDict( validate_assignment=True, extra="forbid", - json_schema_extra=_example_raffle_user_view, + # json_schema_extra=json_example_raffle_user_view, ) user_amount: int | USDCent = Field( @@ -319,3 +378,38 @@ class RaffleUserView(RaffleContest, ContestUserView): # todo: others in self.entry_rule ... min_completes, id_verified, etc. return True, "" + + @classmethod + def example(cls) -> RaffleUserView: + product_id = "1108d053e4fa47c5b0dbdcd03a7981e7" + return cls( + name="Win an iPhone", + description="iPhone winner will be drawn in proportion to entry " + "amount. Contest ends once $800 has been entered.", + contest_type=ContestType.RAFFLE, + end_condition=ContestEndCondition(target_entry_amount=USDCent(800_00)), + prizes=[ + ContestPrize( + kind=ContestPrizeKind.PHYSICAL, + name="iPhone 16", + estimated_cash_value=USDCent(800_00), + ) + ], + starts_at="2025-06-12T21:12:58.061170Z", + terms_and_conditions=None, + entry_rule=ContestEntryRule( + max_entry_amount_per_user=10000, max_daily_entries_per_user=1000 + ), + country_isos={"us", "ca"}, + entry_type=ContestEntryType.CASH, + status=ContestStatus.ACTIVE, + uuid="ce3968b8e18a4b96af62007f262ed7f7", + created_at="2025-06-12T21:12:58.061205Z", + updated_at="2025-06-12T21:12:58.061205Z", + current_amount=4723, + current_participants=12, + product_id=product_id, + user_amount=420, + user_amount_today=0, + product_user_id="test-user", + ) diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 1e2be5b..8fbae4c 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -2,7 +2,7 @@ from __future__ import annotations import ipaddress from datetime import UTC, datetime -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from faker import Faker from grip_client.enums import AccessType @@ -20,7 +20,9 @@ from generalresearch.models.custom_types import ( CountryISOLike, IPvAnyAddressStr, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.ipinfo import IPGeonameManager fake = Faker() @@ -235,14 +237,13 @@ class IPInformation(BaseModel): # --- prefetch_* --- def prefetch_geoname( self, - pg_config: PostgresConfig, + ip_gm: IPGeonameManager, ) -> None: if self.geoname_id is None: raise ValueError("Must provide geoname_id") - from generalresearch.managers.thl.ipinfo import IPGeonameManager - - ip_gm = IPGeonameManager(pg_config=pg_config) + # from generalresearch.managers.thl.ipinfo import IPGeonameManager + # ip_gm = IPGeonameManager(pg_config=pg_config) self._geoname = ip_gm.get_by_id(geoname_id=self.geoname_id) diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index d547cf4..7121ea5 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -42,12 +42,12 @@ from generalresearch.models.thl.definitions import ( WallAdjustedStatus, WallStatusCode2, ) -from generalresearch.models.thl.user import User if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) + from generalresearch.models.thl.user import User logger = logging.getLogger("Wall") diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py index 2773629..812c18f 100644 --- a/generalresearch/models/thl/user_iphistory.py +++ b/generalresearch/models/thl/user_iphistory.py @@ -2,7 +2,7 @@ from __future__ import annotations import ipaddress from datetime import UTC, datetime, timedelta -from typing import Self +from typing import TYPE_CHECKING, Self from faker import Faker from grip_client.enums import AccessType @@ -19,14 +19,14 @@ from generalresearch.models.custom_types import ( CountryISOLike, IPvAnyAddressStr, ) -from generalresearch.models.thl.ipinfo import ( - GeoIPInformation, - normalize_ip, -) -from generalresearch.models.thl.user import User +from generalresearch.models.thl.ipinfo import normalize_ip from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig +if TYPE_CHECKING: + from generalresearch.models.thl.ipinfo import GeoIPInformation + from generalresearch.models.thl.user import User + fake = Faker() diff --git a/generalresearch/models/thl/userhealth.py b/generalresearch/models/thl/userhealth.py index f0b7473..e751c2b 100644 --- a/generalresearch/models/thl/userhealth.py +++ b/generalresearch/models/thl/userhealth.py @@ -2,11 +2,12 @@ from __future__ import annotations from datetime import UTC, datetime from enum import Enum -from typing import Self +from typing import TYPE_CHECKING, Any, Self from pydantic import BaseModel, Field, NonNegativeFloat, PositiveInt -from generalresearch.models.custom_types import AwareDatetimeISO +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO class AuditLogLevel(int, Enum): @@ -73,6 +74,6 @@ class AuditLog(BaseModel): return d @classmethod - def from_mysql(cls, d: dict) -> Self: + def from_mysql(cls, d: dict[str, Any]) -> Self: d["created"] = d["created"].replace(tzinfo=UTC) - return AuditLog.model_validate(d) + return cls.model_validate(d) diff --git a/generalresearch/wall_status_codes/__init__.py b/generalresearch/wall_status_codes/__init__.py index cfccb35..37f3960 100644 --- a/generalresearch/wall_status_codes/__init__.py +++ b/generalresearch/wall_status_codes/__init__.py @@ -1,6 +1,7 @@ +from typing import TYPE_CHECKING + from generalresearch.models import Source from generalresearch.models.thl.definitions import Status, StatusCode1 -from generalresearch.models.thl.session import Wall from generalresearch.wall_status_codes import ( cint, dynata, @@ -16,6 +17,9 @@ from generalresearch.wall_status_codes import ( spectrum, ) +if TYPE_CHECKING: + from generalresearch.models.thl.session import Wall + def annotate_status_code( source: Source, |
