aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorGreg Stupp2026-10-08 04:41:46 +0000
committerGreg Stupp2026-10-08 04:41:46 +0000
commitc3fea9e8dd51e84316cdff985c0b433eb1f97c8f (patch)
treefc2f4d5bca0b4acd15b9153a2329531db55a9569
parentab3b7d17a7b6d7e5033e2f1fd2008d16e356fcfe (diff)
parentb9ec91f4b64c6d3575287c1ea6dcc1182f3157c0 (diff)
downloadgeneralresearch-3.7.1.tar.gz
generalresearch-3.7.1.zip
Merges pull request #7 v3.7.1
Product Manager cached db filed + filters, orders, pagination
-rw-r--r--generalresearch/managers/base.py17
-rw-r--r--generalresearch/managers/thl/product.py542
-rw-r--r--generalresearch/models/thl/finance.py51
-rw-r--r--generalresearch/models/thl/payout.py7
-rw-r--r--generalresearch/models/thl/product.py197
-rw-r--r--generalresearch/models/thl/supplier_tag.py3
-rw-r--r--generalresearch/thl_django/userprofile/models.py15
-rw-r--r--pyproject.toml2
-rw-r--r--tests/managers/thl/test_product.py13
-rw-r--r--tests/managers/thl/test_product_prod.py4
-rw-r--r--tests/models/thl/test_product.py11
11 files changed, 608 insertions, 254 deletions
diff --git a/generalresearch/managers/base.py b/generalresearch/managers/base.py
index 2227d21..257032f 100644
--- a/generalresearch/managers/base.py
+++ b/generalresearch/managers/base.py
@@ -19,7 +19,22 @@ class Permission(int, Enum):
class Manager:
- pass
+ DEFAULT_PAGE_SIZE = 50
+
+ def validate_pagination(
+ self,
+ page: int | None = 1,
+ size: int | None = None,
+ ):
+ page = page if page is not None else 1
+ assert type(page) is int
+ assert page >= 1, "page starts at 1"
+ size = size if size is not None else self.DEFAULT_PAGE_SIZE
+ assert type(size) is int
+ assert 1 <= size <= 100
+ offset = (page - 1) * size
+ paginated_filter_str = f"LIMIT {size} OFFSET {offset}"
+ return paginated_filter_str
class SqlManager(Manager):
diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py
index 66dd131..b730f92 100644
--- a/generalresearch/managers/thl/product.py
+++ b/generalresearch/managers/thl/product.py
@@ -7,16 +7,13 @@ from collections.abc import Collection
from datetime import UTC, datetime
from decimal import Decimal
from threading import Lock
-from typing import TYPE_CHECKING
-from uuid import UUID
+from typing import TYPE_CHECKING, Any
from cachetools import TTLCache, cachedmethod, keys
-from more_itertools import chunked
-from psycopg import Cursor
-from pydantic import ValidationError
+from psycopg import Connection, sql
+from pydantic import NonNegativeInt
from sentry_sdk import capture_exception
-from generalresearch.decorators import LOG
from generalresearch.managers.base import (
PostgresManager,
)
@@ -43,6 +40,22 @@ logger = logging.getLogger()
class ProductManager(PostgresManager):
+ CACHED_FIELDS: frozenset[str] = frozenset({
+ "balance",
+ "user_wallet_balance",
+ "payouts",
+ "pop_financial",
+ "users_active_7d",
+ "task_completes_7d",
+ "balance_net_7d",
+ })
+ CACHED_FIELDS_JSON: frozenset[str] = frozenset({
+ "balance",
+ "user_wallet_balance",
+ "payouts",
+ "pop_financial",
+ })
+
def __init__(
self,
pg_config: PostgresConfig,
@@ -66,53 +79,48 @@ class ProductManager(PostgresManager):
product_uuid: UUIDStr,
) -> Product:
assert is_valid_uuid(product_uuid), "invalid uuid"
- res = self.fetch_uuids(
+ res = self.filter_by(
product_uuids=[product_uuid],
)
- # do this so we uniformly raise AssertionErrors
assert len(res) == 1, "product not found"
return res[0]
+ def get_by_uuid_if_exists(
+ self,
+ product_uuid: UUIDStr,
+ ) -> Product | None:
+ # Do not attach the cache decorator here, since we don't
+ # want to cache a None. The interior call is cached if
+ # the product exists.
+ try:
+ return self.get_by_uuid(product_uuid=product_uuid)
+ except AssertionError as e:
+ if "product not found" in str(e):
+ return None
+ raise
+
def get_by_uuids(
self,
product_uuids: list[UUIDStr],
) -> list[Product]:
-
- res = self.fetch_uuids(
+ res = self.filter_paginated(
product_uuids=product_uuids,
)
assert len(product_uuids) == len(res), "incomplete product response"
return res
- @cachedmethod(
- operator.attrgetter("uuid_cache"), lock=operator.attrgetter("uuid_lock")
- )
- def get_by_uuid_if_exists(
- self,
- product_uuid: UUIDStr,
- ) -> Product | None:
- # many=False, raise_on_error=False
- try:
- return self.fetch_uuids(
- product_uuids=[product_uuid],
- )[0]
- except AssertionError:
- return None
- except IndexError:
- return None
-
def get_by_uuids_if_exists(
self,
product_uuids: list[UUIDStr],
) -> list[Product]:
# Same as .get_by_uuids but doesn't raise Exception if len(product_uuids) != len(res)
- return self.fetch_uuids(
+ return self.filter_paginated(
product_uuids=product_uuids,
)
def get_all(self, rand_limit: int | None) -> list[Product]:
product_uuids = self.get_all_uuids(rand_limit=rand_limit)
- return self.fetch_uuids(product_uuids=product_uuids)
+ return self.filter_paginated(product_uuids=product_uuids)
def get_all_uuids(self, rand_limit: int | None) -> list[UUIDStr]:
@@ -128,143 +136,317 @@ class ProductManager(PostgresManager):
)
else:
- res = self.pg_config.execute_sql_query(query="""
+ res = self.pg_config.execute_sql_query(
+ query="""
SELECT p.id::uuid
FROM userprofile_brokerageproduct AS p
- """)
+ """
+ )
return [i["id"] for i in res]
- def fetch_uuids(
+ def filter_paginated(
self,
product_uuids: list[UUIDStr] | None = None,
business_uuids: list[UUIDStr] | None = None,
team_uuids: list[UUIDStr] | None = None,
- ) -> list[Product]:
- LOG.debug(f"PM.fetch_uuids({product_uuids=}, {business_uuids=}, {team_uuids=})")
-
- assert (
- sum(
- bool(x) # This will also be False is the array is empty
- for x in [product_uuids, business_uuids, team_uuids]
- )
- == 1
- ), "Can only provide one set of identifiers"
-
- filter_column = None
- filter_uuids = None
- if bool(product_uuids):
- assert all(is_valid_uuid(v) for v in product_uuids), "invalid uuid passed"
- filter_column = "id"
- filter_uuids = product_uuids
- elif bool(business_uuids):
- assert all(is_valid_uuid(v) for v in business_uuids), "invalid uuid passed"
- filter_column = "business_id"
- filter_uuids = business_uuids
- elif bool(team_uuids):
- assert all(is_valid_uuid(v) for v in team_uuids), "invalid uuid passed"
- filter_column = "team_id"
- filter_uuids = team_uuids
-
- assert filter_column is not None
-
- if filter_uuids is None or len(filter_uuids) == 0:
- return []
-
- with self.pg_config.make_connection() as sql_connection, sql_connection.cursor() as c:
- res = []
- for chunk in chunked(filter_uuids, 500):
- res.extend(
- self.fetch_uuids_(
- c=c, filter_uuids=chunk, filter_column=filter_column
- )
+ name_like: str | None = None,
+ supplier_tags: list[str] | None = None,
+ harmonizer_domain_like: str | None = None,
+ redirect_url_like: str | None = None,
+ user_wallet_enabled: bool | None = None,
+ order_field: str = "created",
+ descending: bool = False,
+ conn: Connection | None = None,
+ ):
+ """
+ Automatically paginate the results of a filter query.
+ """
+ res = []
+ page = 1
+ size = self.DEFAULT_PAGE_SIZE
+ with self.connection(conn) as conn:
+ while True:
+ _res = self.filter_by(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ harmonizer_domain_like=harmonizer_domain_like,
+ redirect_url_like=redirect_url_like,
+ user_wallet_enabled=user_wallet_enabled,
+ page=page,
+ size=size,
+ order_field=order_field,
+ descending=descending,
+ conn=conn,
)
+ res.extend(_res)
+ if len(_res) < size:
+ break
+ page += 1
return res
- def fetch_uuids_(
- self, c: Cursor, filter_uuids: list[UUIDStr], filter_column: str
- ) -> list[Product]:
- from generalresearch.models.thl.product import Product
+ def filter_page(
+ self,
+ product_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ name_like: str | None = None,
+ supplier_tags: list[str] | None = None,
+ harmonizer_domain_like: str | None = None,
+ redirect_url_like: str | None = None,
+ user_wallet_enabled: bool | None = None,
+ page: int | None = None,
+ size: int | None = None,
+ order_field: str = "created",
+ descending: bool = False,
+ conn: Connection | None = None,
+ ) -> tuple[list[Product], NonNegativeInt]:
+ products = self.filter_by(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ harmonizer_domain_like=harmonizer_domain_like,
+ redirect_url_like=redirect_url_like,
+ user_wallet_enabled=user_wallet_enabled,
+ page=page,
+ size=size,
+ order_field=order_field,
+ descending=descending,
+ conn=conn,
+ )
+ count = self.filter_count(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ harmonizer_domain_like=harmonizer_domain_like,
+ redirect_url_like=redirect_url_like,
+ user_wallet_enabled=user_wallet_enabled,
+ )
+ return products, count
- assert len(filter_uuids) <= 500, "chunk me"
- assert filter_column in {"id", "business_id", "team_id"}
+ def filter_by(
+ self,
+ product_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ name_like: str | None = None,
+ supplier_tags: list[str] | None = None,
+ harmonizer_domain_like: str | None = None,
+ redirect_url_like: str | None = None,
+ user_wallet_enabled: bool | None = None,
+ page: int | None = None,
+ size: int | None = None,
+ order_field: str = "created",
+ descending: bool = False,
+ conn: Connection | None = None,
+ ):
+ from generalresearch.models.thl.product import Product
- # Step 1: Retrieve the basic columns from the "Product table"
+ # Enforce explicit pagination if we've querying for more than 1 item
+ identifiers = [product_uuids, business_uuids, team_uuids]
+ if (
+ (sum(len(x) for x in identifiers if x is not None) > 1)
+ and page is None
+ and size is None
+ ):
+ raise ValueError("must paginate")
+ paginated_filter_str = ""
+ if page is not None or size is not None:
+ paginated_filter_str = self.validate_pagination(page=page, size=size)
+ filter_str, params = self.make_filter_str(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ harmonizer_domain_like=harmonizer_domain_like,
+ redirect_url_like=redirect_url_like,
+ user_wallet_enabled=user_wallet_enabled,
+ )
+ order_fields = {
+ "name": "bp.name",
+ "created": "bp.created",
+ # bp_payment_credit is also called "payout"
+ "bp_payment_credit": "(bp.balance ->> 'bp_payment_credit')::bigint",
+ "adjustment_credit": "(bp.balance ->> 'adjustment_credit')::bigint",
+ "adjustment_debit": "(bp.balance ->> 'adjustment_debit')::bigint",
+ "net": "(bp.balance ->> 'net')::bigint",
+ "payment": "(bp.balance ->> 'payment')::bigint",
+ "balance": "(bp.balance ->> 'balance')::bigint",
+ "available_balance": "(bp.balance ->> 'available_balance')::bigint",
+ "adjustment_percent": "(bp.balance ->> 'adjustment_percent')::bigint",
+ }
+ order_by_sql = (
+ f"{order_fields[order_field]} "
+ f"{'DESC' if descending else 'ASC'} "
+ "NULLS LAST, bp.id"
+ )
query = f"""
- SELECT
- bp.id,
- bp.id_int,
- bp.name,
- bp.enabled,
- bp.created::timestamptz,
- bp.team_id::uuid,
- bp.business_id::uuid,
+ WITH selected_products AS MATERIALIZED (
+ SELECT
+ bp.id,
+ bp.id_int,
+ bp.name,
+ bp.team_id,
+ bp.business_id,
+ bp.created,
+ bp.enabled,
+ bp.payments_enabled,
+ bp.commission,
+ bp.redirect_url,
+ bp.grs_domain,
+ bp.profiling_config,
+ bp.user_health_config,
+ bp.yield_man_config,
+ bp.offerwall_config,
+ bp.session_config,
+ bp.payout_config,
+ bp.user_create_config,
+ bp.balance,
+ bp.user_wallet_balance,
+ bp.users_active_7d,
+ bp.task_completes_7d,
+ bp.balance_net_7d
+ FROM userprofile_brokerageproduct bp
+ {filter_str}
+ ORDER BY {order_by_sql}
+ {paginated_filter_str}
+ )
+ SELECT
+ bp.*,
bp.commission AS commission_pct,
- bp.grs_domain as harmonizer_domain,
- bp.redirect_url,
- bp.session_config::jsonb,
- bp.payout_config::jsonb,
- bp.user_create_config::jsonb,
- bp.offerwall_config::jsonb,
- bp.profiling_config::jsonb,
- bp.user_health_config::jsonb,
- bp.yield_man_config::jsonb,
- t.tags
- FROM userprofile_brokerageproduct AS bp
- LEFT JOIN (
- SELECT product_id, STRING_AGG(tag, ',') as tags
- FROM userprofile_brokerageproducttag
- GROUP BY product_id
- ) t ON t.product_id = bp.id_int
- WHERE {filter_column} = ANY(%s)
+ bp.grs_domain AS harmonizer_domain,
+ COALESCE(t.tags, ARRAY[]::varchar[]) AS tags,
+ sources.value -> 'sources_config' AS sources_config,
+ wallet.value -> 'user_wallet' AS user_wallet_config
+ FROM selected_products bp
+ LEFT JOIN LATERAL (
+ SELECT array_agg(pt.tag) AS tags
+ FROM userprofile_brokerageproducttag pt
+ WHERE pt.product_id = bp.id_int
+ ) t ON true
+ LEFT JOIN userprofile_brokerageproductconfig sources
+ ON sources.product_id = bp.id
+ AND sources.key = 'sources_config'
+ LEFT JOIN userprofile_brokerageproductconfig wallet
+ ON wallet.product_id = bp.id
+ AND wallet.key = 'user_wallet'
+ ORDER BY {order_by_sql}
"""
- c.execute(query, [list(filter_uuids)])
-
- res = c.fetchall()
-
- if len(res) == 0:
- return []
- for x in res:
- x["id"] = UUID(x["id"]).hex
- x["team_id"] = UUID(x["team_id"]).hex if x["team_id"] else None
- x["business_id"] = UUID(x["business_id"]).hex if x["business_id"] else None
- x["tags"] = set(x["tags"].split(",")) if x["tags"] else set()
+ with self.connection(conn) as conn, conn.cursor() as c:
+ c.execute(query, params)
+ res = c.fetchall()
+ products = [Product.model_validate(x) for x in res]
+ return products
- res1 = {i["id"]: i for i in res}
-
- # Step 2: Retrieve additional metadata from the "Product Config table"
- c.execute(
- query="""
- SELECT bpc.product_id::uuid as product_id, bpc.key, bpc.value::jsonb
- FROM userprofile_brokerageproductconfig AS bpc
- WHERE product_id = ANY(%s)
- AND key IN ('sources_config', 'user_wallet')
- """,
- # Pulling from keys b/c no reason to try to retrieve any config
- # k,v rows for products that we know aren't in the other table.
- params=[list(res1.keys())],
+ def filter_count(
+ self,
+ product_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ name_like: str | None = None,
+ supplier_tags: list[str] | None = None,
+ harmonizer_domain_like: str | None = None,
+ redirect_url_like: str | None = None,
+ user_wallet_enabled: bool | None = None,
+ conn: Connection | None = None,
+ ) -> NonNegativeInt:
+ filter_str, params = self.make_filter_str(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ harmonizer_domain_like=harmonizer_domain_like,
+ redirect_url_like=redirect_url_like,
+ user_wallet_enabled=user_wallet_enabled,
)
- kv_res = c.fetchall()
- for item in kv_res:
- item["value"] = item["value"][item["key"]]
- if item["key"] == "user_wallet":
- item["key"] = "user_wallet_config"
-
- # Step 2.1: go through them all, and add the key,vals to the correct
- # Product in the dictionary
- for item in kv_res:
- k: str = item["key"]
- product_id: str = UUID(item["product_id"]).hex
- res1[product_id][k] = item["value"]
- r = []
- for k, v in res1.items():
- try:
- r.append(Product.model_validate(v))
- except ValidationError:
- logger.info(f"failed to parse product: {k}")
- raise
+ query = f"""
+ SELECT COUNT(1) AS cnt
+ FROM userprofile_brokerageproduct AS bp
+ {filter_str}
+ """
+ with self.connection(conn) as conn, conn.cursor() as c:
+ c.execute(query, params)
+ res = c.fetchone()
+ return int(res["cnt"])
- return r
+ @staticmethod
+ def make_filter_str(
+ product_uuids: list[UUIDStr] | None = None,
+ business_uuids: list[UUIDStr] | None = None,
+ team_uuids: list[UUIDStr] | None = None,
+ name_like: str | None = None,
+ supplier_tags: list[str] | None = None,
+ harmonizer_domain_like: str | None = None,
+ redirect_url_like: str | None = None,
+ user_wallet_enabled: bool | None = None,
+ ) -> tuple[str, dict[str, Any]]:
+ params = {}
+ filters = []
+ identifiers = [product_uuids, business_uuids, team_uuids]
+ assert sum([bool(x) for x in identifiers]) == 1, (
+ "Must provide exactly one set of identifiers"
+ )
+ identifier = next(x for x in identifiers if x)
+ assert all(is_valid_uuid(v) for v in identifier), "invalid uuid"
+
+ if product_uuids is not None:
+ filters.append("bp.id = ANY(%(product_uuids)s)")
+ params["product_uuids"] = list(product_uuids)
+ if business_uuids is not None:
+ filters.append("bp.business_id = ANY(%(business_uuids)s)")
+ params["business_uuids"] = list(business_uuids)
+ if team_uuids is not None:
+ filters.append("bp.team_id = ANY(%(team_uuids)s)")
+ params["team_uuids"] = list(team_uuids)
+ if name_like:
+ filters.append("bp.name ILIKE '%%' || %(name_like)s || '%%'")
+ params["name_like"] = name_like
+ if harmonizer_domain_like:
+ filters.append(
+ "bp.grs_domain ILIKE '%%' || %(harmonizer_domain_like)s || '%%'"
+ )
+ params["harmonizer_domain_like"] = harmonizer_domain_like
+ if redirect_url_like:
+ filters.append(
+ "bp.redirect_url ILIKE '%%' || %(redirect_url_like)s || '%%'"
+ )
+ params["redirect_url_like"] = redirect_url_like
+ if user_wallet_enabled is not None:
+ filters.append("""
+ EXISTS (
+ SELECT 1
+ FROM userprofile_brokerageproductconfig AS wallet_filter
+ WHERE wallet_filter.product_id = bp.id
+ AND wallet_filter.key = 'user_wallet'
+ AND (
+ wallet_filter.value
+ -> 'user_wallet'
+ ->> 'enabled'
+ )::boolean = %(user_wallet_enabled)s
+ )
+ """)
+ params["user_wallet_enabled"] = user_wallet_enabled
+ if supplier_tags:
+ filters.append("""
+ EXISTS (
+ SELECT 1
+ FROM userprofile_brokerageproducttag pt
+ WHERE pt.product_id = bp.id_int
+ AND pt.tag = ANY(%(supplier_tags)s::text[])
+ )
+ """)
+ params["supplier_tags"] = list(supplier_tags)
+ assert len(filters) > 0, "must pass at least 1 filter"
+ return "WHERE " + " AND ".join(filters), params
def create(
self,
@@ -363,10 +545,16 @@ class ProductManager(PostgresManager):
insert_data["payments_enabled"] = instance.payments_enabled
try:
- insert_data["id_int"] = next(iter(self.pg_config.execute_sql_query(query="""
+ insert_data["id_int"] = next(
+ iter(
+ self.pg_config.execute_sql_query(
+ query="""
SELECT COALESCE(MAX(id_int), 0) + 1 as id_int
FROM userprofile_brokerageproduct
- """)))["id_int"]
+ """
+ )
+ )
+ )["id_int"]
instance.id_int = insert_data["id_int"]
query = """
@@ -400,7 +588,6 @@ class ProductManager(PostgresManager):
# from pymysql import IntegrityError
# except IntegrityError as e:
except Exception as e:
-
try:
return self.get_by_uuid(product_uuid=instance.id)
except AssertionError:
@@ -527,3 +714,60 @@ class ProductManager(PostgresManager):
conn.commit()
self.cache_clear(product_uuid)
+
+ def update_cached_fields(self, products: list[Product]) -> None:
+ """
+ Use this to update any of the cache_* keys on a product
+ or any of the calculated fields such as users_active_7d.
+ All supported fields are updated at once
+ """
+ updates_by_field = {}
+ for product in products:
+ data = product.model_dump(mode="json", include=set(self.CACHED_FIELDS))
+ for k, v in data.items():
+ if k in self.CACHED_FIELDS_JSON and v is not None:
+ v = json.dumps(v)
+ updates_by_field.setdefault(k, {})[product.id] = v
+ for field, data in updates_by_field.items():
+ self.update_field_bulk(field=field, data=data)
+
+ def update_field_bulk(self, field: str, data: dict[str, Any]) -> None:
+ """
+ Update one field for multiple products.
+
+ data maps product UUID -> new value.
+ """
+ if not data:
+ return
+
+ assert field in self.CACHED_FIELDS, f"Unsupported field: {field!r}"
+
+ value_rows = sql.SQL(", ").join(sql.SQL("(%s, %s)") for _ in data)
+ value_expression = (
+ sql.SQL("updates.value::jsonb")
+ if field in self.CACHED_FIELDS_JSON
+ else sql.SQL("updates.value::integer")
+ )
+
+ query = sql.SQL("""
+ UPDATE userprofile_brokerageproduct AS bp
+ SET {field} = {value_expression}
+ FROM (VALUES {value_rows}) AS updates(product_id, value)
+ WHERE bp.id = updates.product_id::uuid
+ """).format(
+ field=sql.Identifier(field),
+ value_expression=value_expression,
+ value_rows=value_rows,
+ )
+
+ params = [
+ item for product_id, value in data.items() for item in (product_id, value)
+ ]
+
+ with self.pg_config.make_connection() as conn:
+ with conn.cursor() as cursor:
+ cursor.execute(query, params)
+ conn.commit()
+
+ for product_id in data:
+ self.cache_clear(product_id)
diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py
index 0eb4f8f..f5e8745 100644
--- a/generalresearch/models/thl/finance.py
+++ b/generalresearch/models/thl/finance.py
@@ -2,7 +2,7 @@ from __future__ import annotations
import random
from collections.abc import Iterable
-from datetime import UTC
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
from uuid import uuid4
@@ -239,6 +239,13 @@ class POPFinancial(BaseModel):
return res
+class ProductPOPFinancials(BaseModel):
+ periods: list[POPFinancial] = Field(default_factory=list)
+ updated_at: AwareDatetimeISO = Field(
+ default_factory=lambda: datetime.now(tz=UTC),
+ )
+
+
class UserWalletBalances(BaseModel):
"""Cumulative ledger activity and balance for one user's USD wallet.
@@ -408,9 +415,7 @@ class UserWalletBalances(BaseModel):
data.update(
product_id=product_id,
product_user_id=product_user_id,
- last_event=(
- None if wallet_rows.empty else wallet_rows["time_idx"].max()
- ),
+ last_event=(None if wallet_rows.empty else wallet_rows["time_idx"].max()),
)
return cls.model_validate(data)
@@ -937,6 +942,15 @@ class ProductBalances(BaseModel):
"balances elsewhere.",
)
+ commission: NonNegativeInt | None = Field(
+ default=None,
+ description=(
+ "Net commission revenue from this Brokerage Product. "
+ "Positive adjustments increase this value and reversals decrease it."
+ ),
+ examples=[5_038],
+ )
+
# --- Validate ---
@model_validator(mode="after")
def check_unknown_fields(self) -> ProductBalances:
@@ -1001,7 +1015,11 @@ class ProductBalances(BaseModel):
)
@property
def expense(self) -> int:
- return self.user_bonus_credit - self.user_bonus_debit - self.user_payout_complete_debit
+ return (
+ self.user_bonus_credit
+ - self.user_bonus_debit
+ - self.user_payout_complete_debit
+ )
# --- Properties: account related ---
@computed_field(
@@ -1226,6 +1244,11 @@ class ProductBalances(BaseModel):
"issued_payment",
),
(
+ "commission_usd",
+ "Net commission revenue from the Product.",
+ "commission",
+ ),
+ (
"payout_usd",
"Total task payouts earned by the Product.",
"payout",
@@ -1278,9 +1301,12 @@ class ProductBalances(BaseModel):
labels=["product_id"],
)
for product in products:
+ value = getattr(product, attribute)
+ if value is None:
+ continue
metric.add_metric(
[str(product.product_id)],
- int(getattr(product, attribute)) / 100,
+ int(value) / 100,
)
yield metric
@@ -1350,19 +1376,6 @@ class ProductBalances(BaseModel):
).replace("$-", "-$")
-class PrivateProductBalances(ProductBalances):
- """Product balances with internal revenue visible to administrative APIs."""
-
- commission: int = Field(
- default=0,
- description=(
- "Net commission revenue earned by GRL from this Brokerage Product. "
- "Positive adjustments increase this value and reversals decrease it."
- ),
- examples=[5_038],
- )
-
-
class BusinessBalances(BaseModel):
product_balances: list[ProductBalances] = Field(default_factory=list)
diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py
index ab231bc..ecc9b0c 100644
--- a/generalresearch/models/thl/payout.py
+++ b/generalresearch/models/thl/payout.py
@@ -238,6 +238,13 @@ class BrokerageProductPayoutEvent(PayoutEvent):
return self.amount_usd.to_usd_str()
+class ProductPayouts(BaseModel):
+ events: list[BrokerageProductPayoutEvent] = Field(default_factory=list)
+ updated_at: AwareDatetimeISO = Field(
+ default_factory=lambda: datetime.now(tz=UTC),
+ )
+
+
class BusinessPayoutEventCreate(BaseModel):
"""A single payout event to a supplier Business."""
diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py
index 0330207..9d826fa 100644
--- a/generalresearch/models/thl/product.py
+++ b/generalresearch/models/thl/product.py
@@ -48,13 +48,11 @@ from generalresearch.models.custom_types import (
from generalresearch.models.definitions import Source
from generalresearch.models.thl.finance import (
POPFinancial,
- PrivateProductBalances,
ProductBalances,
+ ProductPOPFinancials,
ProductUserWalletBalances,
)
-from generalresearch.models.thl.payout import (
- BrokerageProductPayoutEvent,
-)
+from generalresearch.models.thl.payout import ProductPayouts
from generalresearch.models.thl.payout_format import (
PayoutFormatType,
format_payout_format,
@@ -74,6 +72,9 @@ if TYPE_CHECKING:
from dask.distributed import Client
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.incite.mergers.foundations.enriched_session import (
+ EnrichedSessionMerge,
+ )
from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
@@ -85,7 +86,6 @@ if TYPE_CHECKING:
PRODUCT_BALANCES_METRICS_CACHE_KEY = "metrics:product_balances"
-PRODUCT_PRIVATE_BALANCES_METRICS_CACHE_KEY = "metrics:private_product_balances"
PRODUCT_USER_WALLET_BALANCES_METRICS_CACHE_KEY = "metrics:product_user_wallet_balances"
@@ -969,22 +969,27 @@ class Product(BaseModel, validate_assignment=True):
# Initialization is deferred until unless it's called
# (see .prebuild_***())
+ bp_account: LedgerAccount | None = Field(default=None)
+
balance: ProductBalances | None = Field(default=None, description="Product Balance")
- private_balance: PrivateProductBalances | None = Field(
- default=None, description="Product Balance including private keys"
- )
user_wallet_balance: ProductUserWalletBalances | None = Field(default=None)
-
- payouts_total_str: str | None = Field(default=None)
- payouts_total: USDCent | None = Field(default=None)
- payouts: list[BrokerageProductPayoutEvent] | None = Field(
+ payouts: ProductPayouts | None = Field(
default=None,
description="Product Payouts. These are the ACH or Wire payments that were sent to the"
"Business on behalf of this specific Product",
)
+ pop_financial: ProductPOPFinancials | None = Field(default=None)
- pop_financial: list[POPFinancial] | None = Field(default=None)
- bp_account: LedgerAccount | None = Field(default=None)
+ users_active_7d: int | None = Field(
+ default=None, description="Count of active users in the past 7 days"
+ )
+ task_completes_7d: int | None = Field(
+ default=None, description="Count of completes in the past 7 days"
+ )
+ balance_net_7d: int | None = Field(
+ default=None,
+ description="Net Earnings over the last 7 days (in USD Cents, this can be positive or negative)",
+ )
# --- Validators ---
@field_validator("harmonizer_domain", mode="before")
@@ -1089,6 +1094,83 @@ class Product(BaseModel, validate_assignment=True):
# --- Prebuild ---
@staticmethod
+ def get_enriched_session_metrics_df(
+ product_ids: Collection[UUIDStr],
+ client: Client,
+ enriched_session: EnrichedSessionMerge,
+ ) -> pd.DataFrame:
+ """Load all EnrichedSession rows needed to cache a batch of Products."""
+ now = pd.Timestamp.now(tz="UTC")
+ cutoff = now - pd.Timedelta(days=7)
+
+ ddf = enriched_session.ddf(
+ include_partial=True,
+ force_rr_latest=False,
+ columns=[
+ "product_id",
+ "user_id",
+ "started",
+ "status",
+ ],
+ filters=[
+ ("started", ">=", cutoff.to_pydatetime()),
+ ("started", "<", now.to_pydatetime()),
+ ("product_id", "in", list(product_ids)),
+ ],
+ )
+ if ddf is None:
+ return pd.DataFrame(
+ columns=[
+ "product_id",
+ "users_active_7d",
+ "task_completes_7d",
+ ]
+ )
+ ddf = ddf.assign(is_complete=ddf["status"].eq("c"))
+
+ users = (
+ ddf[["product_id", "user_id"]]
+ .drop_duplicates()
+ .groupby("product_id")
+ .size()
+ .rename("users_active_7d")
+ )
+ completes = (
+ ddf.loc[ddf["is_complete"], ["product_id"]]
+ .groupby("product_id")
+ .size()
+ .rename("task_completes_7d")
+ )
+ users_series, completes_series = client.compute(
+ [users, completes],
+ sync=True,
+ )
+ metrics_df = (
+ users_series.to_frame()
+ .join(
+ completes_series.to_frame(),
+ how="outer",
+ )
+ .fillna(0)
+ .astype(
+ {
+ "users_active_7d": int,
+ "task_completes_7d": int,
+ }
+ )
+ )
+ return metrics_df
+
+ def prebuild_metrics(self, metrics_df: pd.DataFrame) -> None:
+ if self.id in metrics_df.index:
+ row = metrics_df.loc[self.id]
+ self.users_active_7d = int(row["users_active_7d"])
+ self.task_completes_7d = int(row["task_completes_7d"])
+ else:
+ self.users_active_7d = 0
+ self.task_completes_7d = 0
+
+ @staticmethod
def get_pop_ledger_df(
product_ids: Collection[UUIDStr],
client: Client,
@@ -1203,6 +1285,7 @@ class Product(BaseModel, validate_assignment=True):
LOG.debug(f"Product.prebuild_balance({self.uuid=})")
self.balance = None
+ self.balance_net_7d = None
if self.bp_account is None:
self.prefetch_bp_account(thl_lm=thl_lm)
assert self.bp_account is not None
@@ -1217,44 +1300,36 @@ class Product(BaseModel, validate_assignment=True):
balance_df = balance_df.drop(
columns=["account_id", "product_id", "product_user_id"]
).set_index("time_idx")
+ balance_df.index = pd.to_datetime(balance_df.index, utc=True)
+
+ cutoff = pd.Timestamp.now(tz="UTC") - timedelta(days=7)
+ balance_7d_df = balance_df.loc[balance_df.index >= cutoff]
+ self.balance_net_7d = (
+ 0 if balance_7d_df.empty else ProductBalances.from_pandas(balance_7d_df).net
+ )
+
balance = ProductBalances.from_pandas(balance_df)
balance.product_id = self.uuid
- self.balance = balance
-
- def prebuild_private_balance(
- self,
- thl_lm: ThlLedgerManager,
- pop_ledger_df: pd.DataFrame | None = None,
- ) -> None:
- from generalresearch.models.thl.finance import PrivateProductBalances
commission_account = thl_lm.get_account_or_create_bp_commission_by_uuid(
self.uuid
)
-
- df = pop_ledger_df.loc[pop_ledger_df["account_id"].eq(commission_account.uuid)]
- if df.empty:
- LOG.warning(
- f"Product({self.uuid=}).prebuild_private_balance empty dataframe"
- )
- return
-
- s = df[
+ commission_rows = pop_ledger_df.loc[
+ pop_ledger_df["account_id"].eq(commission_account.uuid)
+ ]
+ commission_totals = commission_rows[
[
"bp_payment.CREDIT",
"bp_adjustment.CREDIT",
"bp_adjustment.DEBIT",
]
].sum()
-
- commission = (
- s["bp_payment.CREDIT"]
- + s["bp_adjustment.CREDIT"]
- - s["bp_adjustment.DEBIT"]
- )
- self.private_balance = PrivateProductBalances.model_validate(
- self.balance.model_dump() | {"commission": commission}
+ balance.commission = int(
+ commission_totals["bp_payment.CREDIT"]
+ + commission_totals["bp_adjustment.CREDIT"]
+ - commission_totals["bp_adjustment.DEBIT"]
)
+ self.balance = balance
def prebuild_user_wallet_balances(
self,
@@ -1320,20 +1395,18 @@ class Product(BaseModel, validate_assignment=True):
]
if df.empty:
- self.pop_financial = []
+ self.pop_financial = ProductPOPFinancials()
return
df = df.groupby(
[pd.Grouper(key="time_idx", freq=rr.interval), "account_id"]
).sum()
- from generalresearch.models.thl.finance import POPFinancial
-
- self.pop_financial = POPFinancial.list_from_pandas(
+ periods = POPFinancial.list_from_pandas(
input_data=df, accounts=[self.bp_account]
)
- return
+ self.pop_financial = ProductPOPFinancials(periods=periods)
def prebuild_payouts(
self,
@@ -1342,18 +1415,25 @@ class Product(BaseModel, validate_assignment=True):
LOG.debug(f"Product.prebuild_payouts({self.uuid=})")
from generalresearch.models.thl.ledger import OrderBy
- self.payouts = bp_pem.get_bp_bp_payout_events_for_products(
+ events = bp_pem.get_bp_bp_payout_events_for_products(
product_uuids=[self.uuid],
order_by=OrderBy.DESC,
)
+ self.payouts = ProductPayouts(events=events)
- self.prebuild_payouts_total()
-
- def prebuild_payouts_total(self) -> None:
- assert self.payouts is not None
+ @computed_field
+ @property
+ def payouts_total(self) -> USDCent | None:
+ if self.payouts is None:
+ return None
+ return USDCent(sum(event.amount for event in self.payouts.events))
- self.payouts_total = USDCent(sum([po.amount for po in self.payouts]))
- self.payouts_total_str = self.payouts_total.to_usd_str()
+ @computed_field
+ @property
+ def payouts_total_str(self) -> str | None:
+ if self.payouts_total is None:
+ return None
+ return self.payouts_total.to_usd_str()
def set_cache(
self,
@@ -1400,15 +1480,10 @@ class Product(BaseModel, validate_assignment=True):
thl_lm=thl_lm,
pop_ledger_df=pop_ledger_df,
)
- if self.balance:
- self.prebuild_private_balance(
- thl_lm=thl_lm,
+ if self.balance and self.user_wallet_enabled:
+ self.prebuild_user_wallet_balances(
pop_ledger_df=pop_ledger_df,
)
- if self.user_wallet_enabled:
- self.prebuild_user_wallet_balances(
- pop_ledger_df=pop_ledger_df,
- )
self.prebuild_payouts(bp_pem=bp_pem)
self.prebuild_pop_financial(
thl_lm=thl_lm,
@@ -1428,12 +1503,6 @@ class Product(BaseModel, validate_assignment=True):
key=self.uuid,
value=self.balance.model_dump_json(),
)
- if self.private_balance is not None:
- pipe.hset(
- name=PRODUCT_PRIVATE_BALANCES_METRICS_CACHE_KEY,
- key=self.uuid,
- value=self.private_balance.model_dump_json(),
- )
if self.user_wallet_balance is not None:
pipe.hset(
name=PRODUCT_USER_WALLET_BALANCES_METRICS_CACHE_KEY,
diff --git a/generalresearch/models/thl/supplier_tag.py b/generalresearch/models/thl/supplier_tag.py
index 739b895..1f65042 100644
--- a/generalresearch/models/thl/supplier_tag.py
+++ b/generalresearch/models/thl/supplier_tag.py
@@ -3,10 +3,7 @@ from enum import StrEnum
class SupplierTag(StrEnum):
"""Available tags which can be used to annotate supplier traffic
-
- Note: should not include commas!
"""
-
MOBILE = "mobile"
JS_OFFERWALL = "js-offerwall"
DOI = "double-opt-in"
diff --git a/generalresearch/thl_django/userprofile/models.py b/generalresearch/thl_django/userprofile/models.py
index b474bef..445a53c 100644
--- a/generalresearch/thl_django/userprofile/models.py
+++ b/generalresearch/thl_django/userprofile/models.py
@@ -79,6 +79,21 @@ class BrokerageProduct(models.Model):
# Store configuration regarding user creation. See: models/thl/product.py:UserCreateConfig
user_create_config = models.JSONField(default=dict)
+ # ProductBalances model
+ balance = models.JSONField(default=None, null=True)
+ # ProductUserWalletBalances model
+ user_wallet_balance = models.JSONField(default=None, null=True)
+ # ProductPayouts model
+ payouts = models.JSONField(default=None, null=True)
+ # ProductPOPFinancials model
+ pop_financial = models.JSONField(default=None, null=True)
+
+ users_active_7d = models.IntegerField(default=None, null=True)
+ task_completes_7d = models.IntegerField(default=None, null=True)
+ # Net Earnings over the last 7 days (in USD Cents, this can be positive or negative)
+ balance_net_7d = models.IntegerField(default=None, null=True)
+
+
class Meta:
db_table = "userprofile_brokerageproduct"
diff --git a/pyproject.toml b/pyproject.toml
index 3424e20..d7dec1b 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "generalresearch"
-version = "3.7.0"
+version = "3.7.1"
description = "Python Utilities for General Research"
readme = "README.md"
requires-python = ">=3.14"
diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py
index 8d72fa5..ebbb8d6 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -86,8 +86,8 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid=product.id)
assert isinstance(instance, Product)
- instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
- assert instance == None
+ instance = product_manager.get_by_uuid_if_exists(product_uuid=uuid4().hex)
+ assert instance is None
def test_get_by_uuids_if_exists(
self, product_factory: Callable[..., Product], product_manager: ProductManager
@@ -128,7 +128,8 @@ class TestProductManagerGetMethods:
):
business_ids = [uuid4().hex for _ in range(5)]
- product_manager.fetch_uuids(business_uuids=business_ids)
+ res = product_manager.filter_paginated(business_uuids=business_ids)
+ assert len(res) == 0
for business_id in business_ids:
product_factory(
@@ -139,10 +140,11 @@ class TestProductManagerGetMethods:
name=f"Test Product ID #{uuid4().hex[:6]}",
user_create_config=None,
)
+ res = product_manager.filter_paginated(business_uuids=business_ids)
+ assert len(res) == len(business_ids)
class TestProductManagerCreation:
-
def test_base(
self, product_factory: Callable[..., Product], product_manager: ProductManager
):
@@ -156,7 +158,6 @@ class TestProductManagerCreation:
class TestProductManagerCreate:
-
def test_create_simple(self, product_manager: ProductManager):
# Always required: product_id, team_id, name, redirect_url
# Required internally - if not passed use default: harmonizer_domain,
@@ -354,7 +355,6 @@ class TestProductManager:
class TestProductManagerUpdate:
-
def test_update(
self, product_factory: Callable[..., Product], product_manager: ProductManager
):
@@ -377,7 +377,6 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
-
def test_cache_clear(
self, product_factory: Callable[..., Product], product_manager: ProductManager
):
diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py
index d584527..ed13718 100644
--- a/tests/managers/thl/test_product_prod.py
+++ b/tests/managers/thl/test_product_prod.py
@@ -51,7 +51,7 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuids(
product_uuids=[p.id for p in products] + ["abc123"]
)
- assert "invalid uuid passed" in str(cm.value)
+ assert "invalid uuid" in str(cm.value)
def test_get_by_uuid_if_exists(
self, product_factory: Callable[..., Product], product_manager: ProductManager
@@ -61,7 +61,7 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid=products[0].id)
assert isinstance(instance, Product)
- instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
+ instance = product_manager.get_by_uuid_if_exists(product_uuid=uuid4().hex)
assert instance is None
def test_get_by_uuids_if_exists(
diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py
index ce3242e..f8f3da3 100644
--- a/tests/models/thl/test_product.py
+++ b/tests/models/thl/test_product.py
@@ -873,8 +873,7 @@ class TestProductFinancials:
body, content_type = p1.balance.to_prometheus()
- p1.prebuild_private_balance(thl_lm=thl_ledger_manager, pop_ledger_df=df)
- assert p1.private_balance.commission == 5 * 2
+ assert p1.balance.commission == 5 * 2
p1.prebuild_user_wallet_balances(pop_ledger_df=df)
assert p1.user_wallet_balance.outstanding_liability == 38 * 2
@@ -952,15 +951,11 @@ class TestProductFinancials:
p1.prebuild_balance(thl_lm=thl_ledger_manager, pop_ledger_df=df)
assert p1.balance.payout == 95
-
- p1.prebuild_private_balance(thl_lm=thl_ledger_manager, pop_ledger_df=df)
- assert p1.private_balance.commission == 5
+ assert p1.balance.commission == 5
p2.prebuild_balance(thl_lm=thl_ledger_manager, pop_ledger_df=df)
assert p2.balance.payout == 57
-
- p2.prebuild_private_balance(thl_lm=thl_ledger_manager, pop_ledger_df=df)
- assert p2.private_balance.commission == 5
+ assert p2.balance.commission == 5
p2.prebuild_user_wallet_balances(pop_ledger_df=df)
assert p2.user_wallet_balance.outstanding_liability == 38