aboutsummaryrefslogtreecommitdiff
path: root/generalresearch/managers/thl/profiling
diff options
context:
space:
mode:
Diffstat (limited to 'generalresearch/managers/thl/profiling')
-rw-r--r--generalresearch/managers/thl/profiling/question.py22
-rw-r--r--generalresearch/managers/thl/profiling/schema.py2
-rw-r--r--generalresearch/managers/thl/profiling/uqa.py2
-rw-r--r--generalresearch/managers/thl/profiling/user_upk.py17
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