aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--generalresearch/__init__.py19
-rw-r--r--generalresearch/config.py6
-rw-r--r--generalresearch/currency.py26
-rw-r--r--generalresearch/grliq/managers/forensic_data.py84
-rw-r--r--generalresearch/grliq/managers/forensic_summary.py8
-rw-r--r--generalresearch/grliq/models/events.py7
-rw-r--r--generalresearch/grliq/models/forensic_data.py19
-rw-r--r--generalresearch/grliq/models/forensic_summary.py7
-rw-r--r--generalresearch/grliq/models/useragents.py4
-rw-r--r--generalresearch/healing_ppe.py2
-rw-r--r--generalresearch/incite/base.py28
-rw-r--r--generalresearch/incite/collections/__init__.py123
-rw-r--r--generalresearch/incite/exceptions.py19
-rw-r--r--generalresearch/incite/mergers/__init__.py26
-rw-r--r--generalresearch/incite/mergers/foundations/__init__.py18
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_session.py26
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_task_adjust.py9
-rw-r--r--generalresearch/incite/mergers/foundations/enriched_wall.py2
-rw-r--r--generalresearch/incite/mergers/ym_survey_wall.py4
-rw-r--r--generalresearch/incite/mergers/ym_wall_summary.py13
-rw-r--r--generalresearch/incite/schemas/admin_responses.py73
-rw-r--r--generalresearch/locales/__init__.py4
-rw-r--r--generalresearch/locales/setup_json.py124
-rw-r--r--generalresearch/logging.py5
-rw-r--r--generalresearch/managers/cint/survey.py2
-rw-r--r--generalresearch/managers/criteria.py20
-rw-r--r--generalresearch/managers/dynata/survey.py64
-rw-r--r--generalresearch/managers/events.py57
-rw-r--r--generalresearch/managers/gr/authentication.py6
-rw-r--r--generalresearch/managers/innovate/survey.py83
-rw-r--r--generalresearch/managers/morning/survey.py87
-rw-r--r--generalresearch/managers/network/mtr.py7
-rw-r--r--generalresearch/managers/network/nmap.py2
-rw-r--r--generalresearch/managers/network/rdns.py5
-rw-r--r--generalresearch/managers/precision/survey.py51
-rw-r--r--generalresearch/managers/prodege/survey.py51
-rw-r--r--generalresearch/managers/repdata/survey.py81
-rw-r--r--generalresearch/managers/sago/survey.py65
-rw-r--r--generalresearch/managers/spectrum/survey.py63
-rw-r--r--generalresearch/managers/thl/task_adjustment.py8
-rw-r--r--generalresearch/managers/thl/user_manager/mysql_user_manager.py2
-rw-r--r--generalresearch/managers/thl/userhealth.py4
-rw-r--r--generalresearch/managers/thl/wall.py2
-rw-r--r--generalresearch/models/gr/business.py5
-rw-r--r--generalresearch/models/legacy/bucket.py6
-rw-r--r--generalresearch/pg_helper.py11
-rw-r--r--generalresearch/sql_helper.py5
-rw-r--r--tests/incite/collections/test_df_collection_item_thl_web.py12
-rw-r--r--tests/incite/collections/test_df_collection_thl_marketplaces.py2
-rw-r--r--tests/managers/thl/test_contest/test_milestone.py2
50 files changed, 677 insertions, 682 deletions
diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py
index 3b2ec3d..f27385c 100644
--- a/generalresearch/__init__.py
+++ b/generalresearch/__init__.py
@@ -1,20 +1,27 @@
+import logging
import threading
import time
from collections.abc import Callable
from functools import wraps
-from typing import Any, Optional
+from typing import ParamSpec, TypeVar
from decorator import decorator
from wrapt import FunctionWrapper, ObjectProxy
+P = ParamSpec("P")
+R = TypeVar("R")
+
+ExceptionType = type[BaseException]
+ExceptionsArg = ExceptionType | tuple[ExceptionType, ...]
+
def retry(
- exceptions,
+ exceptions: ExceptionsArg,
tries: int = 4,
delay: float = 0.5,
backoff: int = 2,
- logger: Any | None = None,
-) -> Callable:
+ logger: logging.Logger | None = None,
+) -> Callable[[Callable[P, R]], Callable[P, R]]:
"""
https://www.calazan.com/retry-decorator-for-python-3/
Retry calling the decorated function using an exponential backoff.
@@ -29,10 +36,10 @@ def retry(
logger: Logger to use. If None, print.
"""
- def deco_retry(f):
+ def deco_retry(f: Callable[P, R]) -> Callable[P, R]:
@wraps(f)
- def f_retry(*args, **kwargs):
+ def f_retry(*args: P.args, **kwargs: P.kwargs) -> R:
mtries, mdelay = tries, delay
while mtries > 1:
try:
diff --git a/generalresearch/config.py b/generalresearch/config.py
index c6f41e8..c414069 100644
--- a/generalresearch/config.py
+++ b/generalresearch/config.py
@@ -16,9 +16,9 @@ def is_debug() -> bool:
import os
is_developer: bool = os.getenv("USER") in {"nanis", "gstupp"}
- is_pytest1: bool = bool(os.getenv("PYTEST_TEST", False))
- is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST", False))
- is_pytest3: bool = bool(os.getenv("PYTEST_VERSION", False))
+ is_pytest1: bool = bool(os.getenv("PYTEST_TEST"))
+ is_pytest2: bool = bool(os.getenv("PYTEST_CURRENT_TEST"))
+ is_pytest3: bool = bool(os.getenv("PYTEST_VERSION"))
is_debugging1: bool = os.getenv("DEBUG", "").lower() in ("1", "true", "yes")
is_debugging2: bool = os.getenv("PYTHON_DEBUG", "").lower() in ("1", "true", "yes")
is_jenkins: bool = bool(os.getenv("JENKINS_HOME")) or bool(os.getenv("JENKINS_URL"))
diff --git a/generalresearch/currency.py b/generalresearch/currency.py
index 9402e04..716cb0f 100644
--- a/generalresearch/currency.py
+++ b/generalresearch/currency.py
@@ -25,7 +25,7 @@ def format_usd_cent(usd_cent: int) -> str:
class USDCent(int):
- def __new__(cls, value, *args, **kwargs):
+ def __new__(cls, value: int, *args, **kwargs):
if isinstance(value, float):
warnings.warn(
@@ -42,17 +42,17 @@ class USDCent(int):
return super(cls, cls).__new__(cls, value)
- def __add__(self, other):
+ def __add__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__add__(other)
return self.__class__(res)
- def __sub__(self, other):
+ def __sub__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__sub__(other)
return self.__class__(res)
- def __mul__(self, other):
+ def __mul__(self, other: Any):
assert isinstance(other, USDCent)
res = super().__mul__(other)
return self.__class__(res)
@@ -61,14 +61,14 @@ class USDCent(int):
res = super().__abs__()
return self.__class__(res)
- def __truediv__(self, other):
+ def __truediv__(self):
raise ValueError("Division not allowed for USDCent")
def __str__(self):
- return "%d" % int(self)
+ return f"{int(self):d}"
def __repr__(self):
- return "USDCent(%d)" % int(self)
+ return f"USDCent({int(self)})"
@classmethod
def __get_pydantic_core_schema__(
@@ -110,17 +110,17 @@ class USDMill(int):
return super(cls, cls).__new__(cls, value)
- def __add__(self, other):
+ def __add__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__add__(other)
return self.__class__(res)
- def __sub__(self, other):
+ def __sub__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__sub__(other)
return self.__class__(res)
- def __mul__(self, other):
+ def __mul__(self, other: Any):
assert isinstance(other, USDMill)
res = super().__mul__(other)
return self.__class__(res)
@@ -129,14 +129,14 @@ class USDMill(int):
res = super().__abs__()
return self.__class__(res)
- def __truediv__(self, other):
+ def __truediv__(self):
raise ValueError("Division not allowed for USDMill")
def __str__(self):
- return "%d" % int(self)
+ return f"{int(self):d}"
def __repr__(self):
- return "USDMill(%d)" % int(self)
+ return f"USDMill({int(self)})"
@classmethod
def __get_pydantic_core_schema__(
diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py
index 093f7ae..c1eac37 100644
--- a/generalresearch/grliq/managers/forensic_data.py
+++ b/generalresearch/grliq/managers/forensic_data.py
@@ -103,14 +103,13 @@ class GrlIqDataManager:
is_attempt_allowed = %(is_attempt_allowed)s
WHERE uuid = %(uuid)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def update_fingerprint(self, iq_data: GrlIqData) -> None:
# We should only run this if we modified the fingerprint algorithm
@@ -123,14 +122,13 @@ class GrlIqDataManager:
SET fingerprint = %(fingerprint)s
WHERE uuid = %(uuid)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def update_data(self, iq_data: GrlIqData) -> None:
# We should only run this if we structured new fields and want to
@@ -141,14 +139,13 @@ class GrlIqDataManager:
SET data = %(data)s
WHERE id = %(id)s
""")
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(query, data)
- if c.rowcount != 1:
- raise ValueError(
- f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
- )
- conn.commit()
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(query, data)
+ if c.rowcount != 1:
+ raise ValueError(
+ f"Expected 1 row to be updated, but {c.rowcount} rows were affected."
+ )
+ conn.commit()
def get_data_if_exists(
self, forensic_uuid: UUIDStr, load_events: bool = False
@@ -472,7 +469,7 @@ class GrlIqDataManager:
filters.append("d.fingerprint = ANY(%(fingerprints)s)")
if product_ids and len(product_ids) == 1:
- product_id = list(product_ids)[0]
+ product_id = next(iter(product_ids))
product_ids = None
if product_ids:
@@ -536,9 +533,9 @@ class GrlIqDataManager:
[f"(%(bp_{i})s, %(bpuid_{i})s)" for i in range(len(users))]
)
filters.append(f"(d.product_id, d.product_user_id) IN ({user_args})")
- for i, user in enumerate(users):
- params[f"bp_{i}"] = user.product_id
- params[f"bpuid_{i}"] = user.product_user_id
+ for i, _user in enumerate(users):
+ params[f"bp_{i}"] = _user.product_id
+ params[f"bpuid_{i}"] = _user.product_user_id
if phase:
params["phase"] = phase.value
@@ -595,24 +592,19 @@ class GrlIqDataManager:
)
if only_product_id:
- try:
- with self.postgres_config.make_connection() as conn:
- with conn.cursor() as c:
- c.execute(
- query="""
- SELECT count AS c
- FROM grliq_forensicdata_product_counts
- WHERE product_id = %s
- LIMIT 1
- """,
- params=(product_id,),
- )
- res = c.fetchone()
- if res and res["c"] >= 0:
- return int(res["c"])
-
- except Exception:
- pass
+ with self.postgres_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ query="""
+ SELECT count AS c
+ FROM grliq_forensicdata_product_counts
+ WHERE product_id = %s
+ LIMIT 1
+ """,
+ params=(product_id,),
+ )
+ res = c.fetchone()
+ if res and res["c"] >= 0:
+ return int(res["c"])
query = f"""
SELECT COUNT(1) AS c
diff --git a/generalresearch/grliq/managers/forensic_summary.py b/generalresearch/grliq/managers/forensic_summary.py
index 98e00d8..b86e1f5 100644
--- a/generalresearch/grliq/managers/forensic_summary.py
+++ b/generalresearch/grliq/managers/forensic_summary.py
@@ -109,7 +109,7 @@ def calculate_timing_summary(
for k, v in country_distributions.items()
}
- out = dict()
+ out = {}
for country_iso, median_rtts in country_median_rtts.items():
country_stats = country_distributions[country_iso]
z_scores = [
@@ -158,12 +158,14 @@ def run_user_forensic_summary(
)
session_uuids = {x["session_uuid"] for x in res}
- timing_res: list[dict] = iq_em.filter_distinct_timing(session_uuids=session_uuids)
+ timing_res: list[dict[str, Any]] = iq_em.filter_distinct_timing(
+ session_uuids=session_uuids
+ )
country_timing_data_summary = (
calculate_timing_summary(redis_config=redis_config, timing_res=timing_res)
if timing_res
- else dict()
+ else {}
)
s = UserForensicSummary(
diff --git a/generalresearch/grliq/models/events.py b/generalresearch/grliq/models/events.py
index 69b67e5..995a6fa 100644
--- a/generalresearch/grliq/models/events.py
+++ b/generalresearch/grliq/models/events.py
@@ -177,12 +177,7 @@ class TimingData(BaseModel):
@property
def server_location(self) -> str:
- # TODO: when we have more locations ...
- return (
- "fremont_ca"
- if self.server_hostname in {"grliq-web-0", "grliq-web-1"}
- else "fremont_ca"
- )
+ return "fremont_ca"
@property
def has_data(self):
diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py
index 666cb81..69b1760 100644
--- a/generalresearch/grliq/models/forensic_data.py
+++ b/generalresearch/grliq/models/forensic_data.py
@@ -474,9 +474,11 @@ class GrlIqData(BaseModel):
description="Bit-packed string for font support. Each element is 32 bits, with each bit representing T/F for "
"font support.",
examples=[
- "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636"
- "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0"
- "|2147483648|109543424|1880099872|268435471"
+ (
+ "72|768|262144|1073741824|0|0|540672|73728|7340032|1342177280|117446656|256|16|0|543|4290797636"
+ "|1677723648|4168998400|0|1048576|262144|268500994|1342177280|262144|125829376|37888000|0|435363842|0"
+ "|2147483648|109543424|1880099872|268435471"
+ )
],
)
@@ -558,19 +560,21 @@ class GrlIqData(BaseModel):
@cached_property
def audio_codecs_named(self) -> dict[str, bool]:
+ assert self.audio_codecs
return dict(
zip(
AUDIO_CODEC_NAMES,
- [True if x == "3" else False for x in self.audio_codecs.split(",")],
+ [x == "3" for x in self.audio_codecs.split(",")],
)
)
@cached_property
def video_codecs_named(self) -> dict[str, bool]:
+ assert self.video_codecs
return dict(
zip(
VIDEO_CODEC_NAMES,
- [True if x == "3" else False for x in self.video_codecs.split(",")],
+ [x == "3" for x in self.video_codecs.split(",")],
)
)
@@ -779,9 +783,8 @@ class GrlIqData(BaseModel):
minutes=90
), "expired session"
-
def model_dump_sql(self, **kwargs) -> dict[str, Any]:
- d = dict()
+ d = {}
d["uuid"] = self.uuid
d["session_uuid"] = self.mid
d["created_at"] = self.created_at
@@ -801,7 +804,7 @@ class GrlIqData(BaseModel):
return d
@classmethod
- def from_db(cls, d: dict[str, Any]) -> Self:
+ def from_db(cls, d: dict[str, Any]) -> GrlIqData:
res = GrlIqData.model_validate(d["data"])
if d.get("category_result"):
diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py
index f5b0f25..6b0e065 100644
--- a/generalresearch/grliq/models/forensic_summary.py
+++ b/generalresearch/grliq/models/forensic_summary.py
@@ -251,10 +251,9 @@ class CountryRTTDistribution(BaseModel):
Render a boxplot from the RTT percentiles.
"""
try:
- # annoying pycharm error
import matplotlib.pyplot as plt
- except ImportError as e:
- raise e
+ except ImportError:
+ return
p = self.rtt_percentiles
data = {
@@ -266,7 +265,7 @@ class CountryRTTDistribution(BaseModel):
"fliers": [p[0]] + ([p[100]] if p[100] > p[95] else []),
}
- fig, ax = plt.subplots(figsize=(4, 1.5))
+ _, ax = plt.subplots(figsize=(4, 1.5))
ax.bxp([data], showfliers=True, vert=False)
ax.set_title(f"RTT Boxplot for {self.country_iso}")
ax.set_xlabel("RTT (ms)")
diff --git a/generalresearch/grliq/models/useragents.py b/generalresearch/grliq/models/useragents.py
index de63a67..3d5e5ce 100644
--- a/generalresearch/grliq/models/useragents.py
+++ b/generalresearch/grliq/models/useragents.py
@@ -132,7 +132,7 @@ class BrowserInfo(BaseModel):
class DeviceInfo(BaseModel):
family: DeviceModelFamily = Field()
- brand: DeviceBrand = Field()
+ brand: DeviceBrand | None = Field(default=None)
model: DeviceModelFamily = Field()
@field_validator("family", "model", mode="before")
@@ -186,7 +186,7 @@ class GrlUserAgent(BaseModel):
def ua_string_values(self) -> dict[str, str]:
# Returns the raw parsed string values for each of these. To be used
# for db filtering, identifying trends, etc.
- d = dict()
+ d = {}
d["ua_browser_family"] = self.ua_parsed.browser.family
d["ua_browser_version"] = self.ua_parsed.browser.version_string
d["ua_os_family"] = self.ua_parsed.os.family
diff --git a/generalresearch/healing_ppe.py b/generalresearch/healing_ppe.py
index 254a893..dfee689 100644
--- a/generalresearch/healing_ppe.py
+++ b/generalresearch/healing_ppe.py
@@ -73,7 +73,7 @@ def test():
time.sleep(0.5)
# Kill a process in the pool
- pid = list(pool._processes.keys())[0]
+ pid = next(iter(pool._processes.keys()))
os.kill(pid, signal.SIGKILL)
time.sleep(0.5)
diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py
index 9d504bd..060df4e 100644
--- a/generalresearch/incite/base.py
+++ b/generalresearch/incite/base.py
@@ -25,6 +25,7 @@ from uuid import uuid4
import dask
import dask.dataframe as dd
import pandas as pd
+import pandera as pa
import pyarrow.parquet as pq
from distributed import Client as DaskClient
from pandera.pandas import DataFrameSchema
@@ -45,6 +46,7 @@ from pydantic.json_schema import SkipJsonSchema
from sentry_sdk import capture_exception
from generalresearch.config import is_debug
+from generalresearch.incite.collections import DFCollectionItem
from generalresearch.incite.schemas import (
ARCHIVE_AFTER,
empty_dataframe_from_schema,
@@ -61,7 +63,7 @@ if TYPE_CHECKING:
Collection = DFCollection | MergeCollection
logging.basicConfig()
-LOG = logging.getLogger()
+LOG = logging.getLogger(f"{__name__}.incite")
# Item = Union["DFCollectionItem", "MergeCollectionItem"]
Item = Any
@@ -133,6 +135,7 @@ class GRLDatasets(BaseModel):
from generalresearch.incite.mergers import MergeType
folder = "mergers" if isinstance(enum_type, MergeType) else "raw/df-collections"
+ assert self.incite is not None
return Path(
pjoin(self.data_src, self.incite.point, folder, str(enum_type.value))
)
@@ -203,9 +206,6 @@ class CollectionBase(BaseModel):
@model_validator(mode="after")
def check_model_after(self) -> Self:
- if self.offset is None or self.start is None:
- return self
-
offset_total_sec = pd.Timedelta(self.offset).total_seconds()
start_total_sec = (datetime.now(tz=UTC) - self.start).total_seconds()
@@ -230,7 +230,7 @@ class CollectionBase(BaseModel):
return v
try:
pd.Timedelta(v)
- except Exception as e:
+ except (ValueError, TypeError) as e:
capture_exception(error=e)
raise ValueError(
"Invalid offset alias provided. Please review: "
@@ -554,7 +554,7 @@ class CollectionBase(BaseModel):
try:
pq.ParquetDataset(highest_version).read().to_pandas()
- except Exception:
+ except (pa.ArrowInvalid, pa.ArrowIOError, FileNotFoundError):
# 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
@@ -764,11 +764,12 @@ 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):
-
- # Don't include the "broken link" or mmfsymlink text file
- if fn != self.filename and fn != self.partial_filename:
- builds.append(fn)
+ if (
+ fn.startswith(self.filename)
+ and fn != self.filename
+ and fn != self.partial_filename
+ ):
+ builds.append(fn)
if len(builds) == 0:
return None
@@ -872,7 +873,7 @@ class CollectionItemBase(BaseModel):
raise ValueError("Unknown path type.")
df = parquet.read().to_pandas()
- except Exception:
+ except (pa.ArrowInvalid, pa.ArrowIOError, OSError):
LOG.warning(f"Invalid archive {path=}")
df = None
@@ -891,8 +892,7 @@ class CollectionItemBase(BaseModel):
try:
schema: DataFrameSchema = self._collection._schema
return schema.validate(check_obj=df, lazy=True, sample=sample)
- except Exception as e:
- LOG.exception(e)
+ except pa.errors.SchemaErrors as e:
capture_exception(error=e)
return None
diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py
index 42c3d31..f1e1b77 100644
--- a/generalresearch/incite/collections/__init__.py
+++ b/generalresearch/incite/collections/__init__.py
@@ -1,6 +1,5 @@
from __future__ import annotations
-import logging
import os
import subprocess
import time
@@ -12,16 +11,18 @@ from typing import Any
import dask
import dask.dataframe as dd
import pandas as pd
+import pyarrow as pa
import pyarrow.parquet as pq
+from dask.distributed import Client as DaskClient
from dask.distributed import Future
-from distributed import Client, as_completed
+from distributed import as_completed
from more_itertools import chunked
from pandera.pandas import DataFrameSchema
from psycopg import Cursor
from pydantic import Field, FilePath, ValidationInfo, field_validator
from sentry_sdk import capture_exception
-from generalresearch.incite.base import CollectionBase, CollectionItemBase
+from generalresearch.incite.base import LOG, CollectionBase, CollectionItemBase
from generalresearch.incite.schemas import (
ARCHIVE_AFTER,
ORDER_KEY,
@@ -49,9 +50,6 @@ from generalresearch.incite.schemas.thl_web import (
UserHealthIPHistoryWSSchema,
)
from generalresearch.pg_helper import PostgresConfig
-from generalresearch.sql_helper import SqlHelper
-
-LOG = logging.getLogger("incite")
DT_STR = "%Y-%m-%d %H:%M:%S"
@@ -106,18 +104,6 @@ class DFCollectionItem(CollectionItemBase):
# --- Methods ---
- def has_mysql(self) -> bool:
- if self._collection.sql_helper is None:
- return False
-
- connected = True
- try:
- self._collection.sql_helper.execute_sql_query("""SELECT 1;""")
- except:
- connected = False
-
- return connected
-
def has_postgres(self) -> bool:
if self._collection.pg_config is None:
return False
@@ -125,7 +111,7 @@ class DFCollectionItem(CollectionItemBase):
connected = True
try:
self._collection.pg_config.execute_sql_query("""SELECT 1;""")
- except:
+ except AssertionError:
connected = False
return connected
@@ -148,7 +134,7 @@ class DFCollectionItem(CollectionItemBase):
since = max([since, self.start]) # don't allow to query before the item's start
df = df[df[order_key] < since].copy()
- _df = self.from_mysql(since=since)
+ _df = self.from_db(since=since)
if _df is not None:
df = pd.concat([df, _df])
@@ -161,7 +147,7 @@ class DFCollectionItem(CollectionItemBase):
return True
def create_partial_archive(self) -> bool:
- _df = self.from_mysql()
+ _df = self.from_db()
if _df is None:
# Returned no rows, but the period is not closed, so we
# don't want to mark as empty. Do nothing.
@@ -172,71 +158,15 @@ class DFCollectionItem(CollectionItemBase):
def to_dict(self) -> dict[str, Any]:
return self._to_dict()
- def from_mysql(self, since: datetime | None = None) -> pd.DataFrame | None:
+ def from_db(self, since: datetime | None = None) -> pd.DataFrame | None:
if self._collection.data_type == DFCollectionType.LEDGER:
assert since is None, "Shouldn't pass since for Ledger item"
assert self._collection.pg_config is not None
return self.from_postgres_ledger()
else:
- if self._collection.sql_helper:
- return self.from_mysql_standard(since=since)
- else:
- return self.from_postgres_standard(since=since)
-
- def from_mysql_standard(self, since: datetime | None = None) -> pd.DataFrame | None:
-
- assert (
- self._collection.data_type != DFCollectionType.LEDGER
- ), "Can't call from_mysql_standard for Ledger DFCollectionItem"
-
- start, finish = self.start, self.finish
- LOG.debug(
- f"{self._collection.data_type.value}.from_mysql("
- f"start={start.strftime(DT_STR)}, "
- f"finish={finish.strftime(DT_STR)})"
- )
- coll = self._collection
- schema = coll._schema
- sql_helper = coll.sql_helper
-
- start = since or start
- order_key = schema.metadata[ORDER_KEY]
- cols = list(schema.columns.keys()) + [schema.index.name]
- cols_str = ",".join(map(sql_helper._quote, cols))
- db_name = sql_helper.db
-
- try:
- res = sql_helper.execute_sql_query(
- query=f"""
- SELECT {cols_str}
- FROM `{db_name}`.`{coll.data_type.value}`
- WHERE `{order_key}` >= %s AND `{order_key}` < %s;
- """,
- params=[start, finish],
- )
- except Exception as e:
- capture_exception(error=e)
- LOG.error(f"_from_mysql Exception: {e}")
- return None
-
- if not res:
- LOG.warning("_from_mysql query returned nothing")
- # Return an empty df.DataFrame with the correct columns
- return empty_dataframe_from_schema(coll._schema)
-
- df = pd.DataFrame.from_records(res).set_index(coll._schema.index.name)
- df = self.validate_df(df=df)
-
- if df is None:
- LOG.warning("_from_mysql query results failed validation")
- # Schema validation can fail...
- return None
-
- return df
+ return self.from_db_standard(since=since)
- def from_postgres_standard(
- self, since: datetime | None = None
- ) -> pd.DataFrame | None:
+ def from_db_standard(self, since: datetime | None = None) -> pd.DataFrame | None:
assert (
self._collection.data_type != DFCollectionType.LEDGER
), "Can't call from_postgres_standard for Ledger DFCollectionItem"
@@ -250,6 +180,7 @@ class DFCollectionItem(CollectionItemBase):
coll = self._collection
schema = coll._schema
pg_config = coll.pg_config
+ assert pg_config, "Must provide PostgresConfig"
start = since or start
order_key = schema.metadata[ORDER_KEY]
@@ -265,7 +196,7 @@ class DFCollectionItem(CollectionItemBase):
""",
params=[start, finish],
)
- except Exception as e:
+ except AssertionError as e:
capture_exception(error=e)
LOG.error(f"_from_postgres Exception: {e}")
return None
@@ -298,13 +229,14 @@ class DFCollectionItem(CollectionItemBase):
)
coll = self._collection
+ assert coll.pg_config, "Must provide PostgresConfig"
pg_config: PostgresConfig = coll.pg_config
limit = 20000
offset = 0
res = []
while True:
- logging.info(
+ LOG.info(
f"{self._collection.data_type.value}.from_postgres_ledger({limit=}, {offset=})"
)
chunk = pg_config.execute_sql_query(
@@ -399,7 +331,7 @@ class DFCollectionItem(CollectionItemBase):
"""
assert isinstance(ddf, dd.DataFrame), "must pass dask df"
- client: Client | None = self._collection._client
+ client: DaskClient | None = self._collection._client
# client = None
if client:
@@ -455,7 +387,9 @@ class DFCollectionItem(CollectionItemBase):
tmp_path = self.tmp_path()
try:
schema = self._collection._schema
- partition = schema.metadata.get(PARTITION_ON, None)
+ assert schema
+ assert schema.metadata
+ partition = schema.metadata.get(PARTITION_ON)
ddf.to_parquet(
path=tmp_path,
@@ -466,7 +400,7 @@ class DFCollectionItem(CollectionItemBase):
compression="brotli",
)
- except Exception as e:
+ except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e:
LOG.exception(e)
self.delete_archive(tmp_path)
return False
@@ -516,7 +450,8 @@ class DFCollectionItem(CollectionItemBase):
collection = self._collection
schema = collection._schema
- client: Client | None = collection._client
+ assert schema
+ client: DaskClient | None = collection._client
next_numbered_path = self.next_numbered_path(self.partial_path)
partial_path = self.partial_path
@@ -544,7 +479,8 @@ class DFCollectionItem(CollectionItemBase):
return False
try:
- partition = schema.metadata.get(PARTITION_ON, None)
+ assert schema.metadata
+ partition = schema.metadata.get(PARTITION_ON)
ddf.to_parquet(
path=next_numbered_path,
partition_on=partition,
@@ -553,7 +489,7 @@ class DFCollectionItem(CollectionItemBase):
write_metadata_file=True,
compression="brotli",
)
- except Exception as e:
+ except (pa.ArrowInvalid, pa.ArrowIOError, OSError) as e:
LOG.exception(e)
self.delete_archive(next_numbered_path)
return False
@@ -572,7 +508,7 @@ class DFCollectionItem(CollectionItemBase):
assert self.should_archive(), "not ready to archive!"
- df: pd.DataFrame | None = self.from_mysql()
+ df: pd.DataFrame | None = self.from_db()
if df is None:
self.set_empty()
@@ -592,7 +528,6 @@ class DFCollection(CollectionBase):
# --- Private ---
pg_config: PostgresConfig | None = Field(default=None)
- sql_helper: SqlHelper | None = Field(default=None)
def __repr__(self):
res = self.signature() + "\n"
@@ -620,7 +555,7 @@ class DFCollection(CollectionBase):
return res
@field_validator("data_type")
- def check_data_type(cls, data_type, info: ValidationInfo):
+ def check_data_type(cls, data_type: DFCollectionType | None, info: ValidationInfo):
if data_type is None:
raise ValueError("Must explicitly provide a data_type")
@@ -647,7 +582,7 @@ class DFCollection(CollectionBase):
def initial_load(
self,
- client: Client | None = None,
+ client: DaskClient | None = None,
sync: bool = True,
since: datetime | None = None,
client_resources: dict[str, Any] | None = None,
@@ -685,7 +620,7 @@ class DFCollection(CollectionBase):
if sync:
fs = client.compute(fs, sync=False, priority=2, resources=client_resources)
- ac = as_completed(fs, timeout=timeout)
+ _ = as_completed(fs, timeout=timeout)
return fs
else:
@@ -734,7 +669,7 @@ class DFCollection(CollectionBase):
def force_rr_latest(
self,
- client: Client,
+ client: DaskClient,
client_resources: dict[str, Any] | None = None,
sync: bool = True,
) -> list[Future]:
diff --git a/generalresearch/incite/exceptions.py b/generalresearch/incite/exceptions.py
new file mode 100644
index 0000000..112978f
--- /dev/null
+++ b/generalresearch/incite/exceptions.py
@@ -0,0 +1,19 @@
+class BuildItemsError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
+
+
+class BuildError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
+
+
+class FetchError(Exception):
+
+ def __init__(self, message: str):
+ self.message = message
+ super().__init__(self.message)
diff --git a/generalresearch/incite/mergers/__init__.py b/generalresearch/incite/mergers/__init__.py
index 45d3e2d..810bedc 100644
--- a/generalresearch/incite/mergers/__init__.py
+++ b/generalresearch/incite/mergers/__init__.py
@@ -108,13 +108,12 @@ class MergeCollectionItem(CollectionItemBase):
client: Client,
ddf: dd.DataFrame,
is_partial: bool = False,
- client_resources=None,
) -> bool:
assert is_partial is False, "use to_archive_symlink"
return self._to_archive(client=client, ddf=ddf, client_resources=None)
def _to_archive(
- self, client: Client, ddf: dd.DataFrame, client_resources=None
+ self, client: Client, ddf: dd.DataFrame | None, client_resources=None
) -> bool:
"""
For archiving an item. Will write an empty file if ddf is empty.
@@ -125,12 +124,15 @@ class MergeCollectionItem(CollectionItemBase):
if ddf is None:
return False
- row_len = client.compute(collections=ddf.shape[0], sync=True)
+ row_len: int = client.compute(collections=ddf.shape[0], sync=True)
+ assert row_len
assert row_len > 0, "empty ddf"
tmp_path = self.tmp_path()
schema = self._collection._schema
- partition = schema.metadata.get(PARTITION_ON, None)
+ assert schema.metadata
+
+ partition = schema.metadata.get(PARTITION_ON)
f = ddf.to_parquet(
compute=False,
path=tmp_path,
@@ -178,8 +180,7 @@ class MergeCollectionItem(CollectionItemBase):
collection = self._collection
LOG.warning(f"{collection.merge_type.value}.to_archive_symlink()")
- if not isinstance(ddf, dd.DataFrame):
- raise ValueError("must pass a dask df")
+ assert isinstance(ddf, dd.DataFrame), "must pass a dask df"
# We should validate before or after!!!
# _validate_df(self.compute(ddf), coll._schema)
@@ -212,13 +213,10 @@ class MergeCollectionItem(CollectionItemBase):
else:
subprocess.call(["ln", "-sfnT", target, path.as_posix()])
- if validate_after:
- if not self.valid_archive(self.path):
- LOG.error(
- f"{collection.merge_type.value} failed validation: {self.path}"
- )
- self.delete_archive(self.path)
- return False
+ if validate_after and not self.valid_archive(self.path):
+ LOG.error(f"{collection.merge_type.value} failed validation: {self.path}")
+ self.delete_archive(self.path)
+ return False
return True
# todo: unclear what the common interface should be here ... ?
@@ -256,7 +254,7 @@ class MergeCollection(CollectionBase):
return self
@field_validator("merge_type")
- def check_merge_type(cls, merge_type, info: ValidationInfo):
+ def check_merge_type(cls, merge_type: MergeType | None, info: ValidationInfo):
if merge_type is None:
raise ValueError("Must explicitly provide a merge_type")
diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py
index d3100fa..561e759 100644
--- a/generalresearch/incite/mergers/foundations/__init__.py
+++ b/generalresearch/incite/mergers/foundations/__init__.py
@@ -68,11 +68,9 @@ def lookup_product_and_team_id(
assert len(user_ids) <= 1000, "you should chunk this bro"
res: list[dict[str, Any]] = []
- with pg_config.make_connection() as conn:
- try:
- with conn.cursor() as c:
- c.execute(
- query="""
+ with pg_config.make_connection() as conn, conn.cursor() as c:
+ c.execute(
+ query="""
SELECT u.id AS user_id,
u.product_id,
bp.team_id
@@ -81,13 +79,9 @@ def lookup_product_and_team_id(
ON bp.id = u.product_id
WHERE u.id = ANY(%s);
""",
- params=[list(user_ids)],
- )
- res.extend(c.fetchall())
-
- except Exception as e:
- LOG.exception(f"lookup_product_and_team_id: {e}")
- raise
+ params=[list(user_ids)],
+ )
+ res.extend(c.fetchall())
return res
diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py
index 8a707ca..049b1bc 100644
--- a/generalresearch/incite/mergers/foundations/enriched_session.py
+++ b/generalresearch/incite/mergers/foundations/enriched_session.py
@@ -6,8 +6,8 @@ from typing import TYPE_CHECKING, Any, Literal
import dask.dataframe as dd
import pandas as pd
+from dask.distributed import Client as DaskClient
from dask.distributed import as_completed
-from distributed import Client
from more_itertools import chunked, flatten
from generalresearch.incite.collections.thl_web import (
@@ -46,7 +46,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
session_coll: SessionDFCollection,
wall_coll: WallDFCollection,
pg_config: PostgresConfig,
- client: Client | None = None,
+ client: DaskClient | None = None,
client_resources: dict[str, Any] | None = None,
) -> None:
@@ -140,9 +140,9 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
try:
results = client.gather(list(futures))
- except Exception as e:
+ except Exception:
client.cancel(futures, asynchronous=False, force=True)
- raise e
+ raise
dfp = pd.DataFrame(
list(flatten(results)), columns=["user_id", "product_id", "team_id"]
@@ -154,18 +154,13 @@ class EnrichedSessionMergeItem(MergeCollectionItem):
df = df[df["started"].between(start, end)]
is_missing = df[["product_id"]].isna().sum().sum() > 0
- session_is_partial = any([w.should_archive() is False for w in session_items])
+ session_is_partial = any(w.should_archive() is False for w in session_items)
session_is_missing = any(
- [
- w.should_archive() is True and w.has_archive() is False
- for w in session_items
- ]
+ w.should_archive() is True and w.has_archive() is False
+ for w in session_items
)
wall_is_missing = any(
- [
- w.should_archive() is True and w.has_archive() is False
- for w in wall_items
- ]
+ w.should_archive() is True and w.has_archive() is False for w in wall_items
)
is_partial = (
is_missing or session_is_partial or session_is_missing or wall_is_missing
@@ -203,7 +198,7 @@ class EnrichedSessionMerge(MergeCollection):
def build(
self,
- client: Client,
+ client: DaskClient,
session_coll: SessionDFCollection,
wall_coll: WallDFCollection,
pg_config: PostgresConfig,
@@ -232,7 +227,7 @@ class EnrichedSessionMerge(MergeCollection):
def to_admin_response(
self,
rr: ReportRequest,
- client: Client,
+ client: DaskClient,
product_ids: list[UUIDStr] | None = None,
user: User | None = None,
) -> pd.DataFrame:
@@ -243,6 +238,7 @@ class EnrichedSessionMerge(MergeCollection):
filters = []
if user:
+ assert product_ids
assert (
len(product_ids) <= 1
), "Can't search more than 1 Product ID for a specific User"
diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
index 234cd5b..e8a3654 100644
--- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
+++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py
@@ -11,6 +11,7 @@ from sentry_sdk import capture_exception
from generalresearch.incite.collections.thl_web import (
TaskAdjustmentDFCollection,
)
+from generalresearch.incite.exceptions import BuildError, BuildItemsError
from generalresearch.incite.mergers import (
MergeCollection,
MergeCollectionItem,
@@ -61,7 +62,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
]
if len(task_adj_coll_items) == 0:
- raise Exception("TaskAdjColl item collection failed")
+ raise BuildItemsError("TaskAdjColl item collection failed")
ddf: dd.DataFrame | None = task_adj_coll.ddf(
items=task_adj_coll_items,
@@ -83,6 +84,8 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
("started", "<", end),
],
)
+
+ assert isinstance(ddf, pd.DataFrame)
# Naked compute... don't log
# LOG.info(f"TaskAdjustmentDetailMergeCollectionItem.rows: {len(ddf.index)}")
@@ -91,7 +94,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem):
ew_items = [ew for ew in enriched_wall.items if ew.interval.overlaps(ir)]
if len(ew_items) == 0:
- raise Exception(
+ raise BuildItemsError(
"EnrichedWall item collection failed for EnrichedTaskAdjColl"
)
@@ -209,5 +212,5 @@ class EnrichedTaskAdjustMerge(MergeCollection):
enriched_wall=enriched_wall,
pg_config=pg_config,
)
- except Exception as e:
+ except BuildError as e:
capture_exception(error=e)
diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py
index b2ac7bb..a74a556 100644
--- a/generalresearch/incite/mergers/foundations/enriched_wall.py
+++ b/generalresearch/incite/mergers/foundations/enriched_wall.py
@@ -148,7 +148,7 @@ class EnrichedWallMergeItem(MergeCollectionItem):
is_missing = False
df = df.dropna(subset=["product_id", "session_id"], how="any")
- wall_is_partial = any([w.should_archive() is False for w in wall_items])
+ wall_is_partial = any(w.should_archive() is False for w in wall_items)
is_partial = is_missing or wall_is_partial
# Lots of downstream issues with this...
diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py
index 9750b57..a99e8ec 100644
--- a/generalresearch/incite/mergers/ym_survey_wall.py
+++ b/generalresearch/incite/mergers/ym_survey_wall.py
@@ -10,6 +10,7 @@ from distributed import Client
from sentry_sdk import capture_exception
from generalresearch.incite.collections.thl_web import WallDFCollection
+from generalresearch.incite.exceptions import BuildError
from generalresearch.incite.mergers import (
MergeCollection,
MergeCollectionItem,
@@ -109,7 +110,6 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem):
LOG.warning("YMSurveyWallMerge failed validation")
-
class YMSurveyWallMerge(MergeCollection):
merge_type: Literal[MergeType.YM_SURVEY_WALL] = MergeType.YM_SURVEY_WALL
collection_item_class: Literal[YMSurveyWallMergeCollectionItem] = (
@@ -143,7 +143,7 @@ class YMSurveyWallMerge(MergeCollection):
wall_coll=wall_coll,
enriched_session=enriched_session,
)
- except Exception as e:
+ except BuildError as e:
capture_exception(error=e)
item.delete_dangling_partials(keep_latest=2, target_path=item.path)
diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py
index 37fc3b9..69ef5c5 100644
--- a/generalresearch/incite/mergers/ym_wall_summary.py
+++ b/generalresearch/incite/mergers/ym_wall_summary.py
@@ -12,6 +12,7 @@ from generalresearch.incite.collections.thl_web import (
SessionDFCollection,
WallDFCollection,
)
+from generalresearch.incite.exceptions import FetchError
from generalresearch.incite.mergers import (
MergeCollection,
MergeCollectionItem,
@@ -41,6 +42,8 @@ class YMWallSummaryMergeItem(MergeCollectionItem):
ddf = wall_collection.ddf(
items=wall_items, force_rr_latest=False, include_partial=True
)
+ assert isinstance(ddf, pd.DataFrame)
+
ddf = ddf[ddf["started"].between(start, end)]
# Then we need the sessions for these wall events. They'll have started
@@ -88,12 +91,14 @@ class YMWallSummaryMerge(MergeCollection):
@field_validator("offset")
def check_offset_ym_wall_summary(cls, v: str | None):
# the offset MUST be on a whole day, no hourly
+ assert v
assert v.endswith("D"), "offset must be in days"
return v
@field_validator("start")
def check_start_ym_wall_summary(cls, v: datetime | None):
# the start MUST be start on midnight exactly
+ assert v
assert v.time() == time(0, 0, 0, 0), "start must no have a time component"
return v
@@ -117,7 +122,7 @@ class YMWallSummaryMerge(MergeCollection):
# item every time build is run even if it isn't closed
# if item.should_archive():
item.fetch(wall_collection, session_collection, user_id_product)
- except Exception as e:
+ except FetchError as e:
capture_exception(e)
@staticmethod
@@ -176,10 +181,10 @@ class YMWallSummaryMerge(MergeCollection):
# df.to_parquet(str(self.archive_path) + ".all.parquet")
pass
- def get_counts(self, product_id):
+ def get_counts(self, product_id: str):
# examples...
product_id = ""
- df = dd.read_parquet(
+ _ = dd.read_parquet(
str(self.archive_path) + ".all.parquet",
filters=[
("product_id", "=", product_id),
@@ -187,7 +192,7 @@ class YMWallSummaryMerge(MergeCollection):
],
).compute()
country_iso = "de"
- df = dd.read_parquet(
+ _ = dd.read_parquet(
str(self.archive_path) + ".all.parquet",
filters=[
("product_id", "=", product_id),
diff --git a/generalresearch/incite/schemas/admin_responses.py b/generalresearch/incite/schemas/admin_responses.py
index bb2852f..e65c6e2 100644
--- a/generalresearch/incite/schemas/admin_responses.py
+++ b/generalresearch/incite/schemas/admin_responses.py
@@ -1,5 +1,9 @@
-from datetime import datetime
+from __future__ import annotations
+from collections.abc import Callable
+from datetime import UTC, datetime
+
+import pandas as pd
from pandera.pandas import (
Check,
Column,
@@ -14,6 +18,14 @@ BIG_INT32 = 2_147_483_647
SIX_HOUR_SECONDS = 6 * 60 * 6
ROUNDING = 2
+_fillna: Callable[[pd.Series], pd.Series] = lambda s: s.fillna(value=0.00)
+_clip: Callable[[pd.Series], pd.Series] = lambda s: s.clip(
+ lower=0, upper=SIX_HOUR_SECONDS
+)
+_round: Callable[[pd.Series], pd.Series] = lambda s: s.round(decimals=ROUNDING)
+_tz_localize_none: Callable[[pd.Series], pd.Series] = lambda i: i.dt.tz_localize(None)
+
+
AdminPOPSchema = DataFrameSchema(
# Generic: used for Session or Wall
index=MultiIndex(
@@ -25,10 +37,15 @@ AdminPOPSchema = DataFrameSchema(
Index(
name="index0",
dtype=Timestamp,
- parsers=[Parser(lambda i: i.dt.tz_localize(None))],
+ parsers=[Parser(_tz_localize_none)],
checks=[
Check.less_than(
- max_value=datetime(year=datetime.now().year + 1, month=1, day=1)
+ max_value=datetime(
+ year=datetime.now(tz=UTC).year + 1,
+ month=1,
+ day=1,
+ tzinfo=UTC,
+ )
)
],
),
@@ -44,85 +61,85 @@ AdminPOPSchema = DataFrameSchema(
"elapsed_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.clip(lower=0, upper=SIX_HOUR_SECONDS)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_clip),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=SIX_HOUR_SECONDS),
),
"elapsed_total": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"payout_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=100),
),
"payout_total": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"entrances": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"completes": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"users": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"conversion": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0.00, max_value=1.00),
),
"epc": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=100),
),
"eph": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"cpc": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=250),
),
@@ -140,21 +157,21 @@ AdminPOPWallSchema = DataFrameSchema(
"buyers": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"surveys": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
"sessions": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
@@ -169,15 +186,15 @@ AdminPOPSessionSchema = DataFrameSchema(
"attempts_avg": Column(
dtype=float,
parsers=[
- Parser(lambda s: s.fillna(value=0.00)),
- Parser(lambda s: s.round(decimals=ROUNDING)),
+ Parser(_fillna),
+ Parser(_round),
],
checks=Check.between(min_value=0, max_value=25),
),
"attempts_total": Column(
dtype=int,
parsers=[
- Parser(lambda s: s.fillna(value=0)),
+ Parser(_fillna),
],
checks=Check.between(min_value=0, max_value=BIG_INT32),
),
diff --git a/generalresearch/locales/__init__.py b/generalresearch/locales/__init__.py
index 38c1832..4bb10d0 100644
--- a/generalresearch/locales/__init__.py
+++ b/generalresearch/locales/__init__.py
@@ -19,10 +19,6 @@ class Localelator:
EVERYTHING IS LOWERCASE!!! (except this comment)
"""
- lang_alpha2_to_alpha3b = dict()
- lang_alpha3_to_alpha3b = dict()
- languages = set()
-
def __init__(self):
d = json.loads(pkgutil.get_data(__name__, "iso639-3.json"))
self.lang_alpha2_to_alpha3b = {x["alpha_2"]: x["alpha_3b"] for x in d}
diff --git a/generalresearch/locales/setup_json.py b/generalresearch/locales/setup_json.py
index 57caa00..71beb09 100644
--- a/generalresearch/locales/setup_json.py
+++ b/generalresearch/locales/setup_json.py
@@ -1,62 +1,62 @@
-import json
-
-
-def country_default_lang():
- """
- Some marketplaces have no language specified. Surveys are in the "default
- language for that country", whatever that means. This helper is meant to
- provide a reasonable guess as to what language it is.
-
- Derived from: http://download.geonames.org/export/dump/countryInfo.txt
- """
- raise ValueError("no need to run this, I already ran it.")
- import pandas as pd
-
- from generalresearch.locales import Localelator
-
- l = Localelator()
-
- df = pd.read_csv(
- "http://download.geonames.org/export/dump/countryInfo.txt",
- sep="\t",
- skiprows=49,
- )
- df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0]
- df.default_lang = df.default_lang.fillna("en")
- df.default_lang = df.default_lang.map(
- lambda x: l.get_language_iso(x) if x in l.languages else "eng"
- )
- df["#ISO"] = df["#ISO"].str.lower()
- df["country_iso"] = df["#ISO"].map(
- lambda x: l.get_country_iso(x) if x in l.countries else None
- )
- df = df[df.country_iso.notnull()]
- d = df.set_index("country_iso").default_lang.to_dict()
- with open("country_default_lang.json", "w") as f:
- json.dump(d, f, indent=2)
- return d
-
-
-def setup_json():
- # pycountry is 30mb, which makes using this package on AWS lambda problematic.
- # These JSONs are stolen from pycountry and adapted.
-
- raise ValueError("no need to run this, I already ran it.")
-
- # languages
- d = json.load(open("iso639-3.json"))
- d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x]
- for x in d["639-3"]:
- x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"]
- del x["scope"]
- del x["type"]
- with open("iso639-3.json", "w") as f:
- json.dump(d["639-3"], f, indent=2)
-
- # countries
- d = json.load(open("iso3166-1.json"))["3166-1"]
- for x in d:
- x["alpha_2"] = x["alpha_2"].lower()
- x["alpha_3"] = x["alpha_3"].lower()
- with open("iso3166-1.json", "w") as f:
- json.dump(d, f, indent=2)
+# import json
+
+
+# def country_default_lang():
+# """
+# Some marketplaces have no language specified. Surveys are in the "default
+# language for that country", whatever that means. This helper is meant to
+# provide a reasonable guess as to what language it is.
+
+# Derived from: http://download.geonames.org/export/dump/countryInfo.txt
+# """
+# raise ValueError("no need to run this, I already ran it.")
+# import pandas as pd
+
+# from generalresearch.locales import Localelator
+
+# l = Localelator()
+
+# df = pd.read_csv(
+# "http://download.geonames.org/export/dump/countryInfo.txt",
+# sep="\t",
+# skiprows=49,
+# )
+# df["default_lang"] = df.Languages.str.split(",").str[0].str.split("-").str[0]
+# df.default_lang = df.default_lang.fillna("en")
+# df.default_lang = df.default_lang.map(
+# lambda x: l.get_language_iso(x) if x in l.languages else "eng"
+# )
+# df["#ISO"] = df["#ISO"].str.lower()
+# df["country_iso"] = df["#ISO"].map(
+# lambda x: l.get_country_iso(x) if x in l.countries else None
+# )
+# df = df[df.country_iso.notnull()]
+# d = df.set_index("country_iso").default_lang.to_dict()
+# with open("country_default_lang.json", "w") as f:
+# json.dump(d, f, indent=2)
+# return d
+
+
+# def setup_json():
+# # pycountry is 30mb, which makes using this package on AWS lambda problematic.
+# # These JSONs are stolen from pycountry and adapted.
+
+# raise ValueError("no need to run this, I already ran it.")
+
+# # languages
+# d = json.load(open("iso639-3.json"))
+# d["639-3"] = [x for x in d["639-3"] if "alpha_2" in x]
+# for x in d["639-3"]:
+# x["alpha_3b"] = x.pop("bibliographic", None) or x["alpha_3"]
+# del x["scope"]
+# del x["type"]
+# with open("iso639-3.json", "w") as f:
+# json.dump(d["639-3"], f, indent=2)
+
+# # countries
+# d = json.load(open("iso3166-1.json"))["3166-1"]
+# for x in d:
+# x["alpha_2"] = x["alpha_2"].lower()
+# x["alpha_3"] = x["alpha_3"].lower()
+# with open("iso3166-1.json", "w") as f:
+# json.dump(d, f, indent=2)
diff --git a/generalresearch/logging.py b/generalresearch/logging.py
index 9b72e0b..40f2174 100644
--- a/generalresearch/logging.py
+++ b/generalresearch/logging.py
@@ -1,6 +1,7 @@
import decimal
import json
from datetime import date
+from typing import Any
class ThlJsonEncoder(json.JSONEncoder):
@@ -11,11 +12,11 @@ class ThlJsonEncoder(json.JSONEncoder):
datetime/date to isoformat
"""
- def default(self, o):
+ def default(self, o: Any) -> Any:
if isinstance(o, decimal.Decimal):
return str(o)
if isinstance(o, set):
- return sorted(list(o))
+ return sorted(o)
if isinstance(o, date):
return o.isoformat()
return super().default(o)
diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py
index 819ae3d..686a964 100644
--- a/generalresearch/managers/cint/survey.py
+++ b/generalresearch/managers/cint/survey.py
@@ -141,5 +141,5 @@ class CintSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py
index c13b8ac..b8ae6ac 100644
--- a/generalresearch/managers/criteria.py
+++ b/generalresearch/managers/criteria.py
@@ -9,20 +9,21 @@ from more_itertools import chunked
from generalresearch.managers.base import SqlManager
from generalresearch.models.thl.survey import MarketplaceCondition
+DB_FIELDS = [
+ "hash",
+ "question_id",
+ "logical_operator",
+ "values",
+ "value_type",
+ "negate",
+]
+
class CriteriaManager(SqlManager, ABC):
"""
Using the terms "criteria" & "condition" interchangeably!
"""
- DB_FIELDS = [
- "hash",
- "question_id",
- "logical_operator",
- "values",
- "value_type",
- "negate",
- ]
CONDITION_MODEL = None
TABLE_NAME = ""
@@ -59,7 +60,7 @@ class CriteriaManager(SqlManager, ABC):
def update(self, conditions: Collection[MarketplaceCondition]) -> None:
# Add any new hashes into the DB
- this_hashes = set([condition.criterion_hash for condition in conditions])
+ this_hashes = {condition.criterion_hash for condition in conditions}
known_hashes = self.filter_exists(this_hashes)
new_hashes = this_hashes - known_hashes
@@ -95,7 +96,6 @@ class CriteriaManager(SqlManager, ABC):
)
conn.commit()
-
@property
def mysql_fields(self) -> str:
return ", ".join([f"`{k}`" for k in self.DB_FIELDS])
diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py
index 7643dc4..21ec42e 100644
--- a/generalresearch/managers/dynata/survey.py
+++ b/generalresearch/managers/dynata/survey.py
@@ -13,6 +13,34 @@ from generalresearch.models.dynata.survey import DynataCondition, DynataSurvey
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "status",
+ "is_live",
+ "client_id",
+ "bid_loi",
+ "bid_ir",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "expected_count",
+ "project_id",
+ "group_id",
+ "calculation_type",
+ "days_in_field",
+ "order_number",
+ "requirements",
+ "allowed_devices",
+ "category_exclusions",
+ "project_exclusions",
+ "live_link",
+ "category_ids",
+ "filters",
+ "quotas",
+ "used_question_ids",
+ "created",
+]
+
class DynataCriteriaManager(CriteriaManager):
CONDITION_MODEL = DynataCondition
@@ -20,33 +48,6 @@ class DynataCriteriaManager(CriteriaManager):
class DynataSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "status",
- "is_live",
- "client_id",
- "bid_loi",
- "bid_ir",
- "country_iso",
- "language_iso",
- "cpi",
- "expected_count",
- "project_id",
- "group_id",
- "calculation_type",
- "days_in_field",
- "order_number",
- "requirements",
- "allowed_devices",
- "category_exclusions",
- "project_exclusions",
- "live_link",
- "category_ids",
- "filters",
- "quotas",
- "used_question_ids",
- "created",
- ]
def get_survey_library(
self,
@@ -107,7 +108,7 @@ class DynataSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = ["id"] + self.SURVEY_FIELDS + ["last_updated"]
+ create_fields = ["id"] + SURVEY_FIELDS + ["last_updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -124,10 +125,10 @@ class DynataSurveyManager(SurveyManager):
def update(self, surveys: list[DynataSurvey]) -> bool:
now = datetime.now(tz=UTC)
- update_fields = self.SURVEY_FIELDS + ["last_updated"]
+ update_fields = SURVEY_FIELDS + ["last_updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update("dynata_survey", update_fields, survey_data)
return True
@@ -154,5 +155,6 @@ class DynataSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
+
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py
index efc8c0d..6779104 100644
--- a/generalresearch/managers/events.py
+++ b/generalresearch/managers/events.py
@@ -1,16 +1,16 @@
from __future__ import annotations
-import logging
import math
import socket
import threading
import time
from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import TYPE_CHECKING
+from typing import TYPE_CHECKING, Any
from redis.client import PubSub, Redis
+from generalresearch.incite.base import LOG
from generalresearch.managers.base import RedisManager
from generalresearch.models import Source
from generalresearch.models.custom_types import UUIDStr
@@ -216,16 +216,18 @@ class UserStatsManager(RedisManager):
class TaskStatsManager(RedisManager):
- task_stats = [
- "task_created_count_last_1h",
- "task_created_count_last_24h",
- "live_task_count",
- "live_tasks_max_payout",
- "TaskStatsManager:latest",
- ]
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
+
+ self.task_stats = [
+ "task_created_count_last_1h",
+ "task_created_count_last_24h",
+ "live_task_count",
+ "live_tasks_max_payout",
+ "TaskStatsManager:latest",
+ ]
+
self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT)
self.MAX_HASH_LUA = self.redis_client.register_script(MAX_HASH_LUA_SCRIPT)
@@ -435,8 +437,8 @@ class SessionStatsManager(RedisManager):
pipe.hincrby(name, key, 1)
pipe.hexpire(name, ttl, key, nx=True)
# BP-specific tracker
- pipe.hincrby(name + ":" + user.product_id, key, 1)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, 1)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
# We're not returning this, but keep the sums, so we can
# calculate the avg
@@ -444,8 +446,8 @@ class SessionStatsManager(RedisManager):
value = round(session.elapsed.total_seconds())
pipe.hincrby(name, key, value)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, value)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, value)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
pipe.execute()
@@ -476,31 +478,31 @@ class SessionStatsManager(RedisManager):
pipe.hincrby(name, key, 1)
pipe.hexpire(name, ttl, key, nx=True)
# BP-specific tracker
- pipe.hincrby(name + ":" + user.product_id, key, 1)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, 1)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
name = "sum_payouts_" + name_postfix
amount = round(session.payout * 100)
pipe.hincrby(name, key, amount)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, amount)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, amount)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
if session.user_payout:
name = "sum_user_payouts_" + name_postfix
amount = round(session.user_payout * 100)
pipe.hincrby(name, key, amount)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, amount)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, amount)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
# We're not returning this, but keep the sums, so we can calculate the avg
name = "session_complete_loi_sum_" + name_postfix
value = round(session.elapsed.total_seconds())
pipe.hincrby(name, key, value)
pipe.hexpire(name, ttl, key, nx=True)
- pipe.hincrby(name + ":" + user.product_id, key, value)
- pipe.hexpire(name + ":" + user.product_id, ttl, key, nx=True)
+ pipe.hincrby(f"{name}:{user.product_id}", key, value)
+ pipe.hexpire(f"{name}:{user.product_id}", ttl, key, nx=True)
pipe.execute()
@@ -563,6 +565,7 @@ class SessionStatsManager(RedisManager):
res["session_avg_user_payout_last_24h"] = None
res["session_complete_avg_loi_last_24h"] = None
res["session_fail_avg_loi_last_24h"] = None
+
if res["session_completes_last_24h"]:
res["session_avg_payout_last_24h"] = math.ceil(
res["sum_payouts_last_24h"] / res["session_completes_last_24h"]
@@ -630,7 +633,7 @@ class EventManager(StatsManager):
def get_active_subscribers(self) -> set[UUIDStr]:
res = self.redis_client.pubsub_channels(f"{self.cache_prefix}:event-channel:*")
- product_ids = {x.rsplit(":", 1)[-1] for x in res}
+ product_ids = {str(x.rsplit(":", 1)[-1]) for x in res}
return product_ids
def stats_worker(self):
@@ -638,7 +641,7 @@ class EventManager(StatsManager):
try:
self.stats_worker_task()
except Exception as e:
- logging.exception(e)
+ LOG.exception(e)
finally:
time.sleep(60)
@@ -654,14 +657,14 @@ class EventManager(StatsManager):
lock_key = f"{self.cache_prefix}:event-channel-lock"
res = self.redis_client.set(lock_key, 1, ex=120, nx=True)
if not res:
- logging.debug("failed to acquire stats_worker_task lock")
+ LOG.debug("failed to acquire stats_worker_task lock")
return
- logging.info("Acquired stats_worker_task lock")
+ LOG.info("Acquired stats_worker_task lock")
for product_id in self.get_active_subscribers():
if time.monotonic() - now > 120:
- logging.exception("stats_worker_task is taking too long")
+ LOG.exception("stats_worker_task is taking too long")
break
channel = self.get_channel_name(product_id)
msg = self.get_stats_message(product_id=product_id)
@@ -680,7 +683,7 @@ class EventManager(StatsManager):
return
- def make_influx_point(self, channel: str, numsub: int):
+ def make_influx_point(self, channel: str, numsub: int) -> dict[str, Any]:
return {
"measurement": "redis_pubsub_subscribers",
"tags": {"hostname": socket.gethostname(), "channel": channel},
diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py
index 721895e..851b88a 100644
--- a/generalresearch/managers/gr/authentication.py
+++ b/generalresearch/managers/gr/authentication.py
@@ -160,7 +160,7 @@ class GRUserManager(PostgresManagerWithRedis):
res = thl_pg_config.execute_sql_query(
query="""
- SELECT bp.id
+ SELECT bp.id::uuid as uuid
FROM userprofile_brokerageproduct AS bp
WHERE bp.business_id = ANY(%s)
""",
@@ -234,10 +234,10 @@ class GRTokenManager(PostgresManager):
res = c.fetchall()
if len(res) == 0:
- raise Exception(f"No GRUser with token of '{api_key}'")
+ raise ValueError(f"No GRUser with token of '{api_key}'")
if len(res) > 1:
- raise Exception(f"Too many GRUsers found with token of '{api_key}'")
+ raise ValueError(f"Too many GRUsers found with token of '{api_key}'")
item = res[0]
diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py
index f6d00a8..a4e36c1 100644
--- a/generalresearch/managers/innovate/survey.py
+++ b/generalresearch/managers/innovate/survey.py
@@ -16,6 +16,44 @@ from generalresearch.models.innovate.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "buyer_id",
+ "job_id",
+ "survey_name",
+ "desired_count",
+ "remaining_count",
+ "supplier_completes_achieved",
+ "global_completes",
+ "global_starts",
+ "global_median_loi",
+ "global_conversion",
+ "bid_loi",
+ "bid_ir",
+ "allowed_devices",
+ "entry_link",
+ "category",
+ "requires_pii",
+ "excluded_surveys",
+ "duplicate_check_level",
+ "exclude_pids",
+ "include_pids",
+ "is_revenue_sharing",
+ "group_type",
+ "off_hour_traffic",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+ "modified_api",
+ "created_api",
+ "expected_end_date",
+]
+
class InnovateCriteriaManager(CriteriaManager):
CONDITION_MODEL = InnovateCondition
@@ -23,43 +61,6 @@ class InnovateCriteriaManager(CriteriaManager):
class InnovateSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "buyer_id",
- "job_id",
- "survey_name",
- "desired_count",
- "remaining_count",
- "supplier_completes_achieved",
- "global_completes",
- "global_starts",
- "global_median_loi",
- "global_conversion",
- "bid_loi",
- "bid_ir",
- "allowed_devices",
- "entry_link",
- "category",
- "requires_pii",
- "excluded_surveys",
- "duplicate_check_level",
- "exclude_pids",
- "include_pids",
- "is_revenue_sharing",
- "group_type",
- "off_hour_traffic",
- "qualifications",
- "quotas",
- "used_question_ids",
- "is_live",
- "modified_api",
- "created_api",
- "expected_end_date",
- ]
def get_survey_library(
self,
@@ -104,7 +105,7 @@ class InnovateSurveyManager(SurveyManager):
assert filters, "Must set at least 1 filter"
filter_str = " AND ".join(filters)
filter_str = "WHERE " + filter_str if filter_str else ""
- fields = set(self.SURVEY_FIELDS) | {"created", "updated"}
+ fields = set(SURVEY_FIELDS) | {"created", "updated"}
if exclude_fields:
fields -= exclude_fields
@@ -126,7 +127,7 @@ class InnovateSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -143,10 +144,10 @@ class InnovateSurveyManager(SurveyManager):
def update(self, surveys: list[InnovateSurvey]) -> bool:
now = datetime.now(tz=UTC)
- update_fields = self.SURVEY_FIELDS + ["updated"]
+ update_fields = SURVEY_FIELDS + ["updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update(
table_name="innovate_survey",
field_names=update_fields,
diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py
index 2488cd8..b43cd71 100644
--- a/generalresearch/managers/morning/survey.py
+++ b/generalresearch/managers/morning/survey.py
@@ -14,6 +14,50 @@ from generalresearch.models.morning.survey import MorningBid, MorningCondition
logger = logging.getLogger()
+STAT_FIELDS = [
+ "obs_median_loi",
+ "qualified_conversion",
+ "num_available",
+ "num_completes",
+ "num_failures",
+ "num_in_progress",
+ "num_over_quotas",
+ "num_qualified",
+ "num_quality_terminations",
+ "num_timeouts",
+]
+STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"]
+BID_FIELDS = (
+ [
+ "id",
+ "status",
+ "country_iso",
+ "language_isos",
+ "buyer_account_id",
+ "buyer_id",
+ "name",
+ "supplier_exclusive",
+ "survey_type",
+ "timeout",
+ "topic_id",
+ "bid_loi",
+ "exclusions",
+ "used_question_ids",
+ "expected_end",
+ "created_api",
+ "is_live",
+ ]
+ + STAT_FIELDS
+ + STAT_EXTENDED_FIELDS
+)
+QUOTA_FIELDS = [
+ "id",
+ "cpi",
+ "condition_hashes",
+] + STAT_FIELDS
+BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`"
+QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`"
+
class MorningCriteriaManager(CriteriaManager):
CONDITION_MODEL = MorningCondition
@@ -21,49 +65,6 @@ class MorningCriteriaManager(CriteriaManager):
class MorningSurveyManager(SurveyManager):
- STAT_FIELDS = [
- "obs_median_loi",
- "qualified_conversion",
- "num_available",
- "num_completes",
- "num_failures",
- "num_in_progress",
- "num_over_quotas",
- "num_qualified",
- "num_quality_terminations",
- "num_timeouts",
- ]
- STAT_EXTENDED_FIELDS = ["system_conversion", "num_entrants", "num_screenouts"]
- BID_FIELDS = (
- [
- "id",
- "status",
- "country_iso",
- "language_isos",
- "buyer_account_id",
- "buyer_id",
- "name",
- "supplier_exclusive",
- "survey_type",
- "timeout",
- "topic_id",
- "bid_loi",
- "exclusions",
- "used_question_ids",
- "expected_end",
- "created_api",
- "is_live",
- ]
- + STAT_FIELDS
- + STAT_EXTENDED_FIELDS
- )
- QUOTA_FIELDS = [
- "id",
- "cpi",
- "condition_hashes",
- ] + STAT_FIELDS
- BID_DB_SOURCE = "`thl-morning`.`morning_surveybid`"
- QUOTA_DB_SOURCE = "`thl-morning`.`morning_surveyquota`"
def get_survey_library(
self,
diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py
index 54d74b7..179b8a9 100644
--- a/generalresearch/managers/network/mtr.py
+++ b/generalresearch/managers/network/mtr.py
@@ -41,8 +41,9 @@ class MTRRunManager(PostgresManager):
c.execute(query, params)
if params_hops:
c.executemany(query_hops, params_hops)
+
else:
- with self.pg_config.make_connection() as conn, conn.cursor() as c:
- c.execute(query, params)
+ with self.pg_config.make_connection() as conn, conn.cursor() as _c:
+ _c.execute(query, params)
if params_hops:
- c.executemany(query_hops, params_hops)
+ _c.executemany(query_hops, params_hops)
diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py
index 84d13ad..574bce1 100644
--- a/generalresearch/managers/network/nmap.py
+++ b/generalresearch/managers/network/nmap.py
@@ -50,7 +50,7 @@ class NmapRunManager(PostgresManager):
if nmap_run.ports:
c.executemany(query_ports, params_ports)
else:
- with self.pg_config.make_connection() as conn, conn.cursor() as c:
+ with self.pg_config.make_connection() as conn, conn.cursor():
c.execute(query, params)
if nmap_run.ports:
c.executemany(query_ports, params_ports)
diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py
index c8ce913..95a1381 100644
--- a/generalresearch/managers/network/rdns.py
+++ b/generalresearch/managers/network/rdns.py
@@ -27,6 +27,7 @@ class RDNSRunManager(PostgresManager):
params = run.model_dump_postgres()
if c:
c.execute(query, params)
+
else:
- with self.pg_config.make_connection() as conn, conn.cursor() as c:
- c.execute(query, params)
+ with self.pg_config.make_connection() as conn, conn.cursor() as _c:
+ _c.execute(query, params)
diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py
index c13dca8..cc28287 100644
--- a/generalresearch/managers/precision/survey.py
+++ b/generalresearch/managers/precision/survey.py
@@ -16,6 +16,30 @@ from generalresearch.models.precision.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ # 'country_iso', 'language_iso', # these come from join table
+ "survey_id",
+ "is_live",
+ "status",
+ "cpi",
+ "group_id",
+ "name",
+ "survey_guid",
+ "buyer_id",
+ "category_id",
+ "bid_loi",
+ "bid_ir",
+ "global_conversion",
+ "desired_count",
+ "achieved_count",
+ "allowed_devices",
+ "entry_link",
+ "excluded_surveys",
+ "quotas",
+ "used_question_ids",
+ "expected_end_date",
+]
+
class PrecisionCriteriaManager(CriteriaManager):
CONDITION_MODEL = PrecisionCondition
@@ -23,29 +47,6 @@ class PrecisionCriteriaManager(CriteriaManager):
class PrecisionSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- # 'country_iso', 'language_iso', # these come from join table
- "survey_id",
- "is_live",
- "status",
- "cpi",
- "group_id",
- "name",
- "survey_guid",
- "buyer_id",
- "category_id",
- "bid_loi",
- "bid_ir",
- "global_conversion",
- "desired_count",
- "achieved_count",
- "allowed_devices",
- "entry_link",
- "excluded_surveys",
- "quotas",
- "used_question_ids",
- "expected_end_date",
- ]
def get_survey_library(
self,
@@ -109,7 +110,7 @@ class PrecisionSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(False)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -241,5 +242,5 @@ class PrecisionSurveyManager(SurveyManager):
if e.args[0] == 1062:
existing_sns.add(sn)
else:
- raise e
+ raise
self.update([surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py
index 983572c..ef98d7b 100644
--- a/generalresearch/managers/prodege/survey.py
+++ b/generalresearch/managers/prodege/survey.py
@@ -9,6 +9,31 @@ from generalresearch.managers.criteria import CriteriaManager
from generalresearch.managers.survey import SurveyManager
from generalresearch.models.prodege.survey import ProdegeCondition, ProdegeSurvey
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_name",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "desired_count",
+ "remaining_count",
+ "achieved_completes",
+ "bid_loi",
+ "bid_ir",
+ "actual_loi",
+ "actual_ir",
+ "conversion_rate",
+ "entrance_url",
+ "max_clicks_settings",
+ "past_participation",
+ "include_psids",
+ "exclude_psids",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+]
+
class ProdegeCriteriaManager(CriteriaManager):
CONDITION_MODEL = ProdegeCondition
@@ -16,30 +41,6 @@ class ProdegeCriteriaManager(CriteriaManager):
class ProdegeSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_name",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "desired_count",
- "remaining_count",
- "achieved_completes",
- "bid_loi",
- "bid_ir",
- "actual_loi",
- "actual_ir",
- "conversion_rate",
- "entrance_url",
- "max_clicks_settings",
- "past_participation",
- "include_psids",
- "exclude_psids",
- "quotas",
- "used_question_ids",
- "is_live",
- ]
def get_survey_library(
self,
@@ -98,7 +99,7 @@ class ProdegeSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py
index fe6f621..ab66374 100644
--- a/generalresearch/managers/repdata/survey.py
+++ b/generalresearch/managers/repdata/survey.py
@@ -15,6 +15,40 @@ from generalresearch.models.repdata.survey import (
RepDataSurveyHashed,
)
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_uuid",
+ "survey_name",
+ "project_uuid",
+ "survey_status",
+ "country_iso",
+ "language_iso",
+ "estimated_loi",
+ "estimated_ir",
+ "collects_pii",
+ "allowed_devices",
+]
+STREAM_FIELDS = [
+ "stream_id",
+ "stream_uuid",
+ "stream_name",
+ "stream_status",
+ "calculation_type",
+ "qualification_hashes",
+ "hashed_quotas",
+ "expected_count",
+ "cpi",
+ "days_in_field",
+ "actual_ir",
+ "actual_loi",
+ "actual_conversion",
+ "actual_complete_count",
+ "actual_count",
+ "used_question_ids",
+ "survey_id",
+ "remaining_count",
+]
+
class RepDataCriteriaManager(CriteriaManager):
CONDITION_MODEL = RepDataCondition
@@ -22,39 +56,6 @@ class RepDataCriteriaManager(CriteriaManager):
class RepDataSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_uuid",
- "survey_name",
- "project_uuid",
- "survey_status",
- "country_iso",
- "language_iso",
- "estimated_loi",
- "estimated_ir",
- "collects_pii",
- "allowed_devices",
- ]
- STREAM_FIELDS = [
- "stream_id",
- "stream_uuid",
- "stream_name",
- "stream_status",
- "calculation_type",
- "qualification_hashes",
- "hashed_quotas",
- "expected_count",
- "cpi",
- "days_in_field",
- "actual_ir",
- "actual_loi",
- "actual_conversion",
- "actual_complete_count",
- "actual_count",
- "used_question_ids",
- "survey_id",
- "remaining_count",
- ]
def get_survey_library(
self,
@@ -127,7 +128,7 @@ class RepDataSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "last_updated"]
+ create_fields = SURVEY_FIELDS + ["created", "last_updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -141,10 +142,10 @@ class RepDataSurveyManager(SurveyManager):
args=survey_data,
)
- fields_str = ", ".join([f"`{x}`" for x in self.STREAM_FIELDS])
- values_str = ", ".join([f"%({x})s" for x in self.STREAM_FIELDS])
+ fields_str = ", ".join([f"`{x}`" for x in STREAM_FIELDS])
+ values_str = ", ".join([f"%({x})s" for x in STREAM_FIELDS])
stream_data = [
- {k: v for k, v in stream.items() if k in self.STREAM_FIELDS}
+ {k: v for k, v in stream.items() if k in STREAM_FIELDS}
for stream in d["streams"]
]
for sd in stream_data:
@@ -161,10 +162,10 @@ class RepDataSurveyManager(SurveyManager):
def update(self, surveys: list[RepDataSurveyHashed]) -> bool:
now = datetime.now(tz=UTC)
- update_fields = self.SURVEY_FIELDS + ["last_updated"]
+ update_fields = SURVEY_FIELDS + ["last_updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update(
table_name="repdata_survey",
field_names=update_fields,
@@ -175,7 +176,7 @@ class RepDataSurveyManager(SurveyManager):
for d in data:
for stream in d["streams"]:
stream["survey_id"] = d["survey_id"]
- stream_data.append([stream[k] for k in self.STREAM_FIELDS])
+ stream_data.append([stream[k] for k in STREAM_FIELDS])
self.sql_helper.bulk_update(
table_name="repdata_surveystream",
diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py
index a13fbce..9228528 100644
--- a/generalresearch/managers/sago/survey.py
+++ b/generalresearch/managers/sago/survey.py
@@ -13,6 +13,31 @@ from generalresearch.models.sago.survey import SagoCondition, SagoSurvey
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "is_live",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "buyer_id",
+ "account_id",
+ "study_type_id",
+ "industry_id",
+ "allowed_devices",
+ "collects_pii",
+ "bid_loi",
+ "bid_ir",
+ "live_link",
+ "survey_exclusions",
+ "ip_exclusions",
+ "remaining_count",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "modified_api",
+]
+
class SagoCriteriaManager(CriteriaManager):
CONDITION_MODEL = SagoCondition
@@ -20,30 +45,6 @@ class SagoCriteriaManager(CriteriaManager):
class SagoSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "is_live",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "buyer_id",
- "account_id",
- "study_type_id",
- "industry_id",
- "allowed_devices",
- "collects_pii",
- "bid_loi",
- "bid_ir",
- "live_link",
- "survey_exclusions",
- "ip_exclusions",
- "remaining_count",
- "qualifications",
- "quotas",
- "used_question_ids",
- "modified_api",
- ]
def get_survey_library(
self,
@@ -85,7 +86,7 @@ class SagoSurveyManager(SurveyManager):
assert filters, "Must set at least 1 filter"
filter_str = " AND ".join(filters)
filter_str = "WHERE " + filter_str if filter_str else ""
- fields = set(self.SURVEY_FIELDS) | {"created", "updated"}
+ fields = set(SURVEY_FIELDS) | {"created", "updated"}
if exclude_fields:
fields -= exclude_fields
fields_str = ", ".join([f"`{v}`" for v in fields])
@@ -106,7 +107,7 @@ class SagoSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["created", "updated"]
+ create_fields = SURVEY_FIELDS + ["created", "updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
@@ -123,10 +124,10 @@ class SagoSurveyManager(SurveyManager):
def update(self, surveys: list[SagoSurvey]) -> bool:
now = datetime.now(tz=UTC)
- update_fields = self.SURVEY_FIELDS + ["updated"]
+ update_fields = SURVEY_FIELDS + ["updated"]
data = [survey.to_mysql() for survey in surveys]
- survey_data = [[d[k] for k in self.SURVEY_FIELDS] + [now] for d in data]
+ survey_data = [[d[k] for k in SURVEY_FIELDS] + [now] for d in data]
self.sql_helper.bulk_update("sago_survey", update_fields, survey_data)
return True
@@ -156,8 +157,8 @@ class SagoSurveyManager(SurveyManager):
return True
def create_or_update(self, surveys: list[SagoSurvey]) -> None:
- surveys = {s.survey_id: s for s in surveys}
- sns = set(surveys.keys())
+ _surveys = {s.survey_id: s for s in surveys}
+ sns = set(_surveys.keys())
existing_sns = {
x["survey_id"]
for x in self.sql_helper.execute_sql_query(
@@ -171,7 +172,7 @@ class SagoSurveyManager(SurveyManager):
}
create_sns = sns - existing_sns
for sn in create_sns:
- survey = surveys[sn]
+ survey = _surveys[sn]
try:
self.create(survey)
except IntegrityError as e:
@@ -181,4 +182,4 @@ class SagoSurveyManager(SurveyManager):
else:
raise
- self.update([surveys[sn] for sn in existing_sns])
+ self.update([_surveys[sn] for sn in existing_sns])
diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py
index 987ce7b..5059716 100644
--- a/generalresearch/managers/spectrum/survey.py
+++ b/generalresearch/managers/spectrum/survey.py
@@ -16,6 +16,37 @@ from generalresearch.models.spectrum.survey import (
logger = logging.getLogger()
+SURVEY_FIELDS = [
+ "survey_id",
+ "survey_name",
+ "status",
+ "country_iso",
+ "language_iso",
+ "cpi",
+ "field_end_date",
+ "category_code",
+ "calculation_type",
+ "requires_pii",
+ "buyer_id",
+ "survey_exclusions",
+ "exclusion_period",
+ "bid_loi",
+ "bid_ir",
+ "last_block_loi",
+ "last_block_ir",
+ "overall_ir",
+ "overall_loi",
+ "project_last_complete_date",
+ "include_psids",
+ "exclude_psids",
+ "qualifications",
+ "quotas",
+ "used_question_ids",
+ "is_live",
+ "modified_api",
+ "created_api",
+]
+
class SpectrumCriteriaManager(CriteriaManager):
CONDITION_MODEL = SpectrumCondition
@@ -23,36 +54,6 @@ class SpectrumCriteriaManager(CriteriaManager):
class SpectrumSurveyManager(SurveyManager):
- SURVEY_FIELDS = [
- "survey_id",
- "survey_name",
- "status",
- "country_iso",
- "language_iso",
- "cpi",
- "field_end_date",
- "category_code",
- "calculation_type",
- "requires_pii",
- "buyer_id",
- "survey_exclusions",
- "exclusion_period",
- "bid_loi",
- "bid_ir",
- "last_block_loi",
- "last_block_ir",
- "overall_ir",
- "overall_loi",
- "project_last_complete_date",
- "include_psids",
- "exclude_psids",
- "qualifications",
- "quotas",
- "used_question_ids",
- "is_live",
- "modified_api",
- "created_api",
- ]
def get_survey_library(
self,
@@ -115,7 +116,7 @@ class SpectrumSurveyManager(SurveyManager):
conn: pymysql.Connection = self.sql_helper.make_connection()
conn.autocommit(True)
c = conn.cursor()
- create_fields = self.SURVEY_FIELDS + ["updated"]
+ create_fields = SURVEY_FIELDS + ["updated"]
fields_str = ", ".join([f"`{x}`" for x in create_fields])
values_str = ", ".join([f"%({x})s" for x in create_fields])
diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py
index 60bade7..2d89334 100644
--- a/generalresearch/managers/thl/task_adjustment.py
+++ b/generalresearch/managers/thl/task_adjustment.py
@@ -24,6 +24,9 @@ from generalresearch.models.thl.session import (
)
from generalresearch.models.thl.task_adjustment import TaskAdjustmentEvent
+logging.basicConfig()
+logger = logging.getLogger(__name__)
+
class TaskAdjustmentManager(PostgresManager):
@@ -129,8 +132,9 @@ class TaskAdjustmentManager(PostgresManager):
user.prefetch_product(self.pg_config)
if adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL:
+ assert wall.cpi
amount_usd = wall.cpi * -1
- adjusted_cpi = 0
+ adjusted_cpi = Decimal(0)
elif adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE:
amount_usd = wall.cpi
adjusted_cpi = wall.cpi
@@ -169,7 +173,7 @@ class TaskAdjustmentManager(PostgresManager):
new_adjusted_cpi=new_adjusted_cpi,
)
except AssertionError as e:
- logging.warning(e)
+ logger.warning(e)
return
event = TaskAdjustmentEvent(
diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
index e0a7548..af65d65 100644
--- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py
+++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py
@@ -172,7 +172,7 @@ class MysqlUserManager:
return user
- @lru_cache(maxsize=5000)
+ @lru_cache(maxsize=5_000)
def product_id_exists(self, product_id: str):
# 'id' is the primary key, there can only be 0 or 1
query = """
diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py
index 26f08b4..f28fe0c 100644
--- a/generalresearch/managers/thl/userhealth.py
+++ b/generalresearch/managers/thl/userhealth.py
@@ -374,10 +374,10 @@ class AuditLogManager(PostgresManager):
)
if len(res) == 0:
- raise Exception(f"No AuditLog with id of '{auditlog_id}'")
+ raise ValueError(f"No AuditLog with id of '{auditlog_id}'")
if len(res) > 1:
- raise Exception(f"Too many AuditLog found with id of '{auditlog_id}'")
+ raise ValueError(f"Too many AuditLog found with id of '{auditlog_id}'")
return AuditLog.from_mysql(res[0])
diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py
index ac9fb62..774db9e 100644
--- a/generalresearch/managers/thl/wall.py
+++ b/generalresearch/managers/thl/wall.py
@@ -579,7 +579,7 @@ class WallCacheManager(PostgresManagerWithRedis):
# b as second element and a as third element"
attempts = sorted(attempts, key=lambda x: x.started)
json_res = [attempt.model_dump_json() for attempt in attempts]
- res = self.redis_client.lpush(redis_key, *json_res)
+ _ = self.redis_client.lpush(redis_key, *json_res)
self.redis_client.expire(redis_key, time=60 * 60 * 24)
# So this doesn't grow forever, keep only the most recent 5k
diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py
index 064c200..11a5770 100644
--- a/generalresearch/models/gr/business.py
+++ b/generalresearch/models/gr/business.py
@@ -38,6 +38,9 @@ from generalresearch.redis_helper import RedisConfig
from generalresearch.utils.aggregation import group_by_year
from generalresearch.utils.enum import ReprEnumMeta
+logging.basicConfig()
+logger = logging.getLogger(__name__)
+
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
from generalresearch.incite.mergers.foundations.enriched_session import (
@@ -324,7 +327,7 @@ class Business(BaseModel):
for product_uuid in self.product_uuids:
if product_uuid not in bp_account:
refresh = True
- logging.exception(
+ logger.exception(
f"Business {self.uuid} does not have a BP Wallet Account for Product {product_uuid}. Creating..."
)
product = product_lookup[product_uuid]
diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py
index 3705f4f..2650b90 100644
--- a/generalresearch/models/legacy/bucket.py
+++ b/generalresearch/models/legacy/bucket.py
@@ -441,6 +441,12 @@ class DurationSummary(StatisticalSummary):
@classmethod
def from_bucket(cls, bucket: Bucket) -> DurationSummary:
+ assert bucket.loi_min
+ assert bucket.loi_max
+ assert bucket.loi_q1
+ assert bucket.loi_q2
+ assert bucket.loi_q3
+
return cls(
min=bucket.loi_min.total_seconds(),
max=bucket.loi_max.total_seconds(),
diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py
index a397247..b1ac556 100644
--- a/generalresearch/pg_helper.py
+++ b/generalresearch/pg_helper.py
@@ -3,6 +3,7 @@ from __future__ import annotations
from datetime import UTC
import psycopg
+from psycopg.abc import Query
from psycopg.rows import RowFactory, dict_row
from psycopg.types.datetime import TimestampLoader
from psycopg.types.net import InetLoader
@@ -74,7 +75,9 @@ class PostgresConfig:
self.row_factory = row_factory
@property
- def db(self):
+ def db(self) -> str:
+ assert self.dsn
+ assert self.dsn.path
return self.dsn.path[1:]
def make_connection(self) -> psycopg.Connection:
@@ -98,15 +101,15 @@ class PostgresConfig:
conn.adapters.register_loader("inet", InetHostLoader)
return conn
- def execute_sql_query(self, query, params=None):
+ def execute_sql_query(self, query: Query, params=None):
# This is only intended for SELECT queries
- assert "SELECT" in query.upper(), "Supports SELECTs only"
+ assert "SELECT" in str(query).upper(), "Supports SELECTs only"
with self.make_connection() as conn, conn.cursor() as c:
c.execute(query=query, params=params)
return c.fetchall()
- def execute_write(self, query, params=None) -> int:
+ def execute_write(self, query: Query, params=None) -> int:
cmd = query.lstrip().upper()
assert cmd.startswith(
("INSERT", "UPDATE", "DELETE")
diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py
index ae2b8d8..ef53c3b 100644
--- a/generalresearch/sql_helper.py
+++ b/generalresearch/sql_helper.py
@@ -14,6 +14,9 @@ ListOrTupleOfListOrTuple = (
DataBaseDsn = MySQLDsn | MariaDBDsn | PostgresDsn | None
+logging.basicConfig()
+logger = logging.getLogger(__name__)
+
class MultipleObjectsReturned(Exception):
pass
@@ -131,7 +134,7 @@ class SqlHelper(SqlConnector):
) -> list[dict[str, Any]]:
for param in params if params else []:
if isinstance(param, (tuple, list, set)) and len(param) == 0:
- logging.warning("param is empty. not executing query")
+ logger.warning("param is empty. not executing query")
return []
connection = self.make_connection()
c = connection.cursor()
diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py
index edf90f7..5f9a3f6 100644
--- a/tests/incite/collections/test_df_collection_item_thl_web.py
+++ b/tests/incite/collections/test_df_collection_item_thl_web.py
@@ -250,7 +250,7 @@ class TestDFCollectionItemMethod:
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
delete_df_collection(coll=df_collection)
- df = item.from_mysql()
+ df = item.from_db()
if df_collection.data_type == DFCollectionType.LEDGER:
assert df is None
else:
@@ -260,7 +260,7 @@ class TestDFCollectionItemMethod:
incite_item_factory(user=u1, item=item)
- df = item.from_mysql()
+ df = item.from_db()
assert isinstance(df, pd.DataFrame)
assert not df.empty
assert set(df.columns) == set(df_collection._schema.columns.keys())
@@ -401,7 +401,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods.
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Write the basic archive, the issue is that because it's
@@ -444,7 +444,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Confirm a missing ddf (shouldn't bc of type hint) should
@@ -876,7 +876,7 @@ class TestDFCollectionItemFunctionalTest:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
assert isinstance(ddf, dd.DataFrame)
@@ -946,7 +946,7 @@ class TestDFCollectionItemFunctionalTest:
for item in df_collection.items:
assert not item.has_empty()
- df: pd.DataFrame = item.from_mysql()
+ df: pd.DataFrame = item.from_db()
# We do this check b/c the Ledger returns back None and
# I don't want it to fail when we go to make a ddf
diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py
index b4b5b00..6a7e5c9 100644
--- a/tests/incite/collections/test_df_collection_thl_marketplaces.py
+++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py
@@ -46,7 +46,7 @@ class TestDFCollection_thl_marketplaces:
assert isinstance(data_type, DFCollectionType)
# (1) Can't be totally empty, needs a path...
- with pytest.raises(expected_exception=Exception):
+ with pytest.raises(expected_exception=ValueError):
instance = df_coll()
# (2) Confirm it only needs the archive_path
diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py
index aa738e1..e29ba4c 100644
--- a/tests/managers/thl/test_contest/test_milestone.py
+++ b/tests/managers/thl/test_contest/test_milestone.py
@@ -242,7 +242,7 @@ class TestMilestoneContestUserViews:
def test_list_user_eligible_country(
self,
user_with_wallet: User,
- raffle_contest_factory: Callable[..., Contest],
+ raffle_contest_factory: Callable[..., RaffleContest],
thl_ledger_manager: ThlLedgerManager,
contest_manager: ContestManager,
):