aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--generalresearch/models/thl/payout.py55
-rw-r--r--generalresearch/models/thl/wallet/payout.py218
-rw-r--r--pyproject.toml2
-rw-r--r--tests/managers/thl/test_ledger/test_user_txs.py41
4 files changed, 76 insertions, 240 deletions
diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py
index b0cf880..ab231bc 100644
--- a/generalresearch/models/thl/payout.py
+++ b/generalresearch/models/thl/payout.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import json
from datetime import UTC, datetime
-from typing import TYPE_CHECKING, Self
+from typing import Self
from uuid import uuid4
from pydantic import (
@@ -114,36 +114,49 @@ class PayoutEvent(BaseModel):
self.order_data = order_data
def check_status_change_allowed(self, status: PayoutStatus) -> None:
+ # Allowed status changes:
+ # PENDING -> {APPROVED, REJECTED, CANCELLED, FAILED, COMPLETE}
+ # APPROVED -> {FAILED, COMPLETE}
+ # FAILED -> {APPROVED, REJECTED, CANCELLED, COMPLETE}
+ # {REJECTED, CANCELLED, COMPLETE} are final
- # We may not be changing the status when this method gets called. It's
- # possible to be called when we're updating other attributes so
- # allow immediate bypass if it isn't actually different.
- if self.status == status:
- return
-
- if self.status in {
+ assert self.status not in {
PayoutStatus.REJECTED,
PayoutStatus.CANCELLED,
PayoutStatus.COMPLETE,
- }:
- raise ValueError(f"status {self.status} is final. No changes allowed")
+ }, f"status {self.status} is final. No changes allowed"
- if self.status == PayoutStatus.PENDING:
- assert status != PayoutStatus.PENDING, "status is already PENDING!"
+ # Updating other attributes may leave the status unchanged.
+ if self.status == status:
+ return
- elif self.status == PayoutStatus.APPROVED:
- assert status in {
+ allowed_status_changes = {
+ PayoutStatus.PENDING: {
+ PayoutStatus.APPROVED,
+ PayoutStatus.REJECTED,
+ PayoutStatus.CANCELLED,
PayoutStatus.FAILED,
PayoutStatus.COMPLETE,
- }, f"status APPROVED can only be FAILED or COMPLETED, not {status}"
-
- elif self.status == PayoutStatus.FAILED:
- assert status in {
+ },
+ PayoutStatus.APPROVED: {
+ PayoutStatus.FAILED,
+ PayoutStatus.COMPLETE,
+ },
+ PayoutStatus.FAILED: {
+ PayoutStatus.APPROVED,
+ PayoutStatus.REJECTED,
PayoutStatus.CANCELLED,
PayoutStatus.COMPLETE,
- }, f"status FAILED can only be CANCELLED or COMPLETED, not {status}"
- else:
- raise ValueError("this shouldn't happen")
+ },
+ }
+ assert self.status in allowed_status_changes, (
+ f"status {self.status} cannot transition to {status}"
+ )
+ assert status in allowed_status_changes[self.status], (
+ f"status {self.status} can only be "
+ f"{', '.join(sorted(x.value for x in allowed_status_changes[self.status]))}, "
+ f"not {status}"
+ )
# --- ORM ---
diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py
deleted file mode 100644
index 1fc0f77..0000000
--- a/generalresearch/models/thl/wallet/payout.py
+++ /dev/null
@@ -1,218 +0,0 @@
-from __future__ import annotations
-
-import json
-from collections.abc import Collection
-from datetime import UTC, datetime
-from typing import Any
-from uuid import uuid4
-
-from pydantic import (
- BaseModel,
- Field,
- PositiveInt,
- computed_field,
- field_validator,
-)
-
-from generalresearch.currency import USDCent
-from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr
-from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.wallet.cashout_method import (
- CashMailOrderData,
-)
-from generalresearch.models.thl.wallet.definitions import PayoutType
-
-
-class PayoutEvent(BaseModel, validate_assignment=True):
- """A user has requested to be paid from their wallet balance."""
-
- uuid: UUIDStr = Field(
- default_factory=lambda: uuid4().hex,
- examples=["9453cd076713426cb68d05591c7145aa"],
- )
-
- # This is the LedgerAccount.uuid that this money is being requested
- # from. The user/BP is retrievable through the LedgerAccount.reference_uuid
- debit_account_uuid: UUIDStr = Field(examples=["18298cb1583846fbb06e4747b5310693"])
-
- # These two fields are copied here from the LedgerAccount through the
- # debit_account_uuid for convenience. They will get populated if the
- # PayoutEventManager retrieves a PayoutEvent from the db.
- account_reference_type: str | None = Field(default=None)
- account_reference_uuid: UUIDStr | None = Field(default=None)
-
- # References a row in the account_cashoutmethod table. This is the
- # cashout method that was used to request this payout. (A cashout is
- # the same thing as a payout)
- cashout_method_uuid: UUIDStr = Field(examples=["a6dc1fc1bf934557b952f253dee12813"])
-
- # By default, this will just be the cashout_method.name. This also is
- # populated from the db and so does not need to be set (there is no
- # `description` field in event_payout)
- description: str | None = Field(default=None)
- created: AwareDatetimeISO = Field(default_factory=lambda: datetime.now(tz=UTC))
-
- # In the smallest unit of the currency being transacted. For USD, this
- # is cents.
- amount: PositiveInt = Field(
- lt=2**63 - 1,
- strict=True,
- description="The USDCent amount int. This cannot be 0 or negative",
- examples=[531],
- )
-
- status: PayoutStatus = Field(
- default=PayoutStatus.PENDING,
- description=PayoutStatus.as_openapi(),
- examples=[PayoutStatus.COMPLETE],
- )
-
- # Used for holding an external, payout-type-specific identifier
- ext_ref_id: str | None = Field(default=None)
- payout_type: PayoutType = Field(
- description=PayoutType.as_openapi(), examples=[PayoutType.ACH]
- )
-
- # Stores payout-type-specific information that is used to request this
- # payout from the external provider.
- request_data: dict[str, Any] = Field(default_factory=dict)
-
- # Stores payout-type-specific order information that is returned from
- # the external payout provider.
- order_data: dict[str, Any] | CashMailOrderData | None = Field(default=None)
-
- @field_validator("payout_type", mode="before")
- @classmethod
- def normalize_enum(cls, v):
- if isinstance(v, str):
- try:
- return PayoutType[v.upper()]
- except KeyError:
- raise ValueError(f"Invalid payout_type: {v}")
- return v
-
- def update(
- self,
- status: PayoutStatus,
- ext_ref_id: str | None = None,
- order_data: dict[str, Any] | None = None,
- ) -> None:
- # These 3 things are the only modifiable attributes
- self.check_status_change_allowed(status)
- self.status = status
- self.ext_ref_id = ext_ref_id
- self.order_data = order_data
-
- def check_status_change_allowed(self, status: PayoutStatus) -> None:
- if self.status in {
- PayoutStatus.REJECTED,
- PayoutStatus.CANCELLED,
- PayoutStatus.COMPLETE,
- }:
- raise ValueError(f"status {self.status} is final. No changes allowed")
-
- if self.status == PayoutStatus.PENDING:
- assert status != PayoutStatus.PENDING, "status is already PENDING!"
-
- elif self.status == PayoutStatus.APPROVED:
- assert status in {
- PayoutStatus.FAILED,
- PayoutStatus.COMPLETE,
- }, f"status APPROVED can only be FAILED or COMPLETED, not {status}"
-
- elif self.status == PayoutStatus.FAILED:
- assert status in {
- PayoutStatus.CANCELLED,
- PayoutStatus.COMPLETE,
- }, f"status FAILED can only be CANCELLED or COMPLETED, not {status}"
-
- else:
- raise ValueError("this shouldn't happen")
-
- def model_dump_mysql(self) -> dict[str, Any]:
- d = self.model_dump(mode="json")
-
- if "created" in d:
- d["created"] = self.created.replace(tzinfo=None)
- if d.get("request_data") is not None:
- d["request_data"] = json.dumps(self.request_data)
- if d.get("order_data") is not None:
- assert self.order_data
-
- if isinstance(self.order_data, dict):
- d["order_data"] = json.dumps(self.order_data)
- else:
- d["order_data"] = self.order_data.model_dump_json()
- return d
-
-
-class BPPayoutEvent(BaseModel):
- uuid: UUIDStr = Field(
- title="Brokerage Product Payout ID",
- description="Unique identifier for the Payout Event",
- examples=["9453cd076713426cb68d05591c7145aa"],
- )
-
- product_id: UUIDStr = Field(
- description="The Brokerage Product that was paid out",
- examples=["1108d053e4fa47c5b0dbdcd03a7981e7"],
- )
-
- created: AwareDatetimeISO = Field(
- description="When the Brokerage Product was paid out",
- default_factory=lambda: datetime.now(tz=UTC),
- )
-
- amount: USDCent = Field(
- lt=2**63 - 1,
- strict=True,
- description="The USDCent amount int. This cannot be 0 or negative",
- examples=[531],
- )
-
- status: PayoutStatus | None = Field(
- default=PayoutStatus.PENDING,
- description=PayoutStatus.as_openapi(),
- examples=[PayoutStatus.COMPLETE],
- )
-
- method: PayoutType = Field(
- title="Payout Method",
- description=PayoutType.as_openapi(),
- examples=[PayoutType.ACH],
- )
-
- @computed_field(return_type=str, examples=["$10,000.000"])
- @property
- def amount_usd(self) -> str:
- return self.amount.to_usd_str()
-
- @staticmethod
- def from_pe(
- payout_events: Collection[PayoutEvent],
- account_product_mapping: dict[str, str],
- order_by: str = "ASC",
- ) -> list[BPPayoutEvent]:
- res = []
- for pe in payout_events:
- bp_pe = BPPayoutEvent.model_validate(
- {
- "uuid": pe.uuid,
- "product_id": account_product_mapping[pe.debit_account_uuid],
- "created": pe.created,
- "amount": USDCent(pe.amount),
- "status": pe.status,
- "method": pe.payout_type,
- }
- )
- res.append(bp_pe)
-
- match order_by:
- case "ASC":
- sorted_list = sorted(res, key=lambda x: x.created, reverse=False)
- case "DESC":
- sorted_list = sorted(res, key=lambda x: x.created, reverse=True)
- case _:
- raise ValueError("Invalid order provided..")
-
- return sorted_list
diff --git a/pyproject.toml b/pyproject.toml
index 46ab5c3..95076c7 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "generalresearch"
-version = "3.5.9"
+version = "3.6.0"
description = "Python Utilities for General Research"
readme = "README.md"
requires-python = ">=3.14"
diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py
index 6b6ef5b..4354c43 100644
--- a/tests/managers/thl/test_ledger/test_user_txs.py
+++ b/tests/managers/thl/test_ledger/test_user_txs.py
@@ -139,6 +139,47 @@ def test_user_txs(
assert sorted([tx.amount for tx in tx_adj_c]) == [-38, 76]
+def test_user_payout_cancel_to_user_tx(
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ delete_ledger_db: Callable[..., None],
+ user_payout_event_manager: UserPayoutEventManager,
+ settings: GRLBaseSettings,
+):
+ delete_ledger_db()
+ create_main_accounts()
+
+ user = user_factory(product=product_amt_true)
+ account = thl_ledger_manager.get_account_or_create_user_wallet(user)
+ pe = user_payout_event_manager.create(
+ uuid=uuid4().hex,
+ debit_account_uuid=account.uuid,
+ cashout_method_uuid=settings.amt_bonus_cashout_method_id,
+ amount=100,
+ payout_type=PayoutType.AMT_BONUS,
+ request_data={},
+ )
+ thl_ledger_manager.create_tx_user_payout_request(
+ user=user,
+ payout_event=pe,
+ skip_wallet_balance_check=True,
+ )
+ thl_ledger_manager.create_tx_user_payout_cancelled(user=user, payout_event=pe)
+
+ txs = thl_ledger_manager.get_user_txs(user)
+ cancel = next(
+ tx
+ for tx in txs.transactions
+ if tx.tx_type == TransactionType.USER_PAYOUT_CANCEL
+ )
+ assert cancel.amount == 100
+ assert cancel.description == "Payout Cancelled"
+ assert cancel.payout_id == pe.uuid
+ assert cancel.url == f"https://fsb.generalresearch.com/{user.product_id}/cashout/{pe.uuid}/"
+
+
def test_user_txs_pagination(
user_factory: Callable[..., User],
product_amt_true: Product,