diff options
| author | Max Nanis | 2026-08-28 17:28:10 -0700 |
|---|---|---|
| committer | Max Nanis | 2026-08-28 17:28:10 -0700 |
| commit | 8219bf4814a7aa374a500516f2254f39085f0357 (patch) | |
| tree | 8c5bfd4d01eafdec3e216702a0c94852aa38230e | |
| parent | ba544e2ba31432aad4d2acaba3e1f90c27137ded (diff) | |
| download | generalresearch-8219bf4814a7aa374a500516f2254f39085f0357.tar.gz generalresearch-8219bf4814a7aa374a500516f2254f39085f0357.zip | |
Ruff afternoon
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, ): |
