From a1b6e9fc5d1c8615a16e5ffb11caac7ef19765f5 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 24 Aug 2026 16:45:26 -0600 Subject: add TaskCalculationType 'starts' --- generalresearch/models/__init__.py | 1 + pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/generalresearch/models/__init__.py b/generalresearch/models/__init__.py index 9a2bb9c..560e5d9 100644 --- a/generalresearch/models/__init__.py +++ b/generalresearch/models/__init__.py @@ -92,6 +92,7 @@ class TaskCalculationType(str, Enum): "survey start": cls.STARTS, "survey starts": cls.STARTS, "start": cls.STARTS, + "starts": cls.STARTS, "prescreens": cls.STARTS, "prescreen": cls.STARTS, }[v.lower()] diff --git a/pyproject.toml b/pyproject.toml index 183e271..0fe8772 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.0" +version = "3.4.1" description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.8" -- cgit v1.2.3 From d63d47c69406614a6cf5593bc30e6ae9f6529962 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 24 Aug 2026 17:06:41 -0600 Subject: return_filepath_adapter -> return _filepath_adapter --- generalresearch/incite/base.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index a8088ac..2b45d24 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -53,7 +53,7 @@ from generalresearch.incite.schemas import ( from generalresearch.models.custom_types import AwareDatetimeISO if TYPE_CHECKING: - from generalresearch.incite.collections import DFCollection + from generalresearch.incite.collections import DFCollection, DFCollectionItem from generalresearch.incite.collections.thl_marketplaces import ( DFCollectionType, ) @@ -712,7 +712,7 @@ class CollectionItemBase(BaseModel): @property def path(self) -> FilePath: - return_filepath_adapter.validate_python( + return _filepath_adapter.validate_python( os.path.join(self._collection.archive_path, self.filename) ) -- cgit v1.2.3 From f2f6370a3031b7e0f0376b6628ee0f55ec09d990 Mon Sep 17 00:00:00 2001 From: stuppie Date: Wed, 26 Aug 2026 18:06:48 -0600 Subject: CollectionItemBase.path: 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... Also, its not even a path; it is a dir --- generalresearch/incite/base.py | 82 ++++++++++++++---------------------------- 1 file changed, 27 insertions(+), 55 deletions(-) diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 2b45d24..795aad9 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -104,9 +104,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: @@ -447,26 +447,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() @@ -482,7 +462,7 @@ class CollectionBase(BaseModel): item.cleanup_partials() def clear_tmp_archives(self) -> None: - regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.I) + regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.IGNORECASE) for fn in os.listdir(self.archive_path): if regex.search(fn): @@ -524,7 +504,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) @@ -549,13 +528,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 (Exception,): + 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 @@ -607,7 +586,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: @@ -672,9 +650,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 --- @@ -711,27 +689,24 @@ class CollectionItemBase(BaseModel): return f"{self.filename}.empty" @property - def path(self) -> FilePath: - return _filepath_adapter.validate_python( - os.path.join(self._collection.archive_path, self.filename) - ) + def path(self) -> Path: + # 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) -> FilePath: - return FilePath( - os.path.join(self._collection.archive_path, self.partial_filename) - ) + def partial_path(self) -> Path: + return Path(self._collection.archive_path, self.partial_filename) @property - def empty_path(self) -> FilePath: - return FilePath( - os.path.join(self._collection.archive_path, self.empty_filename) - ) + def empty_path(self) -> Path: + 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 @@ -766,7 +741,6 @@ class CollectionItemBase(BaseModel): builds = [] for fn in os.listdir(coll.archive_path): 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) @@ -788,10 +762,8 @@ class CollectionItemBase(BaseModel): # up as always returning the same tmp filename return f"{self.filename}.{uuid4().hex}" - def tmp_path(self) -> FilePath: - return FilePath( - os.path.join(self._collection.archive_path, self.tmp_filename()) - ) + def tmp_path(self) -> Path: + return Path(self._collection.archive_path, self.tmp_filename()) # --- --- --- --- # If it has a partial, it isn't always going to be a partial. However, @@ -820,7 +792,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) @@ -842,9 +813,9 @@ class CollectionItemBase(BaseModel): return datetime.now(tz=timezone.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" @@ -979,6 +950,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()} -- cgit v1.2.3 From f2a6426642cc07991bb10fb3b09a2d1ee2dc8f8d Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 27 Aug 2026 10:55:09 -0600 Subject: bump --- pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 0fe8772..fd55201 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,10 +4,10 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.1" +version = "3.4.2" description = "Python Utilities for General Research" readme = "README.md" -requires-python = ">=3.8" +requires-python = ">=3.10" dependencies = [ "fastapi", "Faker", -- cgit v1.2.3 From f556f942b2ef1a921167f061585e9e7ec2511a78 Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 27 Aug 2026 14:37:12 -0600 Subject: fix ruff warning on user manager using cachedmethod. Change get_user cache to a 30-second ttl cache --- .../thl/user_manager/memcached_user_manager.py | 49 ---------------------- .../thl/user_manager/mysql_user_manager.py | 21 ++++++---- .../managers/thl/user_manager/user_manager.py | 24 +++++++---- 3 files changed, 29 insertions(+), 65 deletions(-) delete mode 100644 generalresearch/managers/thl/user_manager/memcached_user_manager.py 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 dbed5de..22abaa8 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -1,11 +1,13 @@ from __future__ import annotations import logging +import operator from collections.abc import Collection from datetime import datetime, timezone -from functools import lru_cache +from threading import Lock from uuid import uuid4 +from cachetools import LRUCache, cachedmethod import psycopg from psycopg import sql @@ -22,6 +24,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() @@ -143,7 +147,7 @@ class MysqlUserManager: with conn.cursor() as c: c.execute(query=query, params=params) user_id = c.fetchone()["id"] - except psycopg.IntegrityError as e: + except psycopg.IntegrityError: # Two machines/processes are trying to create this same (product_id, product_user_id) # at the same time. There's a unique index, so mysql will not let two be created. # The 2nd should get an IntegrityError, meaning this already exists, and we can just query it. @@ -160,7 +164,7 @@ class MysqlUserManager: else: # We specifically queried the NON read-replica, and we got an IntegrityError, so # something else must be wrong... - raise e + raise else: user = User( user_id=user_id, @@ -173,8 +177,11 @@ class MysqlUserManager: return user - @lru_cache(maxsize=5000) - 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 @@ -254,10 +261,10 @@ class MysqlUserManager: user_ids and user_uuids ), "Must pass ONE of user_ids, user_uuids" if user_ids: - assert len(user_ids) <= 500, "limit 500 user_ids" assert isinstance( user_ids, (list, set) ), "must pass a collection of user_ids" + assert len(user_ids) <= 500, "limit 500 user_ids" res = self.pg_config.execute_sql_query( query=""" @@ -270,10 +277,10 @@ class MysqlUserManager: params={"user_ids": user_ids}, ) else: - assert len(user_uuids) <= 500, "limit 500 user_uuids" assert isinstance( user_uuids, (list, set) ), "must pass a collection of user_uuids" + assert len(user_uuids) <= 500, "limit 500 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/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 3794020..5da3370 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -1,10 +1,12 @@ 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 cachetools import TTLCache, cachedmethod from pydantic import RedisDsn from generalresearch.managers.base import Permission @@ -78,6 +80,8 @@ 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" @@ -105,14 +109,17 @@ class UserManager: return None - 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, *, @@ -123,7 +130,7 @@ 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) """ @@ -316,8 +323,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 -- cgit v1.2.3 From 57bf54e188b8a18f072c458ba245119998439127 Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 27 Aug 2026 15:05:31 -0600 Subject: add display_name to UserMetadata --- .../thl/user_manager/user_metadata_manager.py | 28 ++++++++++++++-------- generalresearch/models/thl/user_profile.py | 20 ++++++++++------ generalresearch/thl_django/common/models.py | 5 ++++ 3 files changed, 36 insertions(+), 17 deletions(-) diff --git a/generalresearch/managers/thl/user_manager/user_metadata_manager.py b/generalresearch/managers/thl/user_manager/user_metadata_manager.py index ffa8b44..4f394f1 100644 --- a/generalresearch/managers/thl/user_manager/user_metadata_manager.py +++ b/generalresearch/managers/thl/user_manager/user_metadata_manager.py @@ -22,9 +22,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 +48,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 +126,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 +143,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/user_profile.py b/generalresearch/models/thl/user_profile.py index e96266a..064036f 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -1,7 +1,7 @@ from __future__ import annotations import hashlib -from typing import Any +from typing import Annotated, Any from pydantic import ( BaseModel, @@ -12,7 +12,7 @@ from pydantic import ( computed_field, ) from pydantic.json_schema import SkipJsonSchema -from typing_extensions import Annotated, Self +from typing_extensions import Self from generalresearch.models import MAX_INT32, Source from generalresearch.models.custom_types import UUIDStr @@ -23,12 +23,12 @@ 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) + @computed_field def email_md5( self, @@ -86,9 +86,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 = [ -- cgit v1.2.3 From c5aa734cad3d4a8a1f01d113fff05198218e7879 Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 27 Aug 2026 17:46:17 -0600 Subject: Leaderboard id & name optional; set by validators --- generalresearch/models/thl/leaderboard.py | 17 +++++++++++------ pyproject.toml | 2 +- 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py index 399a906..bc55034 100644 --- a/generalresearch/models/thl/leaderboard.py +++ b/generalresearch/models/thl/leaderboard.py @@ -6,6 +6,7 @@ from datetime import datetime, timedelta, timezone from enum import Enum from typing import Literal from uuid import UUID, uuid3 +from zoneinfo import ZoneInfo import pandas as pd from pydantic import ( @@ -17,7 +18,6 @@ from pydantic import ( field_validator, model_validator, ) -from zoneinfo import ZoneInfo from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.legacy.api_status import StatusResponse @@ -78,14 +78,19 @@ class Leaderboard(BaseModel): weekly, and monthly leaderboard. """ - id: UUIDStr = Field( + # Note: id and name get auto-generated by the model_validators, but the fields need + # to be optional with a default or the model can't be inited. + # todo: these should really be computed_fields instead + id: UUIDStr | None = Field( description="Unique ID for this leaderboard", examples=["845b0074ad533df580ebb9c80cc3bce1"], + default=None, ) - name: str = Field( + name: str | None = Field( description="Descriptive name for the leaderboard based on the board_code", examples=["Number of Completes"], + default=None, ) board_code: LeaderboardCode = Field( @@ -266,9 +271,9 @@ class Leaderboard(BaseModel): .to_pydatetime() .replace(tzinfo=self.timezone) ) - assert ( - period_start_local == self.period_start_local - ), f"invalid period_start_local {self.period_start_local}. The period starts at {period_start_local}" + assert period_start_local == self.period_start_local, ( + f"invalid period_start_local {self.period_start_local}. The period starts at {period_start_local}" + ) if self.period_end_local is not None: assert self.period_end_local == period_end_local, "invalid period" else: diff --git a/pyproject.toml b/pyproject.toml index fd55201..2d1d246 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.2" +version = "3.4.3" description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.10" -- cgit v1.2.3 From ec850b9ddd5cf520646f0f8c4a5f0a996b78ac94 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 28 Aug 2026 10:56:58 -0600 Subject: fix broken ruff imports. fastapi isnt a dependency --- generalresearch/models/thl/survey/penalty.py | 5 ++--- pyproject.toml | 4 ++-- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index e9515d4..0f544b8 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -2,10 +2,9 @@ from __future__ import annotations import abc from datetime import datetime, timezone -from typing import Literal +from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter -from typing_extensions import Annotated from generalresearch.models import Source from generalresearch.models.custom_types import ( @@ -59,7 +58,7 @@ class TeamSurveyPenalty(SurveyPenalty): Penalty = Annotated[ - Union[BPSurveyPenalty, TeamSurveyPenalty], + BPSurveyPenalty | TeamSurveyPenalty, Field(discriminator="kind"), ] PenaltyListAdapter = TypeAdapter(list[Penalty]) diff --git a/pyproject.toml b/pyproject.toml index 2d1d246..e855d16 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,12 +4,12 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.3" +version = "3.4.4" description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.10" dependencies = [ - "fastapi", +# "fastapi", "Faker", "PyMySQL", "psycopg", -- cgit v1.2.3 From 97eeea03793cde962c38243e5c0f37ae9e5dfcbe Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 28 Aug 2026 11:28:31 -0600 Subject: fix more typing imports --- generalresearch/incite/base.py | 5 +---- generalresearch/models/cint/survey.py | 14 +++++++------- generalresearch/models/dynata/survey.py | 6 +++--- generalresearch/models/dynata/task_collection.py | 2 +- generalresearch/models/prodege/task_collection.py | 6 +++--- generalresearch/models/thl/contest/__init__.py | 7 ++++--- generalresearch/models/thl/contest/contest_entry.py | 1 + generalresearch/utils/copying_cache.py | 2 +- 8 files changed, 21 insertions(+), 22 deletions(-) diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 795aad9..54d1565 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -7,7 +7,7 @@ import re import shutil import subprocess import warnings -from concurrent.futures import Future +from collections.abc import Callable, Sequence from datetime import datetime, timedelta, timezone from os import R_OK, access, listdir from os.path import isdir @@ -17,12 +17,9 @@ from sys import platform from typing import ( TYPE_CHECKING, Any, - Callable, - Sequence, ) from uuid import uuid4 -import dask import dask.dataframe as dd import pandas as pd import pyarrow.parquet as pq diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index 56384e3..21bc21a 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -4,7 +4,7 @@ import json import logging from datetime import datetime, timezone from decimal import Decimal -from typing import Any, Literal, Type +from typing import Annotated, Any, Literal from more_itertools import flatten from pydantic import ( @@ -15,7 +15,7 @@ from pydantic import ( computed_field, model_validator, ) -from typing_extensions import Annotated, Self +from typing_extensions import Self from generalresearch.locales import Localelator from generalresearch.models import Source, TaskCalculationType @@ -73,9 +73,9 @@ class CintQuota(BaseModel): @model_validator(mode="after") def validate_condition_len(self) -> Self: if self.quota_type == "total": - assert ( - self.condition_hashes is None - ), "total quota should not have conditions" + assert self.condition_hashes is None, ( + "total quota should not have conditions" + ) elif self.quota_type == "client": assert len(self.condition_hashes) > 0, "quota must have conditions" return self @@ -291,7 +291,7 @@ class CintSurvey(MarketplaceTask): return data @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return CintCondition @property @@ -417,7 +417,7 @@ class CintSurvey(MarketplaceTask): return d @classmethod - def from_mysql(cls, d: Dict[str, Any]) -> Self: + def from_mysql(cls, d: dict[str, Any]) -> Self: d["created_at"] = d["created_at"].replace(tzinfo=timezone.utc) d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc) d["qualifications"] = json.loads(d["qualifications"]) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 0e1b3e5..e88491d 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -5,7 +5,7 @@ import logging from datetime import timezone from decimal import Decimal from functools import cached_property -from typing import Any, Literal, Type +from typing import Any, Literal from more_itertools import flatten from pydantic import ( @@ -500,7 +500,7 @@ class DynataSurvey(MarketplaceTask): return res @property - def condition_model(self) -> Type[MarketplaceCondition]: + def condition_model(self) -> type[MarketplaceCondition]: return DynataCondition @property @@ -551,7 +551,7 @@ class DynataSurvey(MarketplaceTask): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> Self: d["created"] = d["created"].replace(tzinfo=timezone.utc) d["last_updated"] = d["last_updated"].replace(tzinfo=timezone.utc) d["filters"] = json.loads(d["filters"]) diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 71cf3db..2b82bfd 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -54,7 +54,7 @@ DynataTaskCollectionSchema = DataFrameSchema( class DynataTaskCollection(TaskCollection): - items: List[DynataSurvey] + items: list[DynataSurvey] _schema = DynataTaskCollectionSchema def to_row(self, s: DynataSurvey) -> dict[str, Any]: diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 19e594f..4544050 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -30,8 +30,8 @@ ProdegeTaskCollectionSchema = DataFrameSchema( "conversion_rate": Column(float, Check.between(0, 1), nullable=True), "created": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), "updated": Column(dtype=pd.DatetimeTZDtype(tz="UTC")), - "used_question_ids": Column(List[str]), - "all_hashes": Column(List[str]), # set >> list for column support + "used_question_ids": Column(list[str]), + "all_hashes": Column(list[str]), # set >> list for column support "is_recontact": Column(bool), # Not including here: entrance_url, max_clicks_settings, past_participation, include_psids, exclude_psids, # quotas, source, conditions diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 363c8c0..c02acbe 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import datetime, timezone +from typing import Any from uuid import uuid4 from pydantic import ( @@ -86,9 +87,9 @@ class ContestPrize(BaseModel): @model_validator(mode="after") def validate_cash_value(self) -> Self: if self.kind == ContestPrizeKind.CASH: - assert ( - self.estimated_cash_value == self.cash_amount - ), "if kind is CASH, cash_amount must equal estimated_cash_value" + assert self.estimated_cash_value == self.cash_amount, ( + "if kind is CASH, cash_amount must equal estimated_cash_value" + ) return self diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index cddae14..2a9ecde 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import datetime, timezone +from typing import Any from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/utils/copying_cache.py b/generalresearch/utils/copying_cache.py index ea13f69..a1cb37c 100644 --- a/generalresearch/utils/copying_cache.py +++ b/generalresearch/utils/copying_cache.py @@ -1,6 +1,6 @@ +from collections.abc import Callable from copy import deepcopy from functools import wraps -from typing import Callable def deepcopy_return(fn: Callable) -> Callable: -- cgit v1.2.3 From 7a2d647577fd62db042a1ec72a8a0696e7552ad4 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 28 Aug 2026 14:25:05 -0600 Subject: Get leaderboard to fill in users' display names --- generalresearch/managers/leaderboard/manager.py | 73 +++++++++++++--------- .../thl/user_manager/user_metadata_manager.py | 18 ++++++ generalresearch/models/thl/leaderboard.py | 4 ++ generalresearch/models/thl/user.py | 6 +- generalresearch/models/thl/user_profile.py | 6 +- 5 files changed, 74 insertions(+), 33 deletions(-) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 0bf0312..c44a461 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime, timedelta, timezone 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=timezone.utc).astimezone( self.timezone @@ -85,28 +87,37 @@ 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, + self, user_metadata_manager: UserMetadataManager, limit: int | 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) ] + 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, + user_metadata_manager: UserMetadataManager, + bp_user_id: str, + limit: int | None = 5, ) -> 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 @@ -123,7 +134,9 @@ class LeaderboardManager: user_idx = user_indices[0][0] user_row = user_indices[0][1] if user_row.rank == max([row.rank for row in rows]): - user_idx = [i for i, row in enumerate(rows) if row.rank == user_row.rank][0] + user_idx = next( + i for i, row in enumerate(rows) if row.rank == user_row.rank + ) start: int = max(user_idx - limit, 0) end: int = min(user_idx + limit + 1, len(rows)) @@ -131,17 +144,20 @@ class LeaderboardManager: def get_leaderboard( self, + user_metadata_manager: UserMetadataManager, limit: int | None = None, bp_user_id: str | None = None, ) -> Leaderboard: 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 +180,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 +207,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/user_metadata_manager.py b/generalresearch/managers/thl/user_manager/user_metadata_manager.py index 4f394f1..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, diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py index bc55034..d73c7cf 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.py b/generalresearch/models/thl/user.py index 355e331..08e59a0 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -4,7 +4,7 @@ import json import logging import re from datetime import datetime, timezone -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Annotated from uuid import UUID, uuid4 from pydantic import ( @@ -19,7 +19,7 @@ from pydantic import ( model_validator, ) from sentry_sdk import set_tag, set_user -from typing_extensions import Annotated, Self +from typing_extensions import Self from generalresearch.models import MAX_INT32 from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr @@ -35,8 +35,6 @@ if TYPE_CHECKING: ) from generalresearch.managers.thl.userhealth import AuditLogManager - # from generalresearch.managers.thl.userhealth import UserIpHistoryManager - logger = logging.getLogger() BPUID_ALLOWED = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[\]^_{|}~" diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index 064036f..13e8af1 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -27,7 +27,11 @@ class UserMetadata(BaseModel): email_address: EmailStr | None = Field(default=None, examples=["contact@mail.com"]) - display_name: str | None = Field(default=None, max_length=255) + 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( -- cgit v1.2.3 From f7b187e553bea8f8a393d5c67c607fa667de1925 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 28 Aug 2026 15:58:42 -0600 Subject: add user manager change_product_user_id --- .../thl/user_manager/mysql_user_manager.py | 36 ++++++++++ .../thl/user_manager/redis_user_manager.py | 1 - .../managers/thl/user_manager/user_manager.py | 84 ++++++++++++++-------- 3 files changed, 91 insertions(+), 30 deletions(-) diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index 22abaa8..109b601 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -40,6 +40,42 @@ 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: + with 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, *, 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 5da3370..7f24029 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -35,9 +35,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, @@ -47,9 +47,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: @@ -87,6 +87,33 @@ class UserManager: 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, user: User, @@ -99,6 +126,7 @@ class UserManager: from generalresearch.models.thl.userhealth import AuditLogLevel alm = AuditLogManager(pg_config=self.mysql_user_manager.pg_config) + assert user.user_id is not None alm.create( user_id=user.user_id, level=AuditLogLevel(level), @@ -107,8 +135,6 @@ class UserManager: event_value=event_value, ) - return None - 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. @@ -134,16 +160,16 @@ class UserManager: 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, @@ -245,13 +271,13 @@ class UserManager: """ 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 if not User.is_valid_ubp( product_id=product_id, product_user_id=product_user_id @@ -275,9 +301,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: @@ -354,9 +380,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" + ) return self.mysql_user_manager_rr.fetch( user_ids=user_ids, user_uuids=user_uuids ) -- cgit v1.2.3 From 55e77052a8f9560f64ccef07c1a9c46b97ce9261 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 31 Aug 2026 10:38:06 -0600 Subject: get leaderboard user_metadata_manager optional + comments --- generalresearch/managers/leaderboard/manager.py | 34 +++++++++++++++++-------- 1 file changed, 23 insertions(+), 11 deletions(-) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index c44a461..27d6a89 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -91,7 +91,9 @@ class LeaderboardManager: return cast(int, self.redis_client.zcard(self.key)) def get_leaderboard_rows( - self, user_metadata_manager: UserMetadataManager, limit: int | None = None + self, + limit: int | None = None, + user_metadata_manager: UserMetadataManager | None = None, ) -> list[LeaderboardRow]: limit = limit or 0 res = self.redis_client.zrange( @@ -106,29 +108,32 @@ class LeaderboardManager: LeaderboardRow(bpuid=bpuid, value=value, rank=rank) for bpuid, value, rank in s.itertuples(index=False, name=None) ] - 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 + 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, - user_metadata_manager: UserMetadataManager, 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] @@ -144,11 +149,18 @@ class LeaderboardManager: def get_leaderboard( self, - user_metadata_manager: UserMetadataManager, 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, -- cgit v1.2.3 From 625a4c09c2399da76e10796bcc2d79c06a9d1ba1 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 31 Aug 2026 12:06:48 -0600 Subject: migration --- .../migrations/0010_thlusermetadata_display_name.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) create mode 100644 generalresearch/thl_django/migrations/0010_thlusermetadata_display_name.py 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), + ), + ] -- cgit v1.2.3 From f6ee73468e27cfddd34b24f00ffd712bccd03347 Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 31 Aug 2026 12:08:30 -0600 Subject: bump --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index e855d16..0cc4227 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "generalresearch" -version = "3.4.4" +version = "3.4.5" description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.10" -- cgit v1.2.3