aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorstuppie2026-09-07 11:36:07 -0600
committerstuppie2026-09-07 11:36:07 -0600
commit242579a44855873d5e054e375440e9d3492cd682 (patch)
tree4c30256ad73a0f15a592e9fbf30a665b15ac0717
parent3338f74a94d0624bf894ebb35bd1bcfca268216e (diff)
parentf6ee73468e27cfddd34b24f00ffd712bccd03347 (diff)
downloadgeneralresearch-242579a44855873d5e054e375440e9d3492cd682.tar.gz
generalresearch-242579a44855873d5e054e375440e9d3492cd682.zip
Merge branch 'master' into dev
-rw-r--r--generalresearch/incite/base.py89
-rw-r--r--generalresearch/managers/leaderboard/manager.py83
-rw-r--r--generalresearch/managers/thl/user_manager/memcached_user_manager.py49
-rw-r--r--generalresearch/managers/thl/user_manager/mysql_user_manager.py89
-rw-r--r--generalresearch/managers/thl/user_manager/redis_user_manager.py1
-rw-r--r--generalresearch/managers/thl/user_manager/user_manager.py111
-rw-r--r--generalresearch/managers/thl/user_manager/user_metadata_manager.py46
-rw-r--r--generalresearch/models/thl/contest/contest_entry.py7
-rw-r--r--generalresearch/models/thl/leaderboard.py4
-rw-r--r--generalresearch/models/thl/user_profile.py20
-rw-r--r--generalresearch/thl_django/common/models.py5
-rw-r--r--generalresearch/thl_django/migrations/0010_thlusermetadata_display_name.py18
-rw-r--r--pyproject.toml4
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",