diff options
| author | stuppie | 2026-10-07 12:56:40 -0600 |
|---|---|---|
| committer | stuppie | 2026-10-07 12:56:40 -0600 |
| commit | 7366bab8263759c932b1ef5d577b743cd2615741 (patch) | |
| tree | 2ed192761318b8ba5143fbee142c2c3da8a1373b | |
| parent | ab3b7d17a7b6d7e5033e2f1fd2008d16e356fcfe (diff) | |
| download | generalresearch-7366bab8263759c932b1ef5d577b743cd2615741.tar.gz generalresearch-7366bab8263759c932b1ef5d577b743cd2615741.zip | |
paginated product manager & tests. Add cache_* fields on Product (model and django)
| -rw-r--r-- | generalresearch/managers/base.py | 17 | ||||
| -rw-r--r-- | generalresearch/managers/thl/product.py | 407 | ||||
| -rw-r--r-- | generalresearch/models/thl/product.py | 11 | ||||
| -rw-r--r-- | generalresearch/models/thl/supplier_tag.py | 3 | ||||
| -rw-r--r-- | generalresearch/thl_django/userprofile/models.py | 17 | ||||
| -rw-r--r-- | tests/managers/thl/test_product.py | 13 | ||||
| -rw-r--r-- | tests/managers/thl/test_product_prod.py | 4 |
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( |
