diff options
| author | Greg Stupp | 2026-08-31 17:47:45 +0000 |
|---|---|---|
| committer | Greg Stupp | 2026-08-31 17:47:45 +0000 |
| commit | b80ad1e82b676889c0db648704f4dc6acadd6ab0 (patch) | |
| tree | c6cd357b6295909d2297bd79449f515b2e480f6b | |
| parent | 97eeea03793cde962c38243e5c0f37ae9e5dfcbe (diff) | |
| parent | 55e77052a8f9560f64ccef07c1a9c46b97ce9261 (diff) | |
| download | generalresearch-b80ad1e82b676889c0db648704f4dc6acadd6ab0.tar.gz generalresearch-b80ad1e82b676889c0db648704f4dc6acadd6ab0.zip | |
Merges pull request #2
Dev greg
9 files changed, 213 insertions, 80 deletions
diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 0bf0312..27d6a89 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,45 +87,61 @@ 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] 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)) @@ -133,15 +151,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 +192,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 +219,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/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 ) 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/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 e96266a..13e8af1 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,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, @@ -86,9 +90,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 = [ |
