diff options
| author | stuppie | 2026-09-07 11:36:07 -0600 |
|---|---|---|
| committer | stuppie | 2026-09-07 11:36:07 -0600 |
| commit | 242579a44855873d5e054e375440e9d3492cd682 (patch) | |
| tree | 4c30256ad73a0f15a592e9fbf30a665b15ac0717 | |
| parent | 3338f74a94d0624bf894ebb35bd1bcfca268216e (diff) | |
| parent | f6ee73468e27cfddd34b24f00ffd712bccd03347 (diff) | |
| download | generalresearch-242579a44855873d5e054e375440e9d3492cd682.tar.gz generalresearch-242579a44855873d5e054e375440e9d3492cd682.zip | |
Merge branch 'master' into dev
13 files changed, 303 insertions, 223 deletions
diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 3e5d5a0..d83886b 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -7,7 +7,6 @@ import shutil import subprocess import warnings from collections.abc import Callable, Sequence -from concurrent.futures import Future from datetime import UTC, datetime, timedelta from os import R_OK, access, listdir from os.path import isdir @@ -21,7 +20,6 @@ from typing import ( ) from uuid import uuid4 -import dask import dask.dataframe as dd import pandas as pd import pandera as pa @@ -52,10 +50,8 @@ from generalresearch.incite.schemas import ( from generalresearch.models.custom_types import AwareDatetimeISO if TYPE_CHECKING: - from generalresearch.incite.collections.base import DFCollection, DFCollectionItem - from generalresearch.incite.collections.thl_marketplaces import ( - DFCollectionType, - ) + from generalresearch.incite.collections import DFCollection, DFCollectionItem + from generalresearch.incite.collections.thl_marketplaces import DFCollectionType from generalresearch.incite.mergers.base import MergeCollection, MergeType Collection = DFCollection | MergeCollection @@ -100,9 +96,9 @@ class GRLDatasets(BaseModel): # Create the base folders and confirm we have read access self.data_src.mkdir(parents=True, exist_ok=True) - assert access( - path=self.data_src, mode=R_OK - ), f"can't access data_src: {self.data_src}" + assert access(path=self.data_src, mode=R_OK), ( + f"can't access data_src: {self.data_src}" + ) for enum_type in [MergeType, DFCollectionType]: for et in enum_type: @@ -439,26 +435,6 @@ class CollectionBase(BaseModel): return ddf # --- Methods: Cleanup --- - def schedule_cleanup( - self, client: DaskClient | None = None, sync: bool = True, client_resources=None - ) -> pd.DataFrame | Future: - LOG.info(f"cleanup(archive_path={self.archive_path})") - - fs = [] - for item in self.items: - fs.append(dask.delayed(item.cleanup_partials)()) - fs.append(dask.delayed(item.clear_corrupt_archive)()) - fs.append(dask.delayed(self.clear_tmp_archives)()) - - assert isinstance(client, DaskClient) - res = client.compute( - collections=fs, - sync=sync, - priority=2, - client_resources=client_resources, - ) - return res - def cleanup(self) -> None: # Same as schedule_cleanup but runs locally self.cleanup_partials() @@ -516,7 +492,6 @@ class CollectionBase(BaseModel): # --- Regular Path --- if os.path.exists(reg_path): - if os.path.isfile(reg_path): # These should never be a file, clean up os.remove(reg_path) @@ -541,13 +516,13 @@ class CollectionBase(BaseModel): # Make sure these are in the same dir. b/c the symlink has to be # relative, not an absolute path - assert ( - item.path.parent == highest_version.parent - ), "Can't have numbered_path in a different directory" + assert item.path.parent == highest_version.parent, ( + "Can't have numbered_path in a different directory" + ) try: pq.ParquetDataset(highest_version).read().to_pandas() - except (pa.ArrowInvalid, pa.ArrowIOError, FileNotFoundError): + except Exception: # If the most recent version isn't valid, we don't want to # create a symlink to it. # TODO: We could try to be smart and iterate down the most recent @@ -599,7 +574,6 @@ class CollectionBase(BaseModel): # TODO: This appears to be a bug. It should be using the # IntervalRange overlaps approach - Max 2024-06-07 if item.start >= since: - # We want to retrieve the item that falls before the # first item, so we aren't missing any partial time ranges if first_match and idx != 0: @@ -664,9 +638,9 @@ class CollectionItemBase(BaseModel): """We don't want to support CollectionItems that start on a fractional second. """ - assert ( - self.start.microsecond == 0 - ), "CollectionItem.start must not have microsecond precision" + assert self.start.microsecond == 0, ( + "CollectionItem.start must not have microsecond precision" + ) return self # --- Properties --- @@ -704,26 +678,23 @@ class CollectionItemBase(BaseModel): @property def path(self) -> Path: - return Path( - os.path.join(self._collection.archive_path, self.filename) - ) + # Do not use _filepath_adapter.validate_python here as that validates + # that the file exists on disk, which is doesn't until we get the path + # we want to write it to... + return Path(self._collection.archive_path, self.filename) @property def partial_path(self) -> Path: - return Path( - os.path.join(self._collection.archive_path, self.partial_filename) - ) + return Path(self._collection.archive_path, self.partial_filename) @property def empty_path(self) -> Path: - return Path( - os.path.join(self._collection.archive_path, self.empty_filename) - ) + return Path(self._collection.archive_path, self.empty_filename) # --- Methods --- @staticmethod - def path_exists(generic_path: FilePath) -> bool: + def path_exists(generic_path: Path) -> bool: return os.path.exists(generic_path) @staticmethod @@ -757,12 +728,10 @@ class CollectionItemBase(BaseModel): # regex = re.compile(r'\.parquet\.[0-9a-f]{32}', re.I) builds = [] for fn in os.listdir(coll.archive_path): - if ( - fn.startswith(self.filename) - and fn != self.filename - and fn != self.partial_filename - ): - builds.append(fn) + if fn.startswith(self.filename): + # Don't include the "broken link" or mmfsymlink text file + if fn != self.filename and fn != self.partial_filename: + builds.append(fn) if len(builds) == 0: return None @@ -782,9 +751,7 @@ class CollectionItemBase(BaseModel): return f"{self.filename}.{uuid4().hex}" def tmp_path(self) -> Path: - return Path( - os.path.join(self._collection.archive_path, self.tmp_filename()) - ) + return Path(self._collection.archive_path, self.tmp_filename()) # --- --- --- --- # If it has a partial, it isn't always going to be a partial. However, @@ -813,7 +780,6 @@ class CollectionItemBase(BaseModel): def delete_archive(generic_path: Path) -> None: # If a partial directory or file exists, delete it. if os.path.exists(generic_path): - if os.path.isfile(generic_path): os.remove(generic_path) @@ -835,9 +801,9 @@ class CollectionItemBase(BaseModel): return datetime.now(tz=UTC) > self.finish + archive_after def set_empty(self): - assert ( - self.should_archive() - ), "Can not set_empty on an item that is not archive-able" + assert self.should_archive(), ( + "Can not set_empty on an item that is not archive-able" + ) assert not self.is_empty(), "set_empty is already set; why are you doing this?" self.empty_path.touch() assert self.is_empty(), "set_empty(): something is wrong" @@ -971,6 +937,7 @@ class CollectionItemBase(BaseModel): # with a symlink. It does not matter if the item is archiveable or not. if target_path is None: target_path = self.partial_path + target_path = Path(target_path) fps = glob.glob(target_path.as_posix() + ".*") fps = {x for x in fps if x.split(".")[-1].isnumeric()} diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index ed13cf2..1673860 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import pandas as pd from pandas import Period @@ -11,6 +11,9 @@ from pydantic import AwareDatetime, NaiveDatetime from redis import Redis from generalresearch.managers.leaderboard import country_timezone +from generalresearch.managers.thl.user_manager.user_metadata_manager import ( + UserMetadataManager, +) from generalresearch.models.thl.leaderboard import ( Leaderboard, LeaderboardCode, @@ -43,7 +46,6 @@ class LeaderboardManager: self.freq = freq self.product_id = product_id self.country_iso = country_iso - self.within_time_aware = None if within_time is None: self.within_time_aware = datetime.now(tz=UTC).astimezone(self.timezone) elif within_time.tzinfo is not None: @@ -83,39 +85,53 @@ class LeaderboardManager: def get_row_count(self) -> int: # How many rows (unique users) does this leaderboard have? - return self.redis_client.zcard(self.key) or 0 + # redis-py types command responses broadly, but ZCARD returns an integer. + return cast(int, self.redis_client.zcard(self.key)) def get_leaderboard_rows( self, limit: int | None = None, + user_metadata_manager: UserMetadataManager | None = None, ) -> list[LeaderboardRow]: - limit = limit if limit else 0 + limit = limit or 0 res = self.redis_client.zrange( self.key, start=0, end=limit - 1, withscores=True, desc=True ) # We re-rank using pandas min value for ties. Redis does not consider ties in ranking. - s = pd.DataFrame(res, columns=["bpuid", "value"]).sort_values( + s = pd.DataFrame(res, columns=pd.Index(["bpuid", "value"])).sort_values( by="value", ascending=False ) s["rank"] = s["value"].rank(method="min", ascending=False) - return [ - LeaderboardRow(bpuid=r.bpuid, value=r.value, rank=r.rank) - for r in s.itertuples() + rows = [ + LeaderboardRow(bpuid=bpuid, value=value, rank=rank) + for bpuid, value, rank in s.itertuples(index=False, name=None) ] + if user_metadata_manager: + um = user_metadata_manager.filter_by_bpuids( + product_id=self.product_id, product_user_ids={r.bpuid for r in rows} + ) + for row in rows: + row.display_name = um[row.bpuid].display_name + return rows def get_personal_leaderboard_rows( - self, bp_user_id: str, limit: int | None = 5 + self, + bp_user_id: str, + limit: int | None = 5, + user_metadata_manager: UserMetadataManager | None = None, ) -> list[LeaderboardRow]: # We can't just grab this user's rank and nearby rows b/c redis does # not handle ties the same way we do (in redis, each value is a # unique rank, we use lowest rank for all ties). So we have to just # grab everything, then filter - limit = limit if limit is not None else 5 - rows = self.get_leaderboard_rows() + rows = self.get_leaderboard_rows( + user_metadata_manager=user_metadata_manager, limit=None + ) rows = sorted(rows, key=lambda x: x.value, reverse=True) user_indices = [ (i, row) for i, row in enumerate(rows) if row.bpuid == bp_user_id ] + limit = limit or 5 if not user_indices: return rows[: limit * 2] user_idx = user_indices[0][0] @@ -133,15 +149,25 @@ class LeaderboardManager: self, limit: int | None = None, bp_user_id: str | None = None, + user_metadata_manager: UserMetadataManager | None = None, ) -> Leaderboard: - + """ + Returns the leaderboard instance with populated rows. + :param limit: Return limit rows. If bp_user_id, default 5 rows above + below, + else default: no limit / all rows. + :param bp_user_id: If passed, the rows surrounding this user are returned. + :param user_metadata_manager: If passed, the user's display_names are looked up + and populated. + """ if bp_user_id: rows = self.get_personal_leaderboard_rows( - bp_user_id=bp_user_id, limit=limit + bp_user_id=bp_user_id, + limit=limit, + user_metadata_manager=user_metadata_manager, ) else: rows = self.get_leaderboard_rows( - limit=limit, + limit=limit, user_metadata_manager=user_metadata_manager ) total = self.get_row_count() @@ -164,25 +190,25 @@ class LeaderboardManager: ) def hit_complete_count(self, product_user_id: str) -> None: - assert ( - self.board_code == LeaderboardCode.COMPLETE_COUNT - ), "wrong kind of leaderboard" + assert self.board_code == LeaderboardCode.COMPLETE_COUNT, ( + "wrong kind of leaderboard" + ) self.redis_client.zincrby(self.key, amount=1, value=product_user_id) self.redis_client.expire(self.key, time=self.expiration) def hit_sum_payouts(self, product_user_id: str, user_payout: Decimal) -> None: - assert ( - self.board_code == LeaderboardCode.SUM_PAYOUTS - ), "wrong kind of leaderboard" + assert self.board_code == LeaderboardCode.SUM_PAYOUTS, ( + "wrong kind of leaderboard" + ) self.redis_client.zincrby( self.key, amount=round(user_payout * 100), value=product_user_id ) self.redis_client.expire(self.key, time=self.expiration) def hit_largest_payout(self, product_user_id: str, user_payout: Decimal) -> None: - assert ( - self.board_code == LeaderboardCode.LARGEST_PAYOUT - ), "wrong kind of leaderboard" + assert self.board_code == LeaderboardCode.LARGEST_PAYOUT, ( + "wrong kind of leaderboard" + ) # Only sets the value if the new value is greater than the existing self.redis_client.zadd( self.key, {product_user_id: round(user_payout * 100)}, gt=True @@ -191,18 +217,19 @@ class LeaderboardManager: def hit(self, session: Session) -> None: user = session.user + assert user.product_user_id is not None match self.board_code: case LeaderboardCode.COMPLETE_COUNT: - return self.hit_complete_count(product_user_id=user.product_user_id) + self.hit_complete_count(product_user_id=user.product_user_id) case LeaderboardCode.SUM_PAYOUTS: - return self.hit_sum_payouts( + assert session.user_payout is not None + self.hit_sum_payouts( product_user_id=user.product_user_id, user_payout=session.user_payout, ) case LeaderboardCode.LARGEST_PAYOUT: - return self.hit_largest_payout( + assert session.user_payout is not None + self.hit_largest_payout( product_user_id=user.product_user_id, user_payout=session.user_payout, ) - - return None diff --git a/generalresearch/managers/thl/user_manager/memcached_user_manager.py b/generalresearch/managers/thl/user_manager/memcached_user_manager.py deleted file mode 100644 index d2c68ed..0000000 --- a/generalresearch/managers/thl/user_manager/memcached_user_manager.py +++ /dev/null @@ -1,49 +0,0 @@ -# from typing import List, Optional -# -# import pylibmc -# -# from generalresearch.models.thl.user import User -# -# -# class MemcachedUserManager: -# def __init__(self, servers: List[str], cache_prefix: Optional[str] = None): -# self.servers = servers -# self.cache_prefix = cache_prefix if cache_prefix else "user-lookup" -# -# def create_client(self): -# # Clients are NOT thread safe. Make a new one each time -# -# # There's a receive_timeout and send_timeout also, but the documentation is incomprehensible, -# # and they don't seem to do anything??? (I tested setting them at 1ms and I can't -# # get it to fail) -# # https://sendapatch.se/projects/pylibmc/behaviors.html -# mc_client = pylibmc.Client(servers=self.servers, binary=True, -# behaviors={'connect_timeout': 100}) -# return mc_client -# -# def get_user(self, *, product_id: str = None, product_user_id: str = None, user_id: int = None, -# user_uuid: UUIDStr = None) -> User: -# # assume we did input validation in user_manager.get_user() function -# mc_client = self.create_client() -# if user_uuid: -# d = mc_client.get(f"{self.cache_prefix}:uuid:{user_uuid}") -# elif user_id: -# d = mc_client.get(f"{self.cache_prefix}:user_id:{user_id}") -# else: -# d = mc_client.get(f"{self.cache_prefix}:ubp:{product_id}:{product_user_id}") -# if d: -# return User.model_validate_json(d) -# -# def set_user(self, user: User): -# d = user.to_json() -# mc_client = self.create_client() -# mc_client.set(f"{self.cache_prefix}:uuid:{user.uuid}", d, time=60 * 60 * 24) -# mc_client.set(f"{self.cache_prefix}:user_id:{user.user_id}", d, time=60 * 60 * 24) -# mc_client.set(f"{self.cache_prefix}:ubp:{user.product_id}:{user.product_user_id}", d, time=60 * 60 * 24) -# -# def clear_user(self, user: User): -# # this should only be used by tests -# mc_client = self.create_client() -# mc_client.delete(f"{self.cache_prefix}:uuid:{user.uuid}") -# mc_client.delete(f"{self.cache_prefix}:user_id:{user.user_id}") -# mc_client.delete(f"{self.cache_prefix}:ubp:{user.product_id}:{user.product_user_id}") diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index 0b4b8a4..d2732ed 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -1,20 +1,21 @@ from __future__ import annotations import logging +import operator from collections.abc import Collection from datetime import UTC, datetime -from functools import lru_cache +from threading import Lock from typing import TYPE_CHECKING from uuid import uuid4 import psycopg +from cachetools import LRUCache, cachedmethod from psycopg import sql from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.pg_helper import PostgresConfig logging.basicConfig() @@ -26,6 +27,8 @@ class MysqlUserManager: def __init__(self, pg_config: PostgresConfig, is_read_replica: bool): self.pg_config = pg_config self.is_read_replica = is_read_replica + self.product_id_exists_cache = LRUCache(maxsize=5000) + self.product_id_exists_cache_lock = Lock() def _set_last_seen(self, user: User) -> None: # Don't call this directly. Use UserManager.set_last_seen() @@ -40,6 +43,39 @@ class MysqlUserManager: params=[now, user.user_id], ) + def _change_product_user_id(self, *, user: User, new_product_user_id: str) -> User: + """Change a user's supplier-provided ID in the primary database.""" + assert not self.is_read_replica + assert user.user_id is not None + assert user.product_id is not None + assert user.product_user_id is not None + + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" + UPDATE thl_user + SET product_user_id = %(new_product_user_id)s + WHERE id = %(user_id)s + AND product_id = %(product_id)s + AND product_user_id = %(old_product_user_id)s + RETURNING id AS user_id, product_id, product_user_id, + uuid, blocked, created, last_seen + """, + params={ + "new_product_user_id": new_product_user_id, + "user_id": user.user_id, + "product_id": user.product_id, + "old_product_user_id": user.product_user_id, + }, + ) + row = c.fetchone() + + if row is None: + raise RuntimeError( + "User was not updated; it may have been changed concurrently" + ) + return User.from_db(row) + def get_user_from_mysql( self, *, @@ -53,16 +89,16 @@ class MysqlUserManager: logger.info( f"get_user_from_mysql: {product_id}, {product_user_id}, {user_id}, {user_uuid}" ) - assert ( - (product_id and product_user_id) or user_id or user_uuid - ), "Must pass either (product_id, product_user_id), or user_id, or uuid" + assert (product_id and product_user_id) or user_id or user_uuid, ( + "Must pass either (product_id, product_user_id), or user_id, or uuid" + ) if product_id or product_user_id: - assert ( - product_id and product_user_id - ), "Must pass both product_id and product_user_id" - assert ( - sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1 - ), "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid" + assert product_id and product_user_id, ( + "Must pass both product_id and product_user_id" + ) + assert sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1, ( + "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid" + ) # Using RR: Assume we check redis first for newly created users if can_use_read_replica is False: @@ -176,8 +212,11 @@ class MysqlUserManager: return user - @lru_cache(maxsize=5_000) - def product_id_exists(self, product_id: str): + @cachedmethod( + operator.attrgetter("product_id_exists_cache"), + lock=operator.attrgetter("product_id_exists_cache_lock"), + ) + def product_id_exists(self, product_id: str) -> bool: # 'id' is the primary key, there can only be 0 or 1 query = """ SELECT id @@ -228,9 +267,9 @@ class MysqlUserManager: assert product_id, "must pass product_id" assert len(product_user_ids) > 0, "must pass 1 or more product_user_ids" assert len(product_user_ids) <= 500, "limit 500 product_user_ids" - assert isinstance( - product_user_ids, (list, set) - ), "must pass a collection of product_user_ids" + assert isinstance(product_user_ids, (list, set)), ( + "must pass a collection of product_user_ids" + ) res = self.pg_config.execute_sql_query( query=""" SELECT id AS user_id, product_id, product_user_id, @@ -253,14 +292,14 @@ class MysqlUserManager: user_ids: Collection[int] | None = None, user_uuids: Collection[str] | None = None, ) -> list[User]: - assert (user_ids or user_uuids) and not ( - user_ids and user_uuids - ), "Must pass ONE of user_ids, user_uuids" + assert (user_ids or user_uuids) and not (user_ids and user_uuids), ( + "Must pass ONE of user_ids, user_uuids" + ) if user_ids: + assert isinstance(user_ids, (list, set)), ( + "must pass a collection of user_ids" + ) assert len(user_ids) <= 500, "limit 500 user_ids" - assert isinstance( - user_ids, (list, set) - ), "must pass a collection of user_ids" res = self.pg_config.execute_sql_query( query=""" @@ -273,10 +312,10 @@ class MysqlUserManager: params={"user_ids": user_ids}, ) else: + assert isinstance(user_uuids, (list, set)), ( + "must pass a collection of user_uuids" + ) assert len(user_uuids) <= 500, "limit 500 user_uuids" - assert isinstance( - user_uuids, (list, set) - ), "must pass a collection of user_uuids" res = self.pg_config.execute_sql_query( query=""" SELECT id AS user_id, product_id, product_user_id, diff --git a/generalresearch/managers/thl/user_manager/redis_user_manager.py b/generalresearch/managers/thl/user_manager/redis_user_manager.py index 2190e6e..730cb9e 100644 --- a/generalresearch/managers/thl/user_manager/redis_user_manager.py +++ b/generalresearch/managers/thl/user_manager/redis_user_manager.py @@ -76,7 +76,6 @@ class RedisUserManager: p.execute() def clear_user(self, user: User) -> None: - # this should only be used by tests with self.client.pipeline(transaction=False) as p: p.delete(f"{self.cache_prefix}:uuid:{user.uuid}") p.delete(f"{self.cache_prefix}:user_id:{user.user_id}") diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 52e0567..df3fb6b 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -1,11 +1,13 @@ from __future__ import annotations import logging +import operator from collections.abc import Collection from datetime import datetime -from functools import lru_cache +from threading import Lock from typing import TYPE_CHECKING +from cachetools import TTLCache, cachedmethod from pydantic import RedisDsn from generalresearch.managers.base import Permission @@ -26,7 +28,6 @@ from generalresearch.models.custom_types import UUIDStr from generalresearch.utils.copying_cache import deepcopy_return if TYPE_CHECKING: - from generalresearch.managers.thl.userhealth import AuditLogManager from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User @@ -41,9 +42,9 @@ auditlog = logging.getLogger("auditlog") class UserManager: def __init__( self, - redis: RedisDsn | None = None, - pg_config: PostgresConfig | None = None, - pg_config_rr: PostgresConfig | None = None, + redis: RedisDsn, + pg_config: PostgresConfig, + pg_config_rr: PostgresConfig, sql_permissions: Collection[Permission] | None = None, cache_prefix: str | None = None, redis_timeout: float | None = None, @@ -53,9 +54,9 @@ class UserManager: sql_permissions = [] if pg_config is not None: - assert ( - pg_config_rr is not None - ), "you should pass RR credentials also for fast lookups" + assert pg_config_rr is not None, ( + "you should pass RR credentials also for fast lookups" + ) assert Permission.DELETE not in sql_permissions, "delete not allowed" if Permission.UPDATE in sql_permissions or Permission.CREATE in sql_permissions: @@ -86,11 +87,40 @@ class UserManager: self.product_manager = ProductManager( pg_config=pg_config, permissions=[Permission.READ] ) + self.get_user_cache = TTLCache(maxsize=10000, ttl=30) + self.get_user_cache_lock = Lock() def set_last_seen(self, user: User) -> None: assert Permission.UPDATE in self.sql_permissions, "permission error" return self.mysql_user_manager._set_last_seen(user) + def change_product_user_id(self, *, user: User, new_product_user_id: str) -> User: + """Change a user's supplier-provided ID and refresh user lookup caches. + + This does not rewrite historical or derived data that copied the old ID, + such as leaderboards and activity counters. + """ + assert Permission.UPDATE in self.sql_permissions, "permission error" + assert self.mysql_user_manager is not None + assert user.product_id is not None + assert user.product_user_id is not None + + if new_product_user_id == user.product_user_id: + return user + if not User.is_valid_ubp( + product_id=user.product_id, product_user_id=new_product_user_id + ): + raise ValueError("invalid product_id/product_user_id") + + updated_user = self.mysql_user_manager._change_product_user_id( + user=user, new_product_user_id=new_product_user_id + ) + + self.cache_clear() + if self.redis_user_manager: + self.redis_user_manager.clear_user(user) + return updated_user + def audit_log( self, alm: AuditLogManager, @@ -102,6 +132,7 @@ class UserManager: ) -> AuditLog: from generalresearch.models.thl.userhealth import AuditLogLevel + assert user.user_id is not None return alm.create( user_id=user.user_id, level=AuditLogLevel(level), @@ -110,14 +141,17 @@ class UserManager: event_value=event_value, ) - def cache_clear(self): - # Generally this is used in testing. This clears the .get_user's lru_cache. - # There is no way of clearing only a specific key from the cache. + def cache_clear(self) -> None: + # Generally this is used in testing. This clears get_user's TTL cache. # It does not clear any redis caches; that has to be done separately. - self.get_user.__wrapped__.cache_clear() + with self.get_user_cache_lock: + self.get_user_cache.clear() @deepcopy_return - @lru_cache(maxsize=10000) + @cachedmethod( + operator.attrgetter("get_user_cache"), + lock=operator.attrgetter("get_user_cache_lock"), + ) def get_user( self, *, @@ -128,20 +162,20 @@ class UserManager: ) -> User: """ Retrieve User from (product_id & product_user_id) or (user_id), or (uuid). - Looks up in lru_cache, then (redis, memcached), then mysql. + Looks up in the 30-second TTL cache, then (redis, memcached), then mysql. Raises UserDoesntExistError if user is not found. (the * makes all arguments keyword-only arguments) """ - assert ( - (product_id and product_user_id) or user_id or user_uuid - ), "Must pass either (product_id, product_user_id), or user_id, or uuid" + assert (product_id and product_user_id) or user_id or user_uuid, ( + "Must pass either (product_id, product_user_id), or user_id, or uuid" + ) if product_id or product_user_id: - assert ( - product_id and product_user_id - ), "Must pass both product_id and product_user_id" - assert ( - sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1 - ), "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid" + assert product_id and product_user_id, ( + "Must pass both product_id and product_user_id" + ) + assert sum(map(bool, [product_id or product_id, user_id, user_uuid])) == 1, ( + "Must pass only 1 of (product_id, product_user_id), or user_id, or uuid" + ) user = self.get_user_inmemory_cache( product_id=product_id, product_user_id=product_user_id, @@ -242,19 +276,19 @@ class UserManager: """ Given a bp_user_id and a product_id, get or create a User """ + from generalresearch.models.thl.user import User + assert Permission.CREATE in self.sql_permissions assert self.mysql_user_manager is not None - assert ( - self.redis_user_manager is not None - ), "need at least redis to synchronize user creation" + assert self.redis_user_manager is not None, ( + "need at least redis to synchronize user creation" + ) - assert ( - self.user_manager_limiter is not None - ), "Need user_manager_limiter to get_or_create_user" + assert self.user_manager_limiter is not None, ( + "Need user_manager_limiter to get_or_create_user" + ) # Attempt to create common_struct solely for validation purposes - from generalresearch.models.thl.user import User - if not User.is_valid_ubp( product_id=product_id, product_user_id=product_user_id ): @@ -277,9 +311,9 @@ class UserManager: created: datetime | None = None, ) -> User: - assert ( - self.user_manager_limiter is not None - ), "Need user_manager_limiter to create_user" + assert self.user_manager_limiter is not None, ( + "Need user_manager_limiter to create_user" + ) assert product_id or product, "Needs a product_id or a Product instance" if product is None: @@ -327,8 +361,7 @@ class UserManager: # If we change something about a user, we should update the in-memory caches self.set_user_inmemory_cache(user) - # There is no way to clear a single key from the lru_cache... - # https://bugs.python.org/issue28178 + # Clear local cached copies so the blocked state is returned immediately. self.cache_clear() return True @@ -361,9 +394,9 @@ class UserManager: user_ids: Collection[int] | None = None, user_uuids: Collection[str] | None = None, ) -> list[User]: - assert (user_ids or user_uuids) and not ( - user_ids and user_uuids - ), "Must pass ONE of user_ids, user_uuids" + assert (user_ids or user_uuids) and not (user_ids and user_uuids), ( + "Must pass ONE of user_ids, user_uuids" + ) assert self.mysql_user_manager_rr return self.mysql_user_manager_rr.fetch( user_ids=user_ids, user_uuids=user_uuids diff --git a/generalresearch/managers/thl/user_manager/user_metadata_manager.py b/generalresearch/managers/thl/user_manager/user_metadata_manager.py index ffa8b44..91ac72e 100644 --- a/generalresearch/managers/thl/user_manager/user_metadata_manager.py +++ b/generalresearch/managers/thl/user_manager/user_metadata_manager.py @@ -7,6 +7,24 @@ from generalresearch.models.thl.user_profile import UserMetadata class UserMetadataManager(PostgresManager): + def filter_by_bpuids( + self, product_id: str, product_user_ids: Collection[str] + ) -> dict[str, UserMetadata]: + query = """ + SELECT email_address, email_sha256, email_sha1, email_md5, display_name, + product_user_id, thl_user.id as user_id + FROM thl_usermetadata + RIGHT OUTER JOIN thl_user on thl_usermetadata.user_id = thl_user.id + WHERE product_id = %(product_id)s + AND product_user_id = ANY(%(product_user_ids)s); + """ + res = self.pg_config.execute_sql_query( + query, + {"product_id": product_id, "product_user_ids": list(product_user_ids)}, + ) + + return {x["product_user_id"]: UserMetadata.from_db(**x) for x in res} + def filter( self, user_ids: Collection[int] | None = None, @@ -22,9 +40,9 @@ class UserMetadataManager(PostgresManager): email_sha1s, email_md5s, ]: - assert arg is None or isinstance( - arg, (set, list) - ), "must pass a collection of objects" + assert arg is None or isinstance(arg, (set, list)), ( + "must pass a collection of objects" + ) filters = [] params = {} @@ -48,7 +66,7 @@ class UserMetadataManager(PostgresManager): filter_str = "WHERE " + " AND ".join(filters) if filters else "" res = self.pg_config.execute_sql_query( f""" - SELECT user_id, email_address, email_sha256, email_sha1, email_md5 + SELECT user_id, email_address, email_sha256, email_sha1, email_md5, display_name FROM thl_usermetadata {filter_str} """, @@ -126,8 +144,12 @@ class UserMetadataManager(PostgresManager): c.execute( """ UPDATE thl_usermetadata - SET email_address = %(email_address)s, email_sha256 = %(email_sha256)s, - email_sha1 = %(email_sha1)s, email_md5 = %(email_md5)s + SET + email_address = %(email_address)s, + email_sha256 = %(email_sha256)s, + email_sha1 = %(email_sha1)s, + email_md5 = %(email_md5)s, + display_name = %(display_name)s WHERE user_id = %(user_id)s; """, params=user_metadata.to_db(), @@ -139,10 +161,14 @@ class UserMetadataManager(PostgresManager): def _create(self, user_metadata: UserMetadata) -> int: return self.pg_config.execute_write( query=""" - INSERT INTO thl_usermetadata - (user_id, email_address, email_sha256, email_sha1, email_md5) - VALUES (%(user_id)s, %(email_address)s, %(email_sha256)s, - %(email_sha1)s, %(email_md5)s); + INSERT INTO thl_usermetadata ( + user_id, email_address, email_sha256, + email_sha1, email_md5, display_name + ) + VALUES ( + %(user_id)s, %(email_address)s, %(email_sha256)s, + %(email_sha1)s, %(email_md5)s, %(display_name)s + ); """, params=user_metadata.to_db(), ) diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 4e90eb5..1a36fa9 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -55,6 +55,7 @@ class ContestEntry(BaseModel): ) # user_id used internally, for DB joins/index + # todo: this should be a UserRef user: User = Field(exclude=True) @model_validator(mode="before") @@ -65,9 +66,9 @@ class ContestEntry(BaseModel): entry_type = data.get("entry_type") if entry_type == ContestEntryType.COUNT: - assert isinstance(amount, int) and not isinstance( - amount, USDCent - ), "amount must be int in ContestEntryType.COUNT" + assert isinstance(amount, int) and not isinstance(amount, USDCent), ( + "amount must be int in ContestEntryType.COUNT" + ) elif entry_type == ContestEntryType.CASH: # This may be coming from the DB, in which case it is an int. diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py index 6e49af2..df60a7f 100644 --- a/generalresearch/models/thl/leaderboard.py +++ b/generalresearch/models/thl/leaderboard.py @@ -67,6 +67,10 @@ class LeaderboardRow(BaseModel): examples=[7], ) + display_name: str | None = Field( + description="Optional public name chosen by the user", default=None + ) + def censor(self): censor_idx = math.ceil(len(self.bpuid) / 2) self.bpuid = self.bpuid[:censor_idx] + ("*" * len(self.bpuid[censor_idx:])) diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index 45dcb40..9df7b23 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -22,12 +22,16 @@ from generalresearch.models.thl.user_streak import UserStreak class UserMetadata(BaseModel): model_config = ConfigDict(extra="forbid", validate_assignment=True) - user_id: SkipJsonSchema[PositiveInt | None] = Field( - exclude=True, default=None, lt=MAX_INT32 - ) + user_id: SkipJsonSchema[PositiveInt] = Field(exclude=True, lt=MAX_INT32) email_address: EmailStr | None = Field(default=None, examples=["contact@mail.com"]) + display_name: str | None = Field( + default=None, + max_length=255, + description="A public name chosen by the user. Can be used in leaderboards or event stream.", + ) + @computed_field def email_md5( self, @@ -85,9 +89,15 @@ class UserMetadata(BaseModel): return res @classmethod - def from_db(cls, user_id, email_address, **kwargs) -> Self: + def from_db(cls, user_id, email_address, display_name, **kwargs) -> Self: # If the hashes are passed, just validate that they match - obj = cls.model_validate({"user_id": user_id, "email_address": email_address}) + obj = cls.model_validate( + { + "user_id": user_id, + "email_address": email_address, + "display_name": display_name, + } + ) if kwargs.get("email_md5") is not None: assert obj.email_md5 == kwargs["email_md5"], "email_md5 mismatch" diff --git a/generalresearch/thl_django/common/models.py b/generalresearch/thl_django/common/models.py index a41d2eb..ffd9662 100644 --- a/generalresearch/thl_django/common/models.py +++ b/generalresearch/thl_django/common/models.py @@ -374,6 +374,11 @@ class THLUserMetadata(models.Model): email_sha1 = models.CharField(max_length=40, null=True) email_md5 = models.CharField(max_length=32, null=True) + # Not unique within a BP, or anything like that. A user + # can set this to whatever they like. No index + # as we will not ever look up a user by their name. + display_name = models.CharField(max_length=255, null=True) + class Meta: db_table = "thl_usermetadata" indexes = [ diff --git a/generalresearch/thl_django/migrations/0010_thlusermetadata_display_name.py b/generalresearch/thl_django/migrations/0010_thlusermetadata_display_name.py new file mode 100644 index 0000000..1524e3b --- /dev/null +++ b/generalresearch/thl_django/migrations/0010_thlusermetadata_display_name.py @@ -0,0 +1,18 @@ +# Generated by Django 6.1 on 2026-08-31 18:05 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ('thl_django', '0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more'), + ] + + operations = [ + migrations.AddField( + model_name='thlusermetadata', + name='display_name', + field=models.CharField(max_length=255, null=True), + ), + ] diff --git a/pyproject.toml b/pyproject.toml index 13fa584..5719a43 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,10 +4,10 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.0" +version = "3.4.5" description = "Python Utilities for General Research" readme = "README.md" -requires-python = ">=3.8" +requires-python = ">=3.10" dependencies = [ "Faker", "PyMySQL", |
