aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorstuppie2026-10-07 12:56:40 -0600
committerstuppie2026-10-07 12:56:40 -0600
commit7366bab8263759c932b1ef5d577b743cd2615741 (patch)
tree2ed192761318b8ba5143fbee142c2c3da8a1373b
parentab3b7d17a7b6d7e5033e2f1fd2008d16e356fcfe (diff)
downloadgeneralresearch-7366bab8263759c932b1ef5d577b743cd2615741.tar.gz
generalresearch-7366bab8263759c932b1ef5d577b743cd2615741.zip
paginated product manager & tests. Add cache_* fields on Product (model and django)
-rw-r--r--generalresearch/managers/base.py17
-rw-r--r--generalresearch/managers/thl/product.py407
-rw-r--r--generalresearch/models/thl/product.py11
-rw-r--r--generalresearch/models/thl/supplier_tag.py3
-rw-r--r--generalresearch/thl_django/userprofile/models.py17
-rw-r--r--tests/managers/thl/test_product.py13
-rw-r--r--tests/managers/thl/test_product_prod.py4
7 files changed, 310 insertions, 162 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..c10be38 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
+from pydantic import NonNegativeInt
from sentry_sdk import capture_exception
-from generalresearch.decorators import LOG
from generalresearch.managers.base import (
PostgresManager,
)
@@ -66,53 +63,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 +120,255 @@ 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,
+ 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,
+ 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,
+ page: int | None = None,
+ size: int | None = None,
+ order_field: str = "created",
+ descending: bool = False,
+ conn: Connection | None = None,
+ ) -> tuple[list[Product], int]:
+ products = self.filter_by(
+ product_uuids=product_uuids,
+ business_uuids=business_uuids,
+ team_uuids=team_uuids,
+ name_like=name_like,
+ supplier_tags=supplier_tags,
+ 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,
+ )
+ 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,
+ 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,
+ )
+ order_fields = {
+ "name": "bp.name",
+ "created": "bp.created",
+ }
+ order_by_sql = (
+ f"{order_fields[order_field]} {'DESC' if descending else 'ASC'}, 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.cache_balance,
+ bp.cache_user_wallet_balance,
+ bp.cache_updated_at,
+ 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,
+ bp.cache_balance AS balance,
+ bp.cache_balance AS private_balance,
+ bp.cache_user_wallet_balance AS user_wallet_balance,
+ 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()
+ 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
- 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()
-
- 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,
+ 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,
)
- 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,
+ ) -> 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 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 +467,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 +510,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:
diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py
index 0330207..7188390 100644
--- a/generalresearch/models/thl/product.py
+++ b/generalresearch/models/thl/product.py
@@ -986,6 +986,17 @@ class Product(BaseModel, validate_assignment=True):
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")
def harmonizer_domain_https(cls, s: str | None):
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..813dac0 100644
--- a/generalresearch/thl_django/userprofile/models.py
+++ b/generalresearch/thl_django/userprofile/models.py
@@ -79,6 +79,23 @@ class BrokerageProduct(models.Model):
# Store configuration regarding user creation. See: models/thl/product.py:UserCreateConfig
user_create_config = models.JSONField(default=dict)
+ # PrivateProductBalances model (can be converted to ProductBalances)
+ cache_balance = models.JSONField(default=None, null=True)
+ # ProductUserWalletBalances model
+ cache_user_wallet_balance = models.JSONField(default=None, null=True)
+ # list[BrokerageProductPayoutEvent]
+ cache_payouts = models.JSONField(default=None, null=True)
+ # list[POPFinancial]
+ cache_pop_financial = models.JSONField(default=None, null=True)
+ # We should update all the cache_* fields at once
+ cache_updated_at = models.DateTimeField(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/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(