aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorstuppie2026-08-19 16:58:43 -0600
committerstuppie2026-08-19 16:58:43 -0600
commit0e60d93cf9dfff98bae4a459288c06c6fd7f2c07 (patch)
tree487ba56058447867c18f3d799602a740c0153d8c
parent91dba28268657c41ebbaec1572258b927c5ae24b (diff)
downloadgeneralresearch-0e60d93cf9dfff98bae4a459288c06c6fd7f2c07.tar.gz
generalresearch-0e60d93cf9dfff98bae4a459288c06c6fd7f2c07.zip
working on BusinessPayoutEvent. create_from_ach_or_wire -> stage payout events all at once. BusinessPayoutEvent standalone model instead of computed from bp_payouts
-rw-r--r--generalresearch/managers/base.py6
-rw-r--r--generalresearch/managers/thl/payout.py202
-rw-r--r--generalresearch/models/thl/payout.py146
-rw-r--r--generalresearch/thl_django/event/models.py5
-rw-r--r--test_utils/managers/ledger/conftest.py4
-rw-r--r--tests/managers/thl/test_payout.py29
6 files changed, 247 insertions, 145 deletions
diff --git a/generalresearch/managers/base.py b/generalresearch/managers/base.py
index b935413..ba42e9e 100644
--- a/generalresearch/managers/base.py
+++ b/generalresearch/managers/base.py
@@ -1,6 +1,7 @@
from __future__ import annotations
from collections.abc import Collection
+from contextlib import nullcontext
from enum import Enum
from generalresearch.pg_helper import PostgresConfig
@@ -45,6 +46,11 @@ class PostgresManager(Manager):
self.pg_config = pg_config
self.permissions = set(permissions) if permissions else set()
+ def connection(self, conn=None):
+ if conn is not None:
+ return nullcontext(conn)
+ return self.pg_config.make_connection()
+
class RedisManager(Manager):
CACHE_PREFIX = None
diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py
index 364e1aa..22856f6 100644
--- a/generalresearch/managers/thl/payout.py
+++ b/generalresearch/managers/thl/payout.py
@@ -58,11 +58,13 @@ class PayoutEventManager(PostgresManagerWithRedis):
access
"""
- res = self.pg_config.execute_sql_query(query=f"""
+ res = self.pg_config.execute_sql_query(
+ query=f"""
SELECT uuid, reference_uuid
FROM ledger_account
WHERE qualified_name LIKE '{thl_lm.currency.value}:bp_wallet:%'
- """)
+ """
+ )
account_to_product = {i["uuid"]: i["reference_uuid"] for i in res}
product_to_account = {i["reference_uuid"]: i["uuid"] for i in res}
@@ -103,13 +105,15 @@ class PayoutEventManager(PostgresManagerWithRedis):
payout_event.update(status=status, ext_ref_id=ext_ref_id, order_data=order_data)
d = payout_event.model_dump_mysql()
- query = sql.SQL("""
+ query = sql.SQL(
+ """
UPDATE event_payout SET
status = %(status)s,
ext_ref_id = %(ext_ref_id)s,
order_data = %(order_data)s
WHERE uuid = %(uuid)s;
- """)
+ """
+ )
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=d)
@@ -715,6 +719,16 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
skip_wallet_balance_check=skip_wallet_balance_check,
)
+ def create_pending_bp_payout_events(
+ self,
+ product: Product,
+ amount: USDCent,
+ payout_type: PayoutType = PayoutType.ACH,
+ ext_ref_id: str | None = None,
+ created: AwareDatetime | None = None,
+ ):
+ pass
+
def create_bp_payout_event(
self,
thl_ledger_manager: ThlLedgerManager,
@@ -936,50 +950,43 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
tx_entry_ids = [tx_entry.id for tx_entry in tx_entries]
# assert len(tx_entry) == len(transaction_ids)*2
- # (5) Delete records
-
- # DELETE: tx_entry
- self.pg_config.execute_write(
- query="""
+ # (5) Delete records (all in 1 tx)
+ with self.pg_config.make_connection() as conn:
+ with conn.cursor() as c:
+ # DELETE: tx_entry
+ c.execute("""
DELETE
FROM ledger_entry
WHERE transaction_id = ANY(%s)
AND id = ANY(%s)
- """,
- params=[transaction_ids, tx_entry_ids],
- )
+ """, [transaction_ids, tx_entry_ids])
- # DELETE: tx_metadata
- self.pg_config.execute_write(
- query="""
+ # DELETE: tx_metadata
+ c.execute("""
DELETE
FROM ledger_transactionmetadata
WHERE transaction_id = ANY(%s)
AND id = ANY(%s)
- """,
- params=[transaction_ids, list(tx_metadata_ids)],
- )
+ """, [transaction_ids, list(tx_metadata_ids)],
+ )
- # DELETE: transactions
- self.pg_config.execute_write(
- query="""
+ # DELETE: transactions
+ c.execute("""
DELETE
FROM ledger_transaction
WHERE id = ANY(%s)
- """,
- params=[transaction_ids],
- )
+ """, [transaction_ids])
- # DELETE: event_payouts
- self.pg_config.execute_write(
- query="""
+ # DELETE: event_payouts
+ c.execute(
+ """
DELETE
FROM event_payout
WHERE ext_ref_id = %s
AND uuid = ANY(%s)
- """,
- params=[ext_ref_id, event_payout_uuids],
- )
+ """, [ext_ref_id, event_payout_uuids],
+ )
+ conn.commit()
return None
@@ -1140,6 +1147,22 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
return allocation
+ def create_business_payout_event(self, bpe: BusinessPayoutEvent):
+
+ with self.pg_config.make_connection() as conn:
+ with conn.cursor() as c:
+ # Quick check to make sure it doesn't already exist
+ c.execute("""
+ SELECT 1
+ FROM supplier_payout
+ WHERE ext_ref_id = %(ext_ref_id)s
+ """, {'ext_ref_id': bpe.ext_ref_id})
+ assert c.fetchone() is None
+
+ c.execute("""
+ """)
+
+
def create_from_ach_or_wire(
self,
business: Business,
@@ -1170,92 +1193,101 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager):
)
assert amount > 100_00, "Must issue Supplier Payouts at least $100 minimum."
- LOG.warning("Paying out ")
+ LOG.warning(f"Paying out {business.name} {amount.to_usd_str()}")
if created:
LOG.warning("Payouts in the past, require the parquet files to be rebuilt.")
- assert created < datetime.now(tz=timezone.utc)
-
- else:
- created = datetime.now(tz=timezone.utc)
+ assert created.tzinfo == timezone.utc, "created must be UTC"
+ assert created < datetime.now(
+ tz=timezone.utc
+ ), "created must be in the past"
# Gather the total amount available balance from each and put into
# a simple DF. We're using the available balance because we need it
- # to always be positive.. and we never want to get into a negative
+ # to always be positive. We never want to get into a negative
# situation again, so it's best to be extra conservative.
- res = {
+ balances = {
pb.product_id: pb.available_balance
for pb in business.balance.product_balances
}
- df = pd.DataFrame.from_dict(res, orient="index").reset_index()
+ df = pd.DataFrame.from_dict(balances, orient="index").reset_index()
df.columns = ["product_id", "available_balance"]
- res = BusinessPayoutEventManager.recoup_proportional(
+ df = BusinessPayoutEventManager.recoup_proportional(
df=df, target_amount=business.balance.recoup
)
# Can't pay any Products that don't have a remaining balance
- res = res[res["remaining_balance"] > 0]
+ df = df[df["remaining_balance"] > 0].copy()
assert (
- res.deduction.sum() == business.balance.recoup
+ df.deduction.sum() == business.balance.recoup
), "recoup_proportional failure"
- res["issue_amount"] = BusinessPayoutEventManager.distribute_amount(
- df=res, amount=amount
+ df["issue_amount"] = BusinessPayoutEventManager.distribute_amount(
+ df=df, amount=amount
)
- assert res.issue_amount.sum() == amount, "issue_amount failure"
+ assert df.issue_amount.sum() == amount, "issue_amount failure"
# Can't pay any Products that don't have an issue amount
- res = res[res["issue_amount"] > 0]
+ df = df[df["issue_amount"] > 0].copy()
- recouped_amounts: list[dict[str, int]] = res[
- ["product_id", "remaining_balance", "issue_amount"]
- ].to_dict(orient="records")
+ amounts: dict[str, dict[str, int]] = df.set_index("product_id")[
+ ["remaining_balance", "issue_amount"]
+ ].to_dict(orient="index")
- # Get all of the products at once so we're not doing it for every interation
- products = pm.get_by_uuids(
- product_uuids=[i["product_id"] for i in recouped_amounts]
+ products = pm.get_by_uuids(product_uuids=list(amounts.keys()))
+ product_lookup = {p.uuid: p for p in products}
+
+ bpe = BusinessPayoutEvent(
+ uuid=uuid4().hex,
+ business_id=business.uuid,
+ payout_type=PayoutType.ACH,
+ amount=amount,
+ created=created,
+ ext_ref_id=transaction_id,
+ # The ACH payment was sent! We haven't yet recorded it
+ # in the ledger, but it was sent by the bank.
+ status=PayoutStatus.COMPLETE,
)
bp_payouts: list[BrokerageProductPayoutEvent] = []
- for idx, item in enumerate(recouped_amounts):
- product = next((p for p in products if p.uuid == item["product_id"]), None)
- assert product is not None
-
- try:
- bp_pe: BrokerageProductPayoutEvent = self.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(item["issue_amount"]),
- created=created + timedelta(milliseconds=idx + 1),
- ext_ref_id=transaction_id,
- skip_wallet_balance_check=True
- )
-
- assert bp_pe.status == PayoutStatus.COMPLETE
- bp_payouts.append(bp_pe)
+ for product_id, item in amounts.items():
+ product = product_lookup[product_id]
+ bp_payouts.append(BrokerageProductPayoutEvent(
+ created=created,
+ payout_type=PayoutType.ACH,
+ status=PayoutStatus.PENDING,
+ uuid=uuid4().hex,
+ amount=USDCent(item["issue_amount"]),
+ ext_ref_id=transaction_id,
+ product_id=product.uuid,
+ # We will fill these in
+ debit_account_uuid=None,
+ cashout_method_uuid=None,
+ ))
+ bpe.bp_payouts = bp_payouts
- except (Exception,) as e:
- # Cleanup bp_payouts
- print("Exception", e)
- return None
- if bp_pe.status == PayoutStatus.FAILED:
- sleep(1)
- try:
- bp_pe = self.retry_create_bp_payout_event_tx(
- thl_ledger_manager=thl_lm,
- product=product,
- payout_event_uuid=bp_pe.uuid,
- )
- assert bp_pe.status == PayoutStatus.COMPLETE
- bp_payouts.append(bp_pe)
+ return BusinessPayoutEvent.model_validate({"bp_payouts": bp_payouts})
- except (Exception,) as e:
- # Cleanup bp_payouts
- return None
- return BusinessPayoutEvent.model_validate({"bp_payouts": bp_payouts})
+# import duckdb
+# conn = duckdb.connect()
+# conn.execute("""
+# select * from read_parquet('/mnt/thl-incite/raw/df-collections/ledger/*/*.parquet')
+# where event_payout is not null
+# and direction =1
+# and reference_uuid in ?
+# """, [b.product_uuids])
+# df = conn.fetch_df()
+# df['ext_description'].value_counts()
+#
+# tx_ids = [35554404, 37210650]
+# conn.execute("""
+# select * from read_parquet('/mnt/thl-incite/raw/df-collections/ledger/*/*.parquet')
+# where tx_id in ?
+# """, [tx_ids])
+# df = conn.fetch_df()
diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py
index 1a9d534..5f51d01 100644
--- a/generalresearch/models/thl/payout.py
+++ b/generalresearch/models/thl/payout.py
@@ -11,6 +11,8 @@ from pydantic import (
PositiveInt,
computed_field,
field_validator,
+ model_validator,
+ ConfigDict,
)
from typing_extensions import Self
@@ -45,22 +47,22 @@ class PayoutEvent(BaseModel):
examples=["9453cd076713426cb68d05591c7145aa"],
)
- debit_account_uuid: UUIDStr = Field(
+ debit_account_uuid: UUIDStr | None = Field(
description="The LedgerAccount.uuid that money is being requested from. "
"Thie User or Brokerage Product is retrievable through the "
"LedgerAccount.reference_uuid",
examples=["18298cb1583846fbb06e4747b5310693"],
)
- cashout_method_uuid: UUIDStr = Field(
+ cashout_method_uuid: UUIDStr | None= Field(
description="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)",
examples=["a6dc1fc1bf934557b952f253dee12813"],
)
- created: AwareDatetimeISO = Field(
- default_factory=lambda: datetime.now(tz=timezone.utc)
+ created: AwareDatetimeISO | None = Field(
+ default = None
)
# In the smallest unit of the currency being transacted. For USD, this
@@ -276,80 +278,118 @@ class BrokerageProductPayoutEvent(PayoutEvent):
class BusinessPayoutEvent(BaseModel):
- """A single ACH or Wire event to a Business Bank Account"""
+ """A single payout event to a supplier Business."""
+ model_config = ConfigDict(validate_assignment=True)
- bp_payouts: list[BrokerageProductPayoutEvent] = Field(
- description="Here is the list of Brokerage Product Payouts that"
- "this Business Payout includes.",
- min_length=1,
+ uuid: UUIDStr = Field(
+ title="Supplier Payout Unique Identifier",
+ examples=["9453cd076713426cb68d05591c7145aa"],
)
- @computed_field(
+ # Used for holding a *unique*, external, payout-type-specific identifier.
+ ext_ref_id: str = Field(title="Unique external reference ID")
+
+ business_id: UUIDStr = Field(
+ description="The Business receiving this supplier payout.",
+ examples=[uuid4().hex],
+ )
+
+ created: AwareDatetimeISO | None = Field(
+ default=None
+ )
+
+ # In the smallest unit of the currency being transacted. For USD, this
+ # is cents.
+ amount: PositiveInt = Field(
+ lt=2**63 - 1,
+ strict=True,
title="Amount",
- description="The amount issued to the Bank Account",
- examples=[19_823_43],
- return_type=USDCent,
+ description="The amount issued to the supplier.",
+ examples=[1_982_343],
)
- @property
- def amount(self) -> USDCent:
- return USDCent(sum([p.amount for p in self.bp_payouts]))
- @computed_field(
- title="Amount USD Str",
- description="The amount issued to the Bank Account as a USD string",
- examples=["$19,823.43"],
- return_type=str,
+ status: PayoutStatus = Field(
+ default=PayoutStatus.PENDING,
+ description=PayoutStatus.as_openapi(),
+ examples=[PayoutStatus.COMPLETE],
)
- @property
- def amount_usd_str(self) -> str:
- return self.amount.to_usd_str()
- @computed_field(
- title="Created",
- description="This is equal to the created time of the first"
- "Brokerage Product Payout Event.",
- return_type=AwareDatetimeISO,
+ payout_type: PayoutType = Field(
+ description=PayoutType.as_openapi(), examples=[PayoutType.ACH]
)
- @property
- def created(self) -> AwareDatetimeISO:
- return self.bp_payouts[0].created
- @computed_field(
- title="Line Items",
- description="The number of sub-payments",
- return_type=PositiveInt,
+ request_data: dict | None = Field(
+ default=None,
+ description="Stores payout-type-specific information that is used to "
+ "request this payout from the external provider.",
+ )
+
+ order_data: dict | None = Field(
+ default=None,
+ description="Stores payout-type-specific order information that is "
+ "returned from the external payout provider.",
+ )
+
+ bp_payouts: list[BrokerageProductPayoutEvent] | None = Field(
+ default=None,
+ description="The list of Brokerage Product Payouts that this Business Payout includes",
+ min_length=1,
)
- @property
- def line_items(self):
- return len(self.bp_payouts)
@computed_field(
- title="External Reference ID",
- description="ACH Transaction ID",
- return_type=str | None,
+ title="Amount USD Str",
+ description="The amount issued to the supplier as a USD string",
+ examples=["$19,823.43"],
+ return_type=str,
)
@property
- def ext_ref_id(self):
- return self.bp_payouts[0].ext_ref_id
+ def amount_usd_str(self) -> str:
+ return USDCent(self.amount).to_usd_str()
# --- Validators ---
+ @field_validator("payout_type", mode="before")
+ @classmethod
+ def normalize_payout_type(cls, v):
+ if isinstance(v, str):
+ try:
+ return PayoutType[v.upper()]
+ except KeyError:
+ raise ValueError(f"Invalid payout_type: {v}")
+ return v
+
@field_validator("bp_payouts", mode="before")
@classmethod
- def normalize_enum(cls, v):
+ def validate_bp_payouts_type(cls, v):
"""This can be a list of Instances or Python Dictionaries depending
on how it's initialized.
"""
+ if v is None:
+ return v
+
assert isinstance(v, list)
+ return v
- def get_field(obj, field):
- if isinstance(obj, dict):
- return obj.get(field)
- return getattr(obj, field, None)
+ @model_validator(mode="after")
+ def validate_bp_payouts(self) -> Self:
+ if not self.bp_payouts:
+ return self
- assert all(
- get_field(i, "ext_ref_id") == get_field(v[0], "ext_ref_id") for i in v
- ), "Not all group values are the same"
+ bp_payout_amount = sum([p.amount for p in self.bp_payouts])
+ if bp_payout_amount != self.amount:
+ raise ValueError(
+ "BusinessPayoutEvent.amount must equal the sum of "
+ f"bp_payouts amounts ({self.amount=} {bp_payout_amount=})"
+ )
- return v
+ invalid_payout_types = [
+ p.payout_type for p in self.bp_payouts if p.payout_type != self.payout_type
+ ]
+ if invalid_payout_types:
+ raise ValueError(
+ "All BrokerageProductPayoutEvent.payout_type values must equal "
+ f"BusinessPayoutEvent.payout_type ({self.payout_type=})"
+ )
+
+ return self
diff --git a/generalresearch/thl_django/event/models.py b/generalresearch/thl_django/event/models.py
index 51a8e2f..7144bc4 100644
--- a/generalresearch/thl_django/event/models.py
+++ b/generalresearch/thl_django/event/models.py
@@ -64,9 +64,9 @@ class SupplierPayout(models.Model):
# generalresearch/models/thl/payout.py:PayoutStatus
status = models.CharField(max_length=20, null=True)
- # Used for holding an external, payouttype-specific identifier.
+ # Used for holding a unique, external, payouttype-specific identifier.
# For ACH, this is the ACH transaction id.
- ext_ref_id = models.CharField(max_length=64, null=True)
+ ext_ref_id = models.CharField(max_length=64, unique=True)
# The allowed values for `payout_type` are defined in generalresearch:
# generalresearch/models/thl/payout.py:PayoutType
@@ -86,7 +86,6 @@ class SupplierPayout(models.Model):
indexes = [
models.Index(fields=["created"]),
models.Index(fields=["business_id"]),
- models.Index(fields=["ext_ref_id"]),
]
diff --git a/test_utils/managers/ledger/conftest.py b/test_utils/managers/ledger/conftest.py
index 0aa6cb3..852645d 100644
--- a/test_utils/managers/ledger/conftest.py
+++ b/test_utils/managers/ledger/conftest.py
@@ -596,8 +596,8 @@ def session_with_tx_factory(
final_status: Status = Status.COMPLETE,
wall_req_cpi: Decimal = Decimal(".50"),
started: datetime = utc_hour_ago,
- ) -> Session:
- s: Session = session_factory(
+ ) -> "Session":
+ s: "Session" = session_factory(
user=user,
wall_count=2,
final_status=final_status,
diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py
index 31087b8..f118c71 100644
--- a/tests/managers/thl/test_payout.py
+++ b/tests/managers/thl/test_payout.py
@@ -5,6 +5,7 @@ from decimal import Decimal
from random import choice as rand_choice, randint
from typing import Optional
from uuid import uuid4
+import io
import pandas as pd
import pytest
@@ -16,7 +17,10 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
from generalresearch.managers.thl.payout import UserPayoutEventManager
from generalresearch.models.thl.definitions import PayoutStatus
from generalresearch.models.thl.ledger import LedgerEntry, Direction
-from generalresearch.models.thl.payout import BusinessPayoutEvent
+from generalresearch.models.thl.payout import (
+ BusinessPayoutEvent,
+ BrokerageProductPayoutEvent,
+)
from generalresearch.models.thl.payout import UserPayoutEvent
from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.ledger import LedgerAccount
@@ -708,7 +712,6 @@ class TestBusinessPayoutEventManager:
assert int(res.deduction.sum()) == 0
def test_distribute_amount(self, business_payout_event_manager):
- import io
df = pd.read_csv(
io.StringIO(
@@ -802,6 +805,28 @@ class TestBusinessPayoutEventManager:
)
assert "Must issue Supplier Payouts at least $100 minimum." in str(cm)
+ bpe = BusinessPayoutEvent(
+ business_id=business.uuid,
+ amount=USDCent(100_00),
+ payout_type=PayoutType.ACH,
+ )
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(47_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ),
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(53_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ),
+ ]
+
def test_ach_payment(
self,
product,