diff options
| author | stuppie | 2026-09-03 13:40:38 -0600 |
|---|---|---|
| committer | stuppie | 2026-09-03 13:40:38 -0600 |
| commit | 0eaba21734a77287d00d5f82e2f34c90f46081ac (patch) | |
| tree | 465af34ab905de56284abbf3cbe692f13d877faf | |
| parent | 48863e9431d50fd86405036b55b333b7734b00ac (diff) | |
| download | generalresearch-0eaba21734a77287d00d5f82e2f34c90f46081ac.tar.gz generalresearch-0eaba21734a77287d00d5f82e2f34c90f46081ac.zip | |
fixing pydantic model rebuild in user wall and session. assert isinstance(v, AwareDatetime) does not work like that; remove
| -rw-r--r-- | generalresearch/models/thl/session.py | 192 | ||||
| -rw-r--r-- | generalresearch/models/thl/user.py | 23 |
2 files changed, 108 insertions, 107 deletions
diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 5812c75..b3cb6af 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -24,16 +24,21 @@ from generalresearch.models.custom_types import ( IPvAnyAddressStr, UUIDStr, ) -from generalresearch.models.definitions import Source +from generalresearch.models.definitions import DeviceType, Source +from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_CODE_1_2, WALL_ALLOWED_STATUS_STATUS_CODE, + ReportValue, SessionAdjustedStatus, + SessionStatusCode2, Status, StatusCode1, WallAdjustedStatus, WallStatusCode2, ) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User from generalresearch.models.thl.utils import ( decimal_to_int_cents, int_cents_to_decimal, @@ -43,14 +48,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - from generalresearch.models.definitions import DeviceType - from generalresearch.models.legacy.bucket import Bucket - from generalresearch.models.thl.definitions import ( - ReportValue, - SessionStatusCode2, - ) - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.user import User logger = logging.getLogger("Wall") @@ -127,9 +124,9 @@ class WallBase(BaseModel): @classmethod def check_cpi_decimal_places(cls, v: Decimal) -> Decimal: if v is not None: - assert ( - v.as_tuple().exponent >= -5 - ), "Must have 5 or fewer decimal places ('XXX.YYYYY')" + assert v.as_tuple().exponent >= -5, ( + "Must have 5 or fewer decimal places ('XXX.YYYYY')" + ) return v @model_validator(mode="before") @@ -161,24 +158,24 @@ class WallBase(BaseModel): assert self.started <= datetime.now(tz=UTC), "Started must not be in the future" if self.finished: assert self.finished > self.started, "Finished must be after started" - assert self.finished - self.started <= timedelta( - minutes=90 - ), "Maximum wall event time is 90 min" + assert self.finished - self.started <= timedelta(minutes=90), ( + "Maximum wall event time is 90 min" + ) return self @model_validator(mode="after") def check_ext_statuses(self): if self.ext_status_code_3 is not None: - assert ( - self.ext_status_code_1 is not None - ), "Set ext_status_code_1 before ext_status_code_3" - assert ( - self.ext_status_code_2 is not None - ), "Set ext_status_code_2 before ext_status_code_3" + assert self.ext_status_code_1 is not None, ( + "Set ext_status_code_1 before ext_status_code_3" + ) + assert self.ext_status_code_2 is not None, ( + "Set ext_status_code_2 before ext_status_code_3" + ) if self.ext_status_code_2 is not None: - assert ( - self.ext_status_code_1 is not None - ), "Set ext_status_code_1 before ext_status_code_2" + assert self.ext_status_code_1 is not None, ( + "Set ext_status_code_1 before ext_status_code_2" + ) return self @model_validator(mode="after") @@ -186,27 +183,27 @@ class WallBase(BaseModel): if self.status in {Status.COMPLETE, Status.FAIL}: assert self.finished is not None, "finished should be set" if self.status == Status.COMPLETE: - assert ( - self.status_code_1 == StatusCode1.COMPLETE - ), "status_code_1 should be COMPLETE" + assert self.status_code_1 == StatusCode1.COMPLETE, ( + "status_code_1 should be COMPLETE" + ) return self @model_validator(mode="after") def check_status_status_code_agreement(self) -> Self: if self.status_code_1: options = WALL_ALLOWED_STATUS_STATUS_CODE.get(self.status, {}) - assert ( - self.status_code_1 in options - ), f"If status is {self.status.value}, status_code_1 should be in {options}" + assert self.status_code_1 in options, ( + f"If status is {self.status.value}, status_code_1 should be in {options}" + ) return self @model_validator(mode="after") def check_status_code1_2_agreement(self) -> Self: if self.status_code_2: options = WALL_ALLOWED_STATUS_CODE_1_2.get(self.status_code_1, {}) - assert ( - self.status_code_2 in options - ), f"If status_code_1 is {self.status_code_1.value}, status_code_2 should be in {options}" + assert self.status_code_2 in options, ( + f"If status_code_1 is {self.status_code_1.value}, status_code_2 should be in {options}" + ) return self # --- Methods --- @@ -417,15 +414,15 @@ class Wall(WallBase): @model_validator(mode="after") def check_adjusted_null(self) -> Self: if self.adjusted_status is not None or self.adjusted_cpi is not None: - assert ( - self.adjusted_cpi is not None - ), "Set adjusted_cpi if the wall has been adjusted" - assert ( - self.adjusted_status is not None - ), "Set adjusted_status if the wall has been adjusted" - assert ( - self.adjusted_timestamp is not None - ), "Set adjusted_timestamp if the wall has been adjusted" + assert self.adjusted_cpi is not None, ( + "Set adjusted_cpi if the wall has been adjusted" + ) + assert self.adjusted_status is not None, ( + "Set adjusted_status if the wall has been adjusted" + ) + assert self.adjusted_timestamp is not None, ( + "Set adjusted_timestamp if the wall has been adjusted" + ) return self @model_validator(mode="after") @@ -460,7 +457,6 @@ class Wall(WallBase): class WallOut(WallBase): - # These get serialized to the enum name instead of the int value (for ease in UI) status_code_1: Annotated[StatusCode1, EnumNameSerializer] | None = Field( default=None, @@ -518,9 +514,9 @@ class WallOut(WallBase): @classmethod def check_cpi_decimal_places(cls, v: Decimal | None) -> Decimal | None: if v is not None: - assert ( - v.as_tuple().exponent >= -5 - ), "Must have 5 or fewer decimal places ('XXX.YYYYY')" + assert v.as_tuple().exponent >= -5, ( + "Must have 5 or fewer decimal places ('XXX.YYYYY')" + ) return v @field_validator("status_code_1", mode="before") @@ -673,9 +669,9 @@ class Session(BaseModel): @classmethod def check_payout_decimal_places(cls, v: Decimal) -> Decimal: if v is not None: - assert ( - v.as_tuple().exponent >= -2 - ), "Must have 2 or fewer decimal places ('XXX.YY')" + assert v.as_tuple().exponent >= -2, ( + "Must have 2 or fewer decimal places ('XXX.YY')" + ) # explicitly make sure it is 2 decimal places, after checking that it is already 2 or less. v = v.quantize(Decimal("0.00")) return v @@ -697,17 +693,21 @@ class Session(BaseModel): StatusCode1.PS_FAIL, StatusCode1.PS_QUALITY, StatusCode1.PS_BLOCKED, - }, f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + }, ( + f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + ) elif self.status in {Status.TIMEOUT, Status.ABANDON}: assert self.status_code_1 in { StatusCode1.PS_ABANDON, StatusCode1.GRS_ABANDON, StatusCode1.BUYER_ABANDON, - }, f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + }, ( + f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + ) elif self.status == Status.COMPLETE: - assert ( - self.status_code_1 == StatusCode1.COMPLETE - ), f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + assert self.status_code_1 == StatusCode1.COMPLETE, ( + f"status_code_1 {self.status_code_1.name} invalid for status {self.status.value}" + ) else: assert self.status_code_1 is None, ( f"status_code_1 {self.status_code_1.name} invalid for status " @@ -730,9 +730,9 @@ class Session(BaseModel): @model_validator(mode="after") def check_payout_when_complete(self): if self.status == Status.COMPLETE: - assert ( - self.payout is not None - ), "there should be a payout if the session is marked complete" + assert self.payout is not None, ( + "there should be a payout if the session is marked complete" + ) return self # @model_validator(mode='after') @@ -758,19 +758,19 @@ class Session(BaseModel): @model_validator(mode="after") def check_adjusted(self): if self.adjusted_status is not None or self.adjusted_payout is not None: - assert ( - self.adjusted_payout is not None - ), "Set adjusted_payout if the session has been adjusted" - assert ( - self.adjusted_status is not None - ), "Set adjusted_status if the session has been adjusted" - assert ( - self.adjusted_timestamp is not None - ), "Set adjusted_timestamp if the session has been adjusted" + assert self.adjusted_payout is not None, ( + "Set adjusted_payout if the session has been adjusted" + ) + assert self.adjusted_status is not None, ( + "Set adjusted_status if the session has been adjusted" + ) + assert self.adjusted_timestamp is not None, ( + "Set adjusted_timestamp if the session has been adjusted" + ) if self.adjusted_user_payout is not None: - assert ( - self.adjusted_payout is not None - ), "Set adjusted_payout if adjusted_user_payout is set" + assert self.adjusted_payout is not None, ( + "Set adjusted_payout if adjusted_user_payout is set" + ) # NOTE: the other way around is NOT required! # (the adjusted_user_payout / user_payout can be null) return self @@ -783,9 +783,9 @@ class Session(BaseModel): "the adjusted_status should be null" ) if self.adjusted_status == SessionAdjustedStatus.ADJUSTED_TO_FAIL: - assert ( - self.status == Status.COMPLETE - ), "Session.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL" + assert self.status == Status.COMPLETE, ( + "Session.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL" + ) return self # --- Properties --- @@ -1090,9 +1090,9 @@ class Session(BaseModel): return False if self.status == Status.COMPLETE: - assert ( - self.adjusted_status != SessionAdjustedStatus.ADJUSTED_TO_COMPLETE - ), "Can't have complete adj to complete" + assert self.adjusted_status != SessionAdjustedStatus.ADJUSTED_TO_COMPLETE, ( + "Can't have complete adj to complete" + ) if self.adjusted_status in { None, SessionAdjustedStatus.PAYOUT_ADJUSTMENT, @@ -1236,19 +1236,19 @@ def check_adjusted_status_consistent( assert adjusted_cpi == cpi, "adjusted_cpi should be equal to the original cpi" elif adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL: - assert ( - status == Status.COMPLETE - ), "Wall.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL" - assert ( - adjusted_cpi == 0 - ), "adjusted_cpi should be 0 if adjusted_status is ADJUSTED_TO_FAIL" + assert status == Status.COMPLETE, ( + "Wall.status must be COMPLETE for the adjusted_status to be ADJUSTED_TO_FAIL" + ) + assert adjusted_cpi == 0, ( + "adjusted_cpi should be 0 if adjusted_status is ADJUSTED_TO_FAIL" + ) elif adjusted_status == WallAdjustedStatus.CPI_ADJUSTMENT: # the original status is allowed to be anything # the adjusted cpi should be something different - assert ( - adjusted_cpi != 0 and adjusted_cpi != cpi - ), "If CPI_ADJUSTMENT, the adjusted_cpi should be different from the original cpi or 0" + assert adjusted_cpi != 0 and adjusted_cpi != cpi, ( + "If CPI_ADJUSTMENT, the adjusted_cpi should be different from the original cpi or 0" + ) elif adjusted_status is None: assert adjusted_cpi is None, "incompatible adjusted values" @@ -1309,21 +1309,21 @@ def _check_adjusted_status_wall_consistent( # status / adjusted_status agreement if status == Status.COMPLETE: - assert ( - new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_COMPLETE - ), "adjusted status can't be ADJUSTED_TO_COMPLETE if the status is already COMPLETE" + assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_COMPLETE, ( + "adjusted status can't be ADJUSTED_TO_COMPLETE if the status is already COMPLETE" + ) elif status == Status.FAIL: - assert ( - new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL - ), "adjusted status can't be ADJUSTED_TO_FAIL if the status is already FAIL" + assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL, ( + "adjusted status can't be ADJUSTED_TO_FAIL if the status is already FAIL" + ) else: # status is None/timeout/abandon, which we treat as a fail anyway - assert ( - new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL - ), "attempt is already a failure" + assert new_adjusted_status != WallAdjustedStatus.ADJUSTED_TO_FAIL, ( + "attempt is already a failure" + ) # adjusted_status / new_adjusted_status agreement if new_adjusted_status == WallAdjustedStatus.CPI_ADJUSTMENT: - assert ( - new_adjusted_cpi != adjusted_cpi - ), f"adjusted_cpi is already {adjusted_cpi}" + assert new_adjusted_cpi != adjusted_cpi, ( + f"adjusted_cpi is already {adjusted_cpi}" + ) diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 4f94270..9944944 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -20,21 +20,23 @@ from pydantic import ( ) from sentry_sdk import set_tag, set_user -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, +) from generalresearch.models.definitions import MAX_INT32 +from generalresearch.models.thl.ipinfo import GeoIPInformation +from generalresearch.models.thl.ledger import LedgerTransaction +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.userhealth import AuditLog if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager - from generalresearch.models.thl.ipinfo import GeoIPInformation - from generalresearch.models.thl.ledger import LedgerTransaction - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.userhealth import AuditLog from generalresearch.pg_helper import PostgresConfig - # from generalresearch.managers.thl.userhealth import UserIpHistoryManager logger = logging.getLogger() @@ -122,27 +124,24 @@ class User(BaseModel): # noinspection PyNestedDecoratorsk @field_validator("created", "last_seen") @classmethod - def check_not_in_future(cls, v: AwareDatetime | None) -> AwareDatetime: + def check_not_in_future(cls, v: AwareDatetime | None) -> AwareDatetime | None: if v is not None: try: assert v < datetime.now(tz=UTC) except AssertionError: raise ValueError("Input is in the future") - - assert isinstance(v, AwareDatetime) return v # noinspection PyNestedDecorators @field_validator("created", "last_seen") @classmethod - def check_after_anno_domini(cls, v: AwareDatetime | None) -> AwareDatetime: + def check_after_anno_domini(cls, v: AwareDatetime | None) -> AwareDatetime | None: if v is not None: try: assert v > datetime(year=2016, month=7, day=13, tzinfo=UTC) except AssertionError: raise ValueError("Input is before Anno Domini") - assert isinstance(v, AwareDatetime) return v @model_validator(mode="after") @@ -320,3 +319,5 @@ BPUIDStr = Annotated[ StringConstraints(min_length=3, max_length=128), AfterValidator(User.check_product_user_id), ] + +User.model_rebuild() |
