diff options
Diffstat (limited to 'generalresearch/managers/thl/profiling')
| -rw-r--r-- | generalresearch/managers/thl/profiling/question.py | 22 | ||||
| -rw-r--r-- | generalresearch/managers/thl/profiling/schema.py | 2 | ||||
| -rw-r--r-- | generalresearch/managers/thl/profiling/uqa.py | 2 | ||||
| -rw-r--r-- | generalresearch/managers/thl/profiling/user_upk.py | 17 |
4 files changed, 28 insertions, 15 deletions
diff --git a/generalresearch/managers/thl/profiling/question.py b/generalresearch/managers/thl/profiling/question.py index 1ad27ac..7b2a7ad 100644 --- a/generalresearch/managers/thl/profiling/question.py +++ b/generalresearch/managers/thl/profiling/question.py @@ -1,15 +1,15 @@ import random import threading -from typing import Collection, List, Tuple +from typing import Any, Collection, Dict, List, Tuple -from cachetools import cached, TTLCache +from cachetools import TTLCache, cached from pydantic import ValidationError from generalresearch.decorators import LOG from generalresearch.managers.base import PostgresManager from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, UPKImportance, + UpkQuestion, ) @@ -21,8 +21,8 @@ class QuestionManager(PostgresManager): FROM marketplace_question WHERE id = ANY(%(question_ids)s); """ - res = self.pg_config.execute_sql_query( - query, {"question_ids": list(question_ids)} + res: List[Dict[str, Any]] = self.pg_config.execute_sql_query( + query=query, params={"question_ids": list(question_ids)} ) for x in res: x["data"]["ext_question_id"] = x["property_code"] @@ -31,6 +31,7 @@ class QuestionManager(PostgresManager): "explanation_fragment_template" ] x["data"].pop("categories", None) + return [UpkQuestion.model_validate(x["data"]) for x in res] @cached( @@ -50,7 +51,7 @@ class QuestionManager(PostgresManager): AND property_code NOT LIKE 'g:%%' AND is_live """ - res = self.pg_config.execute_sql_query( + res: List[Dict[str, Any]] = self.pg_config.execute_sql_query( query=query, params={"country_iso": country_iso, "language_iso": language_iso}, ) @@ -94,9 +95,12 @@ class QuestionManager(PostgresManager): "country_iso": country_iso, "language_iso": language_iso, } - res = self.pg_config.execute_sql_query(query=query, params=params) + res: List[Dict[str, Any]] = self.pg_config.execute_sql_query( + query=query, params=params + ) assert len(res) == 1, f"expected 1, got {len(res)} results" x = res[0] + x["data"]["ext_question_id"] = x["property_code"] x["data"]["explanation_template"] = x["explanation_template"] x["data"]["explanation_fragment_template"] = x["explanation_fragment_template"] @@ -119,7 +123,9 @@ class QuestionManager(PostgresManager): WHERE {where_str} """ flat_params = [item for tup in lookup for item in tup] - res = self.pg_config.execute_sql_query(query, params=flat_params) + res: List[Dict[str, Any]] = self.pg_config.execute_sql_query( + query=query, params=flat_params + ) for x in res: x["data"]["ext_question_id"] = x["property_code"] x["data"]["explanation_template"] = x["explanation_template"] diff --git a/generalresearch/managers/thl/profiling/schema.py b/generalresearch/managers/thl/profiling/schema.py index 581270b..e209067 100644 --- a/generalresearch/managers/thl/profiling/schema.py +++ b/generalresearch/managers/thl/profiling/schema.py @@ -2,7 +2,7 @@ from threading import RLock from typing import List from uuid import UUID -from cachetools import cached, TTLCache +from cachetools import TTLCache, cached from generalresearch.managers.base import PostgresManager from generalresearch.models.thl.profiling.upk_property import ( diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py index 6800d32..1cab6c2 100644 --- a/generalresearch/managers/thl/profiling/uqa.py +++ b/generalresearch/managers/thl/profiling/uqa.py @@ -132,7 +132,7 @@ class UQAManager(PostgresManagerWithRedis): # 1) the cache expired and the user hasn't sent an answer recently # or 2) The user just sent an answer, so we'll make sure it gets put into the results # after this query runs. - query = f""" + query = """ WITH ranked AS ( SELECT uqa.*, diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index 4449afa..b3ea52e 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -1,10 +1,11 @@ import json from collections import defaultdict -from datetime import timedelta, datetime, timezone -from typing import Dict, Union, Set, List, Collection, Optional, Tuple +from datetime import datetime, timedelta, timezone +from typing import Any, Collection, Dict, List, Optional, Set, Tuple, Union from uuid import UUID from psycopg import Cursor +from pydantic import PositiveInt from generalresearch.managers.base import ( Permission, @@ -117,8 +118,9 @@ class UserUpkManager(PostgresManagerWithRedis): return [UpkQuestionAnswer.model_validate(x) for x in res] def get_user_upk_simple( - self, user_id, country_iso="us" + self, user_id: PositiveInt, country_iso: str = "us" ) -> Dict[str, Union[Set[str], str, float]]: + res = self.get_user_upk(user_id=user_id) res = [x for x in res if x.country_iso == country_iso] d: Dict[str, Union[Set[str], str, float]] = defaultdict(set) @@ -127,16 +129,19 @@ class UserUpkManager(PostgresManagerWithRedis): d[x.property_label] = x.value else: d[x.property_label].add(x.value) + return dict(d) def get_age_gender( - self, user_id, country_iso="us" + self, user_id: PositiveInt, country_iso: str = "us" ) -> Tuple[Optional[int], Optional[str]]: + # Returns an integer year for age, and {'male', 'female', 'other_gender'} d = self.get_user_upk_simple(user_id, country_iso) age = d.get("age_in_years") if age is not None: age = int(age) + gender = d.get("gender") return age, gender @@ -145,7 +150,9 @@ class UserUpkManager(PostgresManagerWithRedis): country_iso=country_iso ) - def populate_user_upk_from_dict(self, upk_ans_dict): + def populate_user_upk_from_dict( + self, upk_ans_dict: List[Dict[str, Any]] + ) -> List[UpkQuestionAnswer]: country_isos = {x["country_iso"] for x in upk_ans_dict} assert len(country_isos) == 1 |
