From c6c439970d2167e7afce5f99fef11ab0126162e2 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 20 Aug 2026 09:30:19 -0700 Subject: postgresql auto creation into fixtures. simple tests for db conn checks. wrapper fixture for InternalHostnam (pydantic has MultiHost without host attr annoyances) --- .../collections/test_df_collection_item_thl_web.py | 26 ++++++------- .../test_df_collection_thl_marketplaces.py | 2 +- .../collections/test_df_collection_thl_web.py | 27 ++++++++----- tests/pytest.ini | 3 +- tests/test_postgres.py | 44 ++++++++++++++++++++++ 5 files changed, 75 insertions(+), 27 deletions(-) create mode 100644 tests/test_postgres.py (limited to 'tests') 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 a858fbe..8b8bcbe 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Generator from datetime import datetime, timedelta, timezone from itertools import product as iter_product from os.path import join as pjoin @@ -12,13 +15,7 @@ from distributed import Client, Scheduler, Worker # noinspection PyUnresolvedReferences from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, gen_cluster, - loop, - loop_in_thread, ) from faker import Faker from pandera.pandas import DataFrameSchema @@ -34,7 +31,6 @@ from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.sql_helper import PostgresDsn -from test_utils.incite.conftest import incite_item_factory, mnt_filepath if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets @@ -56,12 +52,12 @@ unsupported_mock_types = { } -def combo_object(): +def combo_object() -> Generator[str, None, None]: for x in iter_product( df_collections, ["15min", "45min", "1H"], ): - yield x + yield from x class TestDFCollectionItemBase: @@ -199,7 +195,7 @@ class TestDFCollectionItemMethod: client_no_amm, incite_item_factory, delete_df_collection, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): assert 1 + 1 == 2 @@ -768,7 +764,7 @@ class TestDFCollectionItemFunctionalTest: product: Product, incite_item_factory, delete_df_collection, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): from generalresearch.models.thl.user import User @@ -818,7 +814,7 @@ class TestDFCollectionItemFunctionalTest: df_collection_data_type, incite_item_factory, delete_df_collection, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): """A functional test to write some Parquet files for the DFCollection and then confirm that the files get written @@ -866,7 +862,7 @@ class TestDFCollectionItemFunctionalTest: df_collection_data_type, incite_item_factory, delete_df_collection, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): from generalresearch.models.thl.user import User @@ -919,7 +915,7 @@ class TestDFCollectionItemFunctionalTest: product: Product, offset: str, duration: timedelta, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): """Don't allow creating an archive for data that will likely be overwritten or updated @@ -960,7 +956,7 @@ class TestDFCollectionItemFunctionalTest: user: User, offset: str, duration: timedelta, - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, ): delete_df_collection(coll=df_collection) diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 8ce8acc..981f62e 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -28,7 +28,7 @@ def combo_object(): ], ["5min", "6H", "30D"], ): - yield x + yield from x @pytest.mark.parametrize("df_coll, offset", combo_object()) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index c64dac8..b09d44c 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Generator from datetime import datetime from itertools import product from typing import TYPE_CHECKING @@ -11,9 +14,13 @@ from generalresearch.incite.collections import DFCollection, DFCollectionType if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.collections import ( + DFCollectionItem, + DFCollectionType, + ) -def combo_object(): +def combo_object() -> Generator[tuple, None, None]: for x in product( [ DFCollectionType.USER, @@ -25,7 +32,7 @@ def combo_object(): ], ["30min", "1H"], ): - yield x + yield from x @pytest.mark.parametrize( @@ -33,7 +40,9 @@ def combo_object(): ) class TestDFCollection_thl_web: - def test_init(self, df_collection_data_type, offset: str, df_collection): + def test_init( + self, df_collection_data_type: DFCollectionType, offset: str, df_collection + ): assert isinstance(df_collection_data_type, DFCollectionType) assert isinstance(df_collection, DFCollection) @@ -43,12 +52,12 @@ class TestDFCollection_thl_web: ) class TestDFCollection_thl_web_Properties: - def test_items(self, df_collection_data_type, offset: str, df_collection): + def test_items(self, df_collection): assert isinstance(df_collection.items, list) for i in df_collection.items: assert i._collection == df_collection - def test__schema(self, df_collection_data_type, offset: str, df_collection): + def test__schema(self, df_collection): assert isinstance(df_collection._schema, DataFrameSchema) @@ -58,16 +67,16 @@ class TestDFCollection_thl_web_Properties: class TestDFCollection_thl_web_BaseProperties: @pytest.mark.skip - def test__interval_range(self, df_collection_data_type, offset: str, df_collection): + def test__interval_range(self, df_collection): pass - def test_interval_start(self, df_collection_data_type, offset: str, df_collection): + def test_interval_start(self, df_collection): assert isinstance(df_collection.interval_start, datetime) - def test_interval_range(self, df_collection_data_type, offset: str, df_collection): + def test_interval_range(self, df_collection): assert isinstance(df_collection.interval_range, list) - def test_progress(self, df_collection_data_type, offset: str, df_collection): + def test_progress(self, df_collection): assert isinstance(df_collection.progress, pd.DataFrame) diff --git a/tests/pytest.ini b/tests/pytest.ini index d280de0..1a5c089 100644 --- a/tests/pytest.ini +++ b/tests/pytest.ini @@ -1,2 +1 @@ -[pytest] -asyncio_mode = auto \ No newline at end of file +[pytest] \ No newline at end of file diff --git a/tests/test_postgres.py b/tests/test_postgres.py new file mode 100644 index 0000000..92eed71 --- /dev/null +++ b/tests/test_postgres.py @@ -0,0 +1,44 @@ +import socket +import subprocess + +from pydantic import PostgresDsn + +from generalresearch.models.custom_types import InternalHostname +from generalresearch.pg_helper import PostgresConfig + + +def is_port_open(host: InternalHostname, port: int = 5432, timeout: int = 3): + try: + with socket.create_connection((host, port), timeout=timeout): + return True + except (socket.timeout, ConnectionRefusedError, OSError): + return False + + +def can_ping(host: InternalHostname): + return ( + subprocess.call( + ["ping", "-c", "1", str(host)], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + == 0 + ) + + +class TestPostgresDSN: + + def test_ping(self, postgres_instance_host: InternalHostname): + assert can_ping(host=postgres_instance_host) + + def test_port(self, postgres_instance_host: InternalHostname): + assert is_port_open(host=postgres_instance_host) + + def test_conn(self, postgres_instance: PostgresDsn): + config = PostgresConfig( + dsn=postgres_instance, + connect_timeout=1, + statement_timeout=1, + ) + res = config.execute_sql_query(query="SELECT 1;") + assert len(res) == 1 -- cgit v1.2.3 From 77e64ac954a7738b93a85fefde25fd6436e75737 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 20 Aug 2026 11:20:00 -0700 Subject: TestPostgresDjangoCreation, pytz to toml --- .gitignore | 3 +- Jenkinsfile | 41 ------------ generalresearch/managers/leaderboard/__init__.py | 4 +- generalresearch/models/custom_types.py | 14 +++-- pyproject.toml | 1 + test_utils/conftest.py | 79 +++++++++++++++--------- tests/models/custom_types/test_aware_datetime.py | 2 +- tests/test_postgres.py | 26 +++++++- 8 files changed, 91 insertions(+), 79 deletions(-) (limited to 'tests') diff --git a/.gitignore b/.gitignore index c7c1d0b..db79a59 100644 --- a/.gitignore +++ b/.gitignore @@ -8,4 +8,5 @@ generalresearch/resources/brokerage_trust_calculated.csv tests/.env.test .env.* .DS_Store -build/ \ No newline at end of file +build/ +*.egg-info \ No newline at end of file diff --git a/Jenkinsfile b/Jenkinsfile index a684caf..e829ba9 100644 --- a/Jenkinsfile +++ b/Jenkinsfile @@ -32,47 +32,6 @@ pipeline { stages { stage('Setup DB') { - steps { - script { - env.DB_NAME = 'unittest-thl-' + UUID.randomUUID().toString().replace('-', '').take(12) - env.THL_WEB_RW_DB = "postgres://${env.DB_USER}:${env.DB_PASSWORD}@${env.DB_POSTGRESQL_HOST}/${env.DB_NAME}" - env.THL_WEB_RR_DB = env.THL_WEB_RW_DB - env.THL_WEB_RO_DB = env.THL_WEB_RW_DB - echo "Using database: ${env.DB_NAME}" - - env.SPECTRUM_DB_NAME = 'unittest-thl-spectrum-' + UUID.randomUUID().toString().replace('-', '').take(12) - env.SPECTRUM_RW_DB = "mariadb://${env.DB_USER}:${env.DB_PASSWORD}@${env.DB_MARIA_HOST}/${env.SPECTRUM_DB_NAME}" - env.SPECTRUM_RR_DB = env.SPECTRUM_RW_DB - echo "Using database: ${env.SPECTRUM_DB_NAME}" - - env.GRLIQ_DB_NAME = 'unittest-grliq-' + UUID.randomUUID().toString().replace('-', '').take(12) - env.GRLIQ_DB = "postgres://${env.DB_USER}:${env.DB_PASSWORD}@${env.DB_POSTGRESQL_HOST}/${env.GRLIQ_DB_NAME}" - echo "Using database: ${env.GRLIQ_DB_NAME}" - - env.GR_DB_NAME = 'unittest-gr-' + UUID.randomUUID().toString().replace('-', '').take(12) - env.GR_DB = "postgres://${env.DB_USER}:${env.DB_PASSWORD}@${env.DB_POSTGRESQL_HOST}/${env.GR_DB_NAME}" - echo "Using database: ${env.GR_DB_NAME}" - } - - sh """ - PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres < str: InternalHostname = Annotated[str, AfterValidator(validate_hostname)] +class PostgresDict(MultiHostHost): + """The path part of this host, or `None`.""" + + # Databasename + name: str | None + + def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str: # By default, datetimes are serialized with the %f optional. We don't # want that because then the deserialization fails if the datetime diff --git a/pyproject.toml b/pyproject.toml index 79b2382..dd2e649 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,7 @@ dependencies = [ "pytest", "pylibmc", "pymemcache", + "pytz", "redis", "requests", "scipy", diff --git a/test_utils/conftest.py b/test_utils/conftest.py index c2c46c5..9c80065 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -11,9 +11,10 @@ import redis from _pytest.config import Config from dotenv import load_dotenv from pydantic import MariaDBDsn, PostgresDsn, TypeAdapter +from pydantic_core import MultiHostHost from redis import Redis -from generalresearch.models.custom_types import InternalHostname +from generalresearch.models.custom_types import InternalHostname, PostgresDict from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig from generalresearch.sql_helper import SqlHelper @@ -115,32 +116,54 @@ def postgres_instance(settings: "GRLBaseSettings") -> Generator[PostgresDsn]: conn = connect(f"{db_path_connect}/postgres") conn.autocommit = True cur = conn.cursor() - cur.execute(SQL("DROP DATABASE {}").format(Identifier(db_name))) + cur.execute(SQL("DROP DATABASE {} WITH (FORCE)").format(Identifier(db_name))) cur.close() conn.close() -@pytest.fixture -def postgres_instance_host( +@pytest.fixture(scope="session") +def postgres_instance_dict( postgres_instance: PostgresDsn, -) -> Generator[InternalHostname]: - host = postgres_instance.hosts()[0]["host"] +) -> Generator[PostgresDict]: + host = postgres_instance.hosts()[0] assert host is not None + msg = "Must have full Postgres details" + assert host["host"], msg + assert host["username"], msg + assert host["password"], msg + + assert postgres_instance.path + + yield PostgresDict( + username=host["username"], + password=host["password"], + host=host["host"], + name=postgres_instance.path.lstrip("/"), + port=5432, + ) + + +@pytest.fixture(scope="session") +def postgres_instance_host( + postgres_instance_dict: PostgresDict, +) -> Generator[InternalHostname]: adapter = TypeAdapter(InternalHostname) - value = adapter.validate_python(host) + value = adapter.validate_python(postgres_instance_dict["host"]) yield value @pytest.fixture(scope="session") -def django_db_setup(postgres_instance: PostgresDsn) -> Callable[..., None]: +def django_db_factory( + postgres_instance: PostgresDsn, postgres_instance_dict: PostgresDict +) -> Callable[..., PostgresDsn]: import django from django.apps import apps from django.conf import settings as django_settings from django.core.management import call_command - def _inner(): + def _inner(django_project: str = "generalresearch.thl_django"): # 1. Bootstrapping Django settings if not django_settings.configured: @@ -148,40 +171,38 @@ def django_db_setup(postgres_instance: PostgresDsn) -> Callable[..., None]: DATABASES={ "default": { "ENGINE": "django.db.backends.postgresql", - # PostgresDsn stores path as "/dbname" - "NAME": str(postgres_instance.path).lstrip("/"), - "USER": postgres_instance["username"], - "PASSWORD": postgres_instance["password"], - "HOST": postgres_instance["host"], - "PORT": "5432", + "NAME": postgres_instance_dict["name"], + "USER": postgres_instance_dict["username"], + "PASSWORD": postgres_instance_dict["password"], + "HOST": postgres_instance_dict["host"], + "PORT": postgres_instance_dict["port"], } }, INSTALLED_APPS=[ "django.contrib.postgres", "django.contrib.contenttypes", - "generalresearch.thl_django", + django_project, ], ) django.setup() - for model in apps.get_models(): - print(f"Discovered model: {model._meta.label}") + # for model in apps.get_models(): + # print(f"Discovered model: {model._meta.label}") # 2. Run migrations directly during fixture activation call_command("migrate") + # 3. Return the Dsn so the factory gives a way to connect + return postgres_instance + return _inner @pytest.fixture(scope="session") -def thl_web_rr(postgres_instance: PostgresDsn, django_db_setup) -> PostgresConfig: - - # Run Migrations now. - # generalresearch/thl_django - django_db_setup() +def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: return PostgresConfig( - dsn=postgres_instance, + dsn=django_db_factory("generalresearch.thl_django"), connect_timeout=1, statement_timeout=5, ) @@ -193,11 +214,13 @@ def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: @pytest.fixture(scope="session") -def gr_db(postgres_instance: PostgresDsn) -> PostgresConfig: +def gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: - # Run Migrations, somehow pull from other repo... - django_db_setup() - return PostgresConfig(dsn=postgres_instance, connect_timeout=1, statement_timeout=5) + return PostgresConfig( + dsn=django_db_factory("gr_carer"), + connect_timeout=1, + statement_timeout=5, + ) @pytest.fixture(scope="session") diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index a23413c..14d1343 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -4,7 +4,7 @@ from typing import Optional import pytest import pytz -from pydantic import BaseModel, ValidationError, Field +from pydantic import BaseModel, Field, ValidationError from generalresearch.models.custom_types import AwareDatetimeISO diff --git a/tests/test_postgres.py b/tests/test_postgres.py index 92eed71..3b3ddd0 100644 --- a/tests/test_postgres.py +++ b/tests/test_postgres.py @@ -1,9 +1,10 @@ import socket import subprocess +from typing import Callable from pydantic import PostgresDsn -from generalresearch.models.custom_types import InternalHostname +from generalresearch.models.custom_types import InternalHostname, PostgresDict from generalresearch.pg_helper import PostgresConfig @@ -42,3 +43,26 @@ class TestPostgresDSN: ) res = config.execute_sql_query(query="SELECT 1;") assert len(res) == 1 + + +class TestPostgresDjangoCreation: + + def test_ping(self, postgres_instance_dict: PostgresDict): + assert can_ping(host=postgres_instance_dict["host"]) + + def test_django_creation( + self, + django_db_factory: Callable[..., None], + ): + + dsn = django_db_factory() + assert isinstance(dsn, PostgresDsn) + + def test_django_tables(self, thl_web_rw: PostgresConfig): + res = thl_web_rw.execute_sql_query(query=""" + SELECT COUNT(*) + FROM information_schema.tables + WHERE table_schema = 'public'; + """) + assert len(res) == 1 + assert res[0]["count"] == 56 -- cgit v1.2.3 From 0c10237223829818752ab4fc431957c0b22a0e23 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 20 Aug 2026 15:40:48 -0700 Subject: moving create_dummy to conftest files, fastapi to requirements for Typing, more Ruff --- generalresearch/grliq/managers/forensic_data.py | 330 +++++++++-------------- generalresearch/incite/base.py | 170 ++++++------ generalresearch/managers/gr/authentication.py | 79 +++--- generalresearch/managers/gr/business.py | 175 ++++-------- generalresearch/managers/gr/team.py | 93 +++---- generalresearch/managers/leaderboard/__init__.py | 4 +- generalresearch/managers/thl/ipinfo.py | 106 +------- pyproject.toml | 1 + test_utils/managers/gr/__init__.py | 0 test_utils/managers/gr/conftest.py | 151 +++++++++++ test_utils/managers/grliq/__init__.py | 0 test_utils/managers/grliq/conftest.py | 61 +++++ test_utils/managers/thl/__init__.py | 0 test_utils/managers/thl/conftest.py | 112 ++++++++ test_utils/models/conftest.py | 201 +++++++------- tests/incite/test_collection_base.py | 6 +- tests/models/custom_types/test_aware_datetime.py | 5 +- tests/models/thl/test_product.py | 53 ++-- 18 files changed, 811 insertions(+), 736 deletions(-) create mode 100644 test_utils/managers/gr/__init__.py create mode 100644 test_utils/managers/gr/conftest.py create mode 100644 test_utils/managers/grliq/__init__.py create mode 100644 test_utils/managers/grliq/conftest.py create mode 100644 test_utils/managers/thl/__init__.py create mode 100644 test_utils/managers/thl/conftest.py (limited to 'tests') diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index d7e362d..739c520 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -1,11 +1,11 @@ -from datetime import datetime, timezone -from typing import Any, Collection, Dict, List, Optional, Tuple -from uuid import uuid4 +from __future__ import annotations + +from datetime import datetime +from typing import Any, Collection from psycopg import sql from pydantic import NonNegativeInt, PositiveInt -from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA from generalresearch.grliq.models.events import PointerMove, TimingData from generalresearch.grliq.models.forensic_data import GrlIqData from generalresearch.grliq.models.forensic_result import ( @@ -23,58 +23,13 @@ class GrlIqDataManager: def __init__(self, postgres_config: PostgresConfig): self.postgres_config = postgres_config - def create_dummy( - self, - is_attempt_allowed: bool = True, - product_id: Optional[str] = None, - product_user_id: Optional[str] = None, - uuid: Optional[str] = None, - mid: Optional[str] = None, - created_at: Optional[datetime] = None, - ) -> GrlIqData: - """ - Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), - and GrlIqForensicCategoryResult (category_results) - :param is_attempt_allowed: Whether the attempt is allowed. - :param product_id: product_id of user - :param product_user_id: product_user_id of user - :param uuid: uuid for the grliq data record - :param mid: the thl_session:uuid / mid for the attempt. - :return: - """ - import copy - - res: GrlIqData = copy.deepcopy(DUMMY_GRLIQ_DATA[int(is_attempt_allowed)]) - - product_id = product_id or uuid4().hex - product_user_id = product_user_id or uuid4().hex - uuid = uuid or uuid4().hex - mid = mid or uuid4().hex - created_at = created_at or datetime.now(tz=timezone.utc) - - res["data"].product_id = product_id - res["data"].product_user_id = product_user_id - res["data"].uuid = uuid - res["data"].mid = mid - res["data"].created_at = created_at - res["result_data"].uuid = uuid - res["category_result"].uuid = uuid - - return self.create( - iq_data=res["data"], - result_data=res["result_data"], - category_result=res["category_result"], - fraud_score=res["category_result"].fraud_score, - is_attempt_allowed=res["category_result"].is_attempt_allowed(), - ) - def create( self, iq_data: GrlIqData, - result_data: Optional[GrlIqCheckerResults] = None, - category_result: Optional[GrlIqForensicCategoryResult] = None, - fraud_score: Optional[int] = None, - is_attempt_allowed: Optional[bool] = None, + result_data: GrlIqCheckerResults | None = None, + category_result: GrlIqForensicCategoryResult | None = None, + fraud_score: int | None = None, + is_attempt_allowed: bool | None = None, ) -> GrlIqData: data = iq_data.model_dump_sql(exclude={"events", "mouse_events", "timing_data"}) @@ -95,8 +50,7 @@ class GrlIqDataManager: data["fraud_score"] = fraud_score data["is_attempt_allowed"] = is_attempt_allowed - query = sql.SQL( - """ + query = sql.SQL(""" INSERT INTO grliq_forensicdata (uuid, session_uuid, created_at, product_id, product_user_id, country_iso, client_ip, ua_browser_family, ua_browser_version, @@ -112,14 +66,12 @@ class GrlIqDataManager: %(fingerprint)s, %(fraud_score)s, %(is_attempt_allowed)s, %(result_data)s, %(category_result)s) RETURNING id - """ - ) + """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, data) - pk = c.fetchone()["id"] # type: ignore - conn.commit() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, data) + pk = c.fetchone()["id"] # type: ignore + conn.commit() iq_data.id = pk @@ -130,9 +82,9 @@ class GrlIqDataManager: uuid: UUIDStr, result_data: GrlIqCheckerResults, category_result: GrlIqForensicCategoryResult, - fingerprint: Optional[str] = None, - fraud_score: Optional[int] = None, - is_attempt_allowed: Optional[bool] = None, + fingerprint: str | None = None, + fraud_score: int | None = None, + is_attempt_allowed: bool | None = None, ) -> None: data = {"uuid": uuid} data["result_data"] = result_data.model_dump_json(exclude_none=True) @@ -141,8 +93,7 @@ class GrlIqDataManager: data["fraud_score"] = fraud_score data["is_attempt_allowed"] = is_attempt_allowed - query = sql.SQL( - """ + query = sql.SQL(""" UPDATE grliq_forensicdata SET result_data = %(result_data)s, category_result = %(category_result)s, @@ -150,8 +101,7 @@ class GrlIqDataManager: fraud_score = %(fraud_score)s, 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) @@ -161,21 +111,17 @@ class GrlIqDataManager: ) conn.commit() - return None - def update_fingerprint(self, iq_data: GrlIqData) -> None: # We should only run this if we modified the fingerprint algorithm if "fingerprint" in iq_data.__dict__: # make sure it's not cached del iq_data.__dict__["fingerprint"] data = {"uuid": iq_data.uuid, "fingerprint": iq_data.fingerprint} - query = sql.SQL( - """ + query = sql.SQL(""" UPDATE grliq_forensicdata 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) @@ -189,13 +135,11 @@ class GrlIqDataManager: # We should only run this if we structured new fields and want to # back-populate them in the db data = {"id": iq_data.id, "data": iq_data.model_dump_sql()["data"]} - query = sql.SQL( - """ + query = sql.SQL(""" UPDATE grliq_forensicdata 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) @@ -207,7 +151,7 @@ class GrlIqDataManager: def get_data_if_exists( self, forensic_uuid: UUIDStr, load_events: bool = False - ) -> Optional[GrlIqData]: + ) -> GrlIqData | None: try: return self.get_data(forensic_uuid=forensic_uuid, load_events=load_events) except AssertionError: @@ -215,8 +159,8 @@ class GrlIqDataManager: def get_data( self, - forensic_id: Optional[PositiveInt] = None, - forensic_uuid: Optional[UUIDStr] = None, + forensic_id: PositiveInt | None = None, + forensic_uuid: UUIDStr | None = None, load_events: bool = False, ) -> GrlIqData: from generalresearch.grliq.managers.forensic_events import ( @@ -230,8 +174,7 @@ class GrlIqDataManager: # forensic items' session, 2) event_start is closest to the # created_at for this forensic item, and within 1 minute. - query = sql.SQL( - """ + query = sql.SQL(""" SELECT d.id, d.data, e.events, e.mouse_events, t.timing_data FROM grliq_forensicdata d -- Closest event_start within 1 minute @@ -251,16 +194,13 @@ class GrlIqDataManager: ORDER BY e2.id DESC LIMIT 1 ) t ON true - """ - ) + """) else: - query = sql.SQL( - """ + query = sql.SQL(""" SELECT d.id, d.data FROM grliq_forensicdata d - """ - ) + """) if forensic_id is not None: column_name = "id" @@ -275,10 +215,9 @@ class GrlIqDataManager: limit_clause = sql.SQL(" LIMIT 1") q1 = sql.Composed([query, where_clause, limit_clause]) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=q1, params=(param_value,)) - x = c.fetchone() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=q1, params=(param_value,)) + x = c.fetchone() assert x is not None, f"GrlIqDataManager.get_data({forensic_uuid=}) not found" @@ -316,10 +255,10 @@ class GrlIqDataManager: def filter_timing_data( self, - created_between: Tuple[datetime, datetime], - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> List[Dict[str, Any]]: + created_between: tuple[datetime, datetime], + limit: int | None = None, + offset: int | None = None, + ) -> list[dict[str, Any]]: # TODO! created_between used to be marked as Optional, but it would # break the query. Evaluate it's use to determine best behavior. @@ -349,10 +288,9 @@ class GrlIqDataManager: {limit_str} {offset_str}; """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - res: List[Dict[str, Any]] = c.fetchall() # type: ignore + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + res: list[dict[str, Any]] = c.fetchall() # type: ignore for x in res: x["timing_data"] = TimingData.model_validate(x["timing_data"]) @@ -369,47 +307,44 @@ class GrlIqDataManager: # This is used for filtering for other forensic posts with a certain # fingerprint, in this product_id, but NOT for this user. - query = sql.SQL( - """ + query = sql.SQL(""" SELECT COUNT(DISTINCT product_user_id) as user_count FROM grliq_forensicdata d WHERE product_id = %(product_id)s AND fingerprint = %(fingerprint)s AND product_user_id != %(product_user_id)s AND created_at > NOW() - INTERVAL '30 DAYS' - """ - ) + """) params = { "product_id": product_id, "fingerprint": fingerprint, "product_user_id": product_user_id_not, } # print(query) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - user_count = c.fetchone()["user_count"] # type: ignore + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + user_count = c.fetchone()["user_count"] # type: ignore return int(user_count) def filter_data( self, - session_uuid: Optional[str] = None, - fingerprint: Optional[str] = None, - fingerprints: Optional[Collection[str]] = None, - product_id: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - uuids: Optional[Collection[str]] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - users: Optional[Collection[User]] = None, - phase: Optional[Phase] = None, + session_uuid: str | None = None, + fingerprint: str | None = None, + fingerprints: Collection[str] | None = None, + product_id: str | None = None, + product_ids: Collection[str] | None = None, + uuids: Collection[str] | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + users: Collection[User] | None = None, + phase: Phase | None = None, order_by: str = "created_at DESC", - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> List[GrlIqData]: + limit: int | None = None, + offset: int | None = None, + ) -> list[GrlIqData]: res = self.filter( select_str="d.id, d.data", @@ -433,18 +368,18 @@ class GrlIqDataManager: def filter_results( self, - session_uuid: Optional[str] = None, - uuid: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - product_id: Optional[str] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - limit: Optional[int] = None, - offset: Optional[int] = None, + session_uuid: str | None = None, + uuid: str | None = None, + product_ids: Collection[str] | None = None, + product_id: str | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + limit: int | None = None, + offset: int | None = None, order_by: str = "created_at DESC", - ) -> List[GrlIqCheckerResults]: + ) -> list[GrlIqCheckerResults]: select_str = ( "id, session_uuid, product_id, product_user_id, created_at, result_data" ) @@ -472,18 +407,18 @@ class GrlIqDataManager: def filter_category_results( self, - session_uuid: Optional[str] = None, - uuid: Optional[str] = None, - product_id: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, + session_uuid: str | None = None, + uuid: str | None = None, + product_id: str | None = None, + product_ids: Collection[str] | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, order_by: str = "created_at DESC", - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> List[GrlIqForensicCategoryResult]: + limit: int | None = None, + offset: int | None = None, + ) -> list[GrlIqForensicCategoryResult]: select_str = ( "id, session_uuid, product_id, product_user_id, created_at, category_result" ) @@ -506,22 +441,22 @@ class GrlIqDataManager: @staticmethod def make_filter_str( - session_uuid: Optional[str] = None, - fingerprint: Optional[str] = None, - fingerprints: Optional[Collection[str]] = None, - uuids: Optional[Collection[str]] = None, - product_id: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - users: Optional[Collection[User]] = None, - phase: Optional[Phase] = None, - ) -> Tuple[str, Dict[str, Any]]: + session_uuid: str | None = None, + fingerprint: str | None = None, + fingerprints: Collection[str] | None = None, + uuids: Collection[str] | None = None, + product_id: str | None = None, + product_ids: Collection[str] | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + users: Collection[User] | None = None, + phase: Phase | None = None, + ) -> tuple[str, dict[str, Any]]: filters = [] - params: Dict[str, Any] = {} + params: dict[str, Any] = {} if session_uuid: params["session_uuid"] = session_uuid @@ -614,19 +549,20 @@ class GrlIqDataManager: def filter_count( self, - session_uuid: Optional[str] = None, - fingerprint: Optional[str] = None, - fingerprints: Optional[Collection[str]] = None, - uuids: Optional[Collection[str]] = None, - product_id: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - users: Optional[Collection[User]] = None, - phase: Optional[Phase] = None, + session_uuid: str | None = None, + fingerprint: str | None = None, + fingerprints: Collection[str] | None = None, + uuids: Collection[str] | None = None, + product_id: str | None = None, + product_ids: Collection[str] | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + users: Collection[User] | None = None, + phase: Phase | None = None, ) -> NonNegativeInt: + filter_str, params = self.make_filter_str( session_uuid=session_uuid, fingerprint=fingerprint, @@ -682,36 +618,35 @@ class GrlIqDataManager: FROM grliq_forensicdata d {filter_str} """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - res = c.fetchone() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + res = c.fetchone() return int(res["c"]) def filter( self, select_str: str, - session_uuid: Optional[str] = None, - fingerprint: Optional[str] = None, - fingerprints: Optional[Collection[str]] = None, - uuids: Optional[Collection[str]] = None, - product_id: Optional[str] = None, - product_ids: Optional[Collection[str]] = None, - created_after: Optional[datetime] = None, - created_before: Optional[datetime] = None, - created_between: Optional[Tuple[datetime, datetime]] = None, - user: Optional[User] = None, - users: Optional[Collection[User]] = None, - phase: Optional[Phase] = None, + session_uuid: str | None = None, + fingerprint: str | None = None, + fingerprints: Collection[str] | None = None, + uuids: Collection[str] | None = None, + product_id: str | None = None, + product_ids: Collection[str] | None = None, + created_after: datetime | None = None, + created_before: datetime | None = None, + created_between: tuple[datetime, datetime] | None = None, + user: User | None = None, + users: Collection[User] | None = None, + phase: Phase | None = None, order_by: str = "created_at DESC", - limit: Optional[int] = None, - offset: Optional[int] = None, - ) -> List[Dict[str, Any]]: + limit: int | None = None, + offset: int | None = None, + ) -> list[dict[str, Any]]: """ Accepts lots of optional filters. """ if not limit: - limit = 5000 + limit = 5_000 if not offset: offset = 0 @@ -746,10 +681,9 @@ class GrlIqDataManager: OFFSET {offset} """ # print(query) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - res: List[Dict[str, Any]] = c.fetchall() # type: ignore + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + res: list[dict[str, Any]] = c.fetchall() # type: ignore for x in res: @@ -777,7 +711,7 @@ class GrlIqDataManager: return res @staticmethod - def temporary_add_missing_fields(d: Dict[str, Any]) -> None: + def temporary_add_missing_fields(d: dict[str, Any]) -> None: # The following fields were added recently, and so we must give them # a value or old db rows won't be parseable. Once logs are backfilled # then this can be removed diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index aa64bf0..a8088ac 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -18,11 +18,7 @@ from typing import ( TYPE_CHECKING, Any, Callable, - List, - Optional, Sequence, - Tuple, - Union, ) from uuid import uuid4 @@ -30,7 +26,7 @@ import dask import dask.dataframe as dd import pandas as pd import pyarrow.parquet as pq -from distributed import Client +from distributed import Client as DaskClient from pandera.pandas import DataFrameSchema from pydantic import ( BaseModel, @@ -38,7 +34,9 @@ from pydantic import ( DirectoryPath, Field, FilePath, + PositiveInt, PrivateAttr, + TypeAdapter, ValidationInfo, field_validator, model_validator, @@ -61,7 +59,7 @@ if TYPE_CHECKING: ) from generalresearch.incite.mergers import MergeCollection, MergeType - Collection = Union[DFCollection, MergeCollection] + Collection = DFCollection | MergeCollection logging.basicConfig() LOG = logging.getLogger() @@ -71,6 +69,9 @@ Item = Any Items = Sequence[Item] DT_STR = "%Y-%m-%d %H:%M:%S" +_dir_adapter = TypeAdapter(DirectoryPath) +_filepath_adapter = TypeAdapter(FilePath) + class NFSMount(BaseModel): address: str = Field(default="127.0.0.1") @@ -89,8 +90,8 @@ class GRLDatasets(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) - data_src: Optional[Path] = Field(default=None) - incite: Optional[NFSMount] = Field(default=None) + data_src: Path | None = Field(default=None) + incite: NFSMount | None = Field(default=None) @model_validator(mode="after") def check_data_src_and_et_path(self) -> Self: @@ -99,6 +100,8 @@ class GRLDatasets(BaseModel): ) from generalresearch.incite.mergers import MergeType + assert self.data_src, "data src must be defined" + # Create the base folders and confirm we have read access self.data_src.mkdir(parents=True, exist_ok=True) assert access( @@ -121,7 +124,7 @@ class GRLDatasets(BaseModel): assert access(path=p, mode=R_OK), f"Cannot read {p}" return self - def archive_path(self, enum_type: Union[MergeType, DFCollectionType]) -> Path: + def archive_path(self, enum_type: MergeType | DFCollectionType) -> Path: """ TODO: Extend this so that it takes any type of Enum and that inputs in the correct parent dir for the respective Enum @@ -135,7 +138,7 @@ class GRLDatasets(BaseModel): pjoin(self.data_src, self.incite.point, folder, str(enum_type.value)) ) - def has_data(self, enum_type: Union[MergeType, DFCollectionType]) -> bool: + def has_data(self, enum_type: MergeType | DFCollectionType) -> bool: path_dir = self.archive_path(enum_type=enum_type) if isdir(path_dir): return bool(listdir(path_dir)) @@ -152,7 +155,7 @@ class CollectionBase(BaseModel): extra="forbid", ) - archive_path: DirectoryPath = Field(default="/tmp/") + archive_path: DirectoryPath = Field(default=_dir_adapter.validate_python("/tmp/")) df: SkipJsonSchema[pd.DataFrame] = Field( default_factory=lambda: pd.DataFrame(), exclude=True ) @@ -169,12 +172,12 @@ class CollectionBase(BaseModel): frozen=True, ) - finished: Optional[AwareDatetimeISO] = Field( + finished: AwareDatetimeISO | None = Field( default=None, description="Finished is only set if we don't want a rolling window", ) - _client: Optional[Client] = PrivateAttr(default=None) + _client: DaskClient | None = PrivateAttr(default=None) # --- Validators --- @model_validator(mode="before") @@ -188,15 +191,14 @@ class CollectionBase(BaseModel): assert isinstance(ap, Path), "check_model_before.isinstance(ap, Path)" if not ap.is_dir(): - raise ValueError(f"Path does not point to a directory") + raise ValueError("Path does not point to a directory") if not access(path=ap, mode=R_OK): - raise ValueError(f"Cannot read archive_path") + raise ValueError("Cannot read archive_path") - df: Optional[pd.DataFrame] = data.get("df", None) - if df is not None: - if not df.empty or len(df.columns) != 0: - raise ValueError("Do not provide a pd.DataFrame") + df: pd.DataFrame | None = data.get("df", None) + if df is not None and (not df.empty or len(df.columns) != 0): + raise ValueError("Do not provide a pd.DataFrame") return data @@ -215,21 +217,21 @@ class CollectionBase(BaseModel): @field_validator("start") def check_start( - cls, start: Optional[datetime], info: ValidationInfo - ) -> Optional[datetime]: + cls, start: datetime | None, info: ValidationInfo + ) -> datetime | None: if start and start.microsecond != 0: raise ValueError("Collection.start must not have microseconds") return start @field_validator("offset") - def check_offset(cls, v: Optional[str], info: ValidationInfo): + def check_offset(cls, v: str | None, info: ValidationInfo): # pd.offsets.__all__ if v is None: # In MergeCollections, offset can be None return v try: pd.Timedelta(v) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) raise ValueError( "Invalid offset alias provided. Please review: " @@ -251,6 +253,7 @@ class CollectionBase(BaseModel): assert end, "an end value must be provided" _start = self.interval_start + assert _start, "a start value must be provided" if end.tzinfo is None: # A Naive end was passed in. We probably did this on purpose. @@ -283,13 +286,13 @@ class CollectionBase(BaseModel): ) @property - def interval_start(self) -> Optional[datetime]: + def interval_start(self) -> datetime | None: # In DFCollections, start must be set, so the interval_start = start. In merged # this may be overridden with different behavior. return self.start @property - def interval_range(self) -> List[Tuple]: + def interval_range(self) -> list[tuple[datetime, datetime]]: """closed='left', so 0 <= x < 5""" end = self.finished or datetime.now(tz=timezone.utc).replace(microsecond=0) iv_r = self._interval_range(end) @@ -302,7 +305,7 @@ class CollectionBase(BaseModel): return pd.DataFrame.from_records(records, index=self._interval_range(end)) @property - def items(self) -> pd.DataFrame: + def items(self) -> Items | None: raise NotImplementedError("Must override") @property @@ -315,19 +318,20 @@ class CollectionBase(BaseModel): def fetch_all_paths( self, - items: Optional[Items] = None, - force_rr_latest=False, - include_partial=False, - ) -> List[FilePath]: + items: Items | None = None, + force_rr_latest: bool = False, + include_partial: bool = False, + ) -> list[FilePath]: LOG.info( f"CollectionBase.fetch_all(items={len(items or [])}, " f"{force_rr_latest=}, {include_partial=})" ) items = items or self.items + assert items # (1) All the originally available archives - sources: List[FilePath] = [ + sources: list[FilePath] = [ i.path for i in items if i.has_archive(include_empty=False) ] @@ -357,14 +361,14 @@ class CollectionBase(BaseModel): def ddf( self, - items: Optional[Items] = None, - force_rr_latest=False, + items: Items | None = None, + force_rr_latest: bool = False, columns=None, filters=None, categories=None, include_partial=False, - graph: Optional[Callable] = None, - ) -> Optional[dd.DataFrame]: + graph: Callable | None = None, + ) -> dd.DataFrame | None: """ Args: @@ -396,7 +400,7 @@ class CollectionBase(BaseModel): """ if isinstance(items, list) and len(items): - sources: List[FilePath] = [ + sources: list[FilePath] = [ i.path for i in items if i.has_archive(include_empty=False) ] @@ -410,7 +414,7 @@ class CollectionBase(BaseModel): ) else: - sources: List[FilePath] = self.fetch_all_paths( + sources: list[FilePath] = self.fetch_all_paths( items=None, force_rr_latest=force_rr_latest, include_partial=include_partial, @@ -444,8 +448,8 @@ class CollectionBase(BaseModel): # --- Methods: Cleanup --- def schedule_cleanup( - self, client=None, sync=True, client_resources=None - ) -> Union[pd.DataFrame, Future]: + self, client: DaskClient | None = None, sync: bool = True, client_resources=None + ) -> pd.DataFrame | Future: LOG.info(f"cleanup(archive_path={self.archive_path})") fs = [] @@ -453,6 +457,8 @@ class CollectionBase(BaseModel): fs.append(dask.delayed(item.cleanup_partials)()) fs.append(dask.delayed(item.clear_corrupt_archive)()) fs.append(dask.delayed(self.clear_tmp_archives)()) + + assert isinstance(client, DaskClient) res = client.compute( collections=fs, sync=sync, @@ -468,15 +474,13 @@ class CollectionBase(BaseModel): self.clear_corrupt_archives() # self.check_empty() # what did this do?? - return None - def cleanup_partials(self) -> None: """If an item is "closed", remove any partial files that may be around...""" + assert self.items + for item in self.items: item.cleanup_partials() - return None - def clear_tmp_archives(self) -> None: regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.I) @@ -487,19 +491,18 @@ class CollectionBase(BaseModel): Path(os.path.join(self.archive_path, fn)) ) - return None - def clear_corrupt_archives(self) -> None: + assert self.items + for item in self.items: item.clear_corrupt_archive() - return None - def rebuild_symlinks(self) -> None: """ When copying "things" between filesystems, and using Sambda mmfsylinks, we can't ensure links are properly shared. """ + assert self.items for item in reversed(self.items): item: CollectionItemBase @@ -513,7 +516,6 @@ class CollectionBase(BaseModel): # Don't "continue" onto the next CollectionItem. Later on, # we may need to create a symlink for the most recent partial - pass # --- Empty Path --- if os.path.exists(empty_path): @@ -574,7 +576,7 @@ class CollectionBase(BaseModel): # `ln` command is run. -- Max 2024-07-26 try: os.remove(item.path.as_posix()) - except FileNotFoundError as e: + except FileNotFoundError: pass if platform == "darwin": @@ -582,16 +584,20 @@ class CollectionBase(BaseModel): else: subprocess.call(["ln", "-sfnT", highest_version, item.path.as_posix()]) - return None - # -- Methods: Source timing def get_item(self, interval: pd.Interval) -> Item: + assert self.items + return next(x for x in self.items if x.interval == interval) def get_item_start(self, start: pd.Timestamp) -> Items: + assert self.items + return next(x for x in self.items if x.interval.left == start) def get_items(self, since: datetime) -> Items: + assert self.items + res = [] first_match = True @@ -610,7 +616,7 @@ class CollectionBase(BaseModel): res.append(item) first_match = False - res: List[Item] = [i for i in res if not i.is_empty()] + res: list[Item] = [i for i in res if not i.is_empty()] if len([1 for i in res if i.should_archive() and not i.has_archive()]): warnings.warn( message="DFCollectionItem has missing archives", @@ -620,7 +626,7 @@ class CollectionBase(BaseModel): return res def get_items_from_year(self, year: int) -> Items: - ts = datetime(year=year, month=1, day=1) + ts = datetime(year=year, month=1, day=1, tzinfo=timezone.utc) return self.get_items(since=ts) def get_items_last90(self) -> Items: @@ -645,10 +651,14 @@ class CollectionItemBase(BaseModel): @property def name(self) -> str: coll = self._collection + if hasattr(coll, "data_type"): + assert coll.data_type name = coll.data_type.value else: + assert coll.merge_type name = coll.merge_type.value + return name def __str__(self): @@ -702,7 +712,9 @@ class CollectionItemBase(BaseModel): @property def path(self) -> FilePath: - return FilePath(os.path.join(self._collection.archive_path, self.filename)) + return_filepath_adapter.validate_python( + os.path.join(self._collection.archive_path, self.filename) + ) @property def partial_path(self) -> FilePath: @@ -732,7 +744,7 @@ class CollectionItemBase(BaseModel): # We assume the target ends with ".####". If not, we'll append .00000 try: - left, right = target.rsplit(".", 1) + _, right = target.rsplit(".", 1) right_int = int(right) except ValueError: return Path(f"{path}.{0:>05}") @@ -740,7 +752,7 @@ class CollectionItemBase(BaseModel): right_int += 1 return Path(f"{path}.{right_int:>05}") - def search_highest_numbered_path(self) -> Optional[Path]: + def search_highest_numbered_path(self) -> Path | None: """This is used for when things are broken, and we want to rebuild our symlinks. We can't trust or use any exist symlinks... so given a path or a partial path... find the highest available "versioned" @@ -766,7 +778,7 @@ class CollectionItemBase(BaseModel): # nums = sorted([b.rsplit(".", 1)[1] for b in builds], reverse=True) # return Path(f"{self.path}.{nums[0]}") - files: List[str] = sorted( + files: list[str] = sorted( builds, key=lambda b: b.rsplit(".", 1)[1], reverse=True ) return Path(os.path.join(coll.archive_path, files[0])) @@ -818,7 +830,6 @@ class CollectionItemBase(BaseModel): shutil.rmtree(generic_path) else: LOG.warning(f"tried removing non-existent file: {generic_path}") - pass def should_archive(self) -> bool: # Determine if enough time has passed to move out of a partial file into an @@ -828,9 +839,7 @@ class CollectionItemBase(BaseModel): if archive_after is None: return False - if datetime.now(tz=timezone.utc) > self.finish + archive_after: - return True - return False + return datetime.now(tz=timezone.utc) > self.finish + archive_after def set_empty(self): assert ( @@ -842,8 +851,8 @@ class CollectionItemBase(BaseModel): def valid_archive( self, - generic_path: Optional[FilePath] = None, - sample: Optional[int] = None, + generic_path: FilePath | None = None, + sample: int | None = None, ) -> bool: """ Attempts to confirm if the parquet file or directory that is @@ -864,7 +873,7 @@ class CollectionItemBase(BaseModel): raise ValueError("Unknown path type.") df = parquet.read().to_pandas() - except (Exception,): + except Exception: LOG.warning(f"Invalid archive {path=}") df = None @@ -876,8 +885,8 @@ class CollectionItemBase(BaseModel): return self.validate_df(df=df, sample=sample) is not None def validate_df( - self, df: pd.DataFrame, sample: Optional[int] = None - ) -> Optional[pd.DataFrame]: + self, df: pd.DataFrame, sample: int | None = None + ) -> pd.DataFrame | None: if sample is not None: sample = min(len(df), sample) try: @@ -910,8 +919,8 @@ class CollectionItemBase(BaseModel): def from_archive( self, include_empty: bool = True, - generic_path: Optional[FilePath] = None, - ) -> Optional[dd.DataFrame]: + generic_path: FilePath | None = None, + ) -> dd.DataFrame | None: if include_empty and self.path_exists(generic_path=self.empty_path): # Return an empty dd.DataFrame with the correct columns @@ -930,15 +939,15 @@ class CollectionItemBase(BaseModel): raise NotImplementedError("Must override") # --- ORM / Data handlers--- - def _to_dict(self, *args, **kwargs) -> dict: - return dict( - should_archive=self.should_archive(), - has_archive=self.has_archive(), - filename=self.filename, - path=self.path, - start=self.start, - finish=self.finish, - ) + def _to_dict(self) -> dict[str, Any]: + return { + "should_archive": self.should_archive(), + "has_archive": self.has_archive(), + "filename": self.filename, + "path": self.path, + "start": self.start, + "finish": self.finish, + } def delete_partial(self): # If a Collection Item is archived, we want to delete the partial file. @@ -961,11 +970,16 @@ class CollectionItemBase(BaseModel): else: self.delete_dangling_partials(keep_latest=2) - def delete_dangling_partials(self, keep_latest=None, target_path=None) -> List[str]: + def delete_dangling_partials( + self, + keep_latest: PositiveInt | None = None, + target_path: Path | str | None = None, + ) -> list[str]: # Specifically looking for numbered partials that are NOT associated # with a symlink. It does not matter if the item is archiveable or not. if target_path is None: target_path = self.partial_path + fps = glob.glob(target_path.as_posix() + ".*") fps = {x for x in fps if x.split(".")[-1].isnumeric()} # Note: if the dir itself is sym-linked, this is going to be wrong. diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index f4185b2..a402693 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -5,7 +5,6 @@ import logging import os from datetime import datetime, timezone from typing import TYPE_CHECKING, Any -from uuid import uuid4 from psycopg import sql from pydantic import AnyHttpUrl, PositiveInt @@ -23,18 +22,6 @@ if TYPE_CHECKING: class GRUserManager(PostgresManagerWithRedis): - def create_dummy( - self, - sub: str | None = None, - is_superuser: bool = False, - ) -> GRUser: - sub = sub or f"{uuid4().hex}-{uuid4().hex}" - - return self.create( - sub=sub, - is_superuser=is_superuser, - ) - def create( self, sub: str, @@ -71,18 +58,17 @@ class GRUserManager(PostgresManagerWithRedis): def get_by_id(self, gr_user_id: int) -> GRUser | None: from generalresearch.models.gr.authentication import GRUser - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" SELECT u.* FROM gr_user AS u WHERE u.id = %s LIMIT 1; """, - params=(gr_user_id,), - ) - res = c.fetchone() + params=(gr_user_id,), + ) + res = c.fetchone() if res is None: raise ValueError("GRUser not found") @@ -99,18 +85,17 @@ class GRUserManager(PostgresManagerWithRedis): def get_by_sub(self, sub: str, raises=True) -> GRUser | None: from generalresearch.models.gr.authentication import GRUser - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" SELECT u.* FROM gr_user AS u WHERE u.sub = %s LIMIT 1; """, - params=(sub,), - ) - res = c.fetchone() + params=(sub,), + ) + res = c.fetchone() if raises and res is None: raise ValueError("GRUser not found") @@ -134,32 +119,30 @@ class GRUserManager(PostgresManagerWithRedis): def get_all(self) -> list[GRUser]: from generalresearch.models.gr.authentication import GRUser - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=""" SELECT u.* FROM gr_user AS u """) - res = c.fetchall() + res = c.fetchall() return [GRUser.from_postgresql(i) for i in res] def get_by_team(self, team_id: PositiveInt) -> list[GRUser]: from generalresearch.models.gr.authentication import GRUser - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" SELECT gru.* FROM common_membership AS membership INNER JOIN gr_user AS gru ON gru.id = membership.user_id WHERE membership.team_id = %s """, - params=(team_id,), - ) - res = c.fetchall() + params=(team_id,), + ) + res = c.fetchall() for item in res: for k, v in item.items(): @@ -176,7 +159,7 @@ class GRUserManager(PostgresManagerWithRedis): return None res = thl_pg_config.execute_sql_query( - query=f""" + query=""" SELECT bp.id FROM userprofile_brokerageproduct AS bp WHERE bp.business_id = ANY(%s) @@ -240,16 +223,15 @@ class GRTokenManager(PostgresManager): return gr_token # API Key - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - query = sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + query = sql.SQL(""" SELECT grk.* FROM gr_token AS grk WHERE grk.key = %s LIMIT 1 """) - c.execute(query=query, params=(api_key,)) - res = c.fetchall() + c.execute(query=query, params=(api_key,)) + res = c.fetchall() if len(res) == 0: raise Exception(f"No GRUser with token of '{api_key}'") @@ -295,9 +277,8 @@ class GRTokenManager(PostgresManager): # therefore, this will only return 0 or 1 GRTokens from generalresearch.models.gr.authentication import GRToken - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - query = sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + query = sql.SQL(""" SELECT grt.* FROM gr_token AS grt LEFT JOIN gr_user AS u @@ -306,9 +287,9 @@ class GRTokenManager(PostgresManager): LIMIT 1; """) - c.execute(query=query, params=(user_id,)) + c.execute(query=query, params=(user_id,)) - result = c.fetchall() + result = c.fetchall() if not result: return None diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py index aa440fb..4da0e7f 100644 --- a/generalresearch/managers/gr/business.py +++ b/generalresearch/managers/gr/business.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, List +from typing import TYPE_CHECKING from uuid import UUID, uuid4 from psycopg import sql @@ -26,28 +26,6 @@ if TYPE_CHECKING: class BusinessBankAccountManager(PostgresManager): - def create_dummy( - self, - business_id: PositiveInt, - uuid: UUIDStr | None = None, - transfer_method: TransferMethod | None = None, - account_number: str | None = None, - routing_number: str | None = None, - iban: str | None = None, - swift: str | None = None, - ): - from generalresearch.models.gr.business import TransferMethod - - return self.create( - business_id=business_id, - uuid=uuid or uuid4().hex, - transfer_method=transfer_method or TransferMethod.ACH, - account_number=account_number or uuid4().hex[:6], - routing_number=routing_number or uuid4().hex[:6], - iban=iban or uuid4().hex[:6], - swift=swift or uuid4().hex[:6], - ) - def create( self, business_id: PositiveInt, @@ -94,59 +72,25 @@ class BusinessBankAccountManager(PostgresManager): ba.id = ba_id return ba - def get_by_business_id(self, business_id: UUIDStr) -> List[BusinessBankAccount]: + def get_by_business_id(self, business_id: UUIDStr) -> list[BusinessBankAccount]: from generalresearch.models.gr.business import BusinessBankAccount - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT ba.* FROM common_bankaccount AS ba WHERE ba.business_id = %s """), - params=(business_id,), - ) - res = c.fetchall() + params=(business_id,), + ) + res = c.fetchall() return [BusinessBankAccount.model_validate(item) for item in res] class BusinessAddressManager(PostgresManager): - def create_dummy( - self, - business_id: PositiveInt, - uuid: UUIDStr | None = None, - line_1: str | None = None, - line_2: str | None = None, - city: str | None = None, - state: str | None = None, - postal_code: str | None = None, - phone_number: PhoneNumber | None = None, - country: str | None = None, - ): - uuid = uuid or uuid4().hex - line_1 = line_1 or "abc" - line_2 = line_2 or "bczx" - city = city or "Downingtown" - state = state or "CA" - postal_code = postal_code or "94041" - phone_number = None - country = country or "US" - - return self.create( - business_id=business_id, - uuid=uuid, - line_1=line_1, - line_2=line_2, - city=city, - state=state, - postal_code=postal_code, - phone_number=phone_number, - country=country, - ) - def create( self, business_id: PositiveInt, @@ -237,24 +181,6 @@ class BusinessManager(PostgresManagerWithRedis): uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number ) - def create_dummy( - self, - uuid: UUIDStr | None = None, - name: str | None = None, - team: Team | None = None, - kind: BusinessType | None = None, - tax_number: str | None = None, - ) -> Business: - from random import randint - - uuid = uuid or uuid4().hex - name = name or "< Unknown >" - tax_number = tax_number or str(randint(1, 999_999_999)) - - return self.create( - uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number - ) - def create( self, name: str, @@ -316,13 +242,12 @@ class BusinessManager(PostgresManagerWithRedis): """ from generalresearch.models.gr.business import Business - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=sql.SQL(""" SELECT b.id, b.uuid, b.kind, b.name, b.tax_number FROM common_business AS b """)) - res = c.fetchall() + res = c.fetchall() response = [] for i in res: @@ -341,20 +266,19 @@ class BusinessManager(PostgresManagerWithRedis): ) -> list[Business]: # conn: psycopg.Connection = GR_POSTGRES_C.make_connection() - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT b.id, b.uuid, b.kind, b.name, b.tax_number FROM common_business AS b INNER JOIN common_team_businesses as tb ON tb.business_id = b.id WHERE tb.team_id = %s """), - params=(team_id,), - ) + params=(team_id,), + ) - res = c.fetchall() + res = c.fetchall() response = [] from generalresearch.models.gr.business import Business @@ -372,10 +296,9 @@ class BusinessManager(PostgresManagerWithRedis): ) -> list[Business]: from generalresearch.models.gr.business import Business - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT b.id, b.uuid, b.kind, b.name, b.tax_number FROM common_business AS b INNER JOIN common_team_businesses AS tb @@ -384,10 +307,10 @@ class BusinessManager(PostgresManagerWithRedis): ON m.team_id = tb.team_id WHERE m.user_id = %s """), - params=(user_id,), - ) + params=(user_id,), + ) - res = c.fetchall() + res = c.fetchall() response = [] for i in res: @@ -402,10 +325,9 @@ class BusinessManager(PostgresManagerWithRedis): :return: Every Business UUIDStr that this GRUser has permission to view """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT b.id FROM common_business AS b INNER JOIN common_team_businesses AS tb @@ -414,10 +336,10 @@ class BusinessManager(PostgresManagerWithRedis): ON tb.team_id = cm.team_id WHERE cm.user_id = %s """), - params=(user_id,), - ) + params=(user_id,), + ) - res = c.fetchall() + res = c.fetchall() return [i["id"] for i in res] @@ -426,10 +348,9 @@ class BusinessManager(PostgresManagerWithRedis): :return: Every Business UUIDStr that this GRUser has permission to view """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT b.uuid FROM common_business AS b INNER JOIN common_team_businesses AS tb @@ -438,10 +359,10 @@ class BusinessManager(PostgresManagerWithRedis): ON tb.team_id = cm.team_id WHERE cm.user_id = %s """), - params=(user_id,), - ) + params=(user_id,), + ) - res = c.fetchall() + res = c.fetchall() return [i["uuid"] for i in res] @@ -453,19 +374,18 @@ class BusinessManager(PostgresManagerWithRedis): assert UUID(hex=business_uuid).hex == business_uuid - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT id, uuid, kind, name, tax_number FROM common_business WHERE uuid = %s LIMIT 1; """), - params=(business_uuid,), - ) + params=(business_uuid,), + ) - res = c.fetchall() + res = c.fetchall() if len(res) == 0: return None @@ -481,19 +401,18 @@ class BusinessManager(PostgresManagerWithRedis): assert isinstance(business_id, int) - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT id, uuid, kind, name, tax_number FROM common_business WHERE id = %s LIMIT 1; """), - params=(business_id,), - ) + params=(business_id,), + ) - res = c.fetchall() + res = c.fetchall() if len(res) == 0: return None diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index ecb1ba4..6de82b0 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -84,19 +84,18 @@ class MembershipManager(PostgresManager): def exists( self, gr_user_id: PositiveInt, team_id: PositiveInt ) -> Membership | None: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT id, uuid, privilege, owner, created, user_id, team_id FROM common_membership WHERE team_id = %s AND user_id = %s LIMIT 1 """), - params=(team_id, gr_user_id), - ) - res = c.fetchone() + params=(team_id, gr_user_id), + ) + res = c.fetchone() if not res: return None @@ -104,36 +103,34 @@ class MembershipManager(PostgresManager): return Membership.model_validate(res) def get_by_team_id(self, team_id: PositiveInt) -> list[Membership]: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT id, uuid, privilege, owner, created, user_id, team_id FROM common_membership WHERE team_id = %s LIMIT 250 """), - params=(team_id,), - ) - res = c.fetchall() + params=(team_id,), + ) + res = c.fetchall() return [Membership.model_validate(i) for i in res] def get_by_gr_user_id(self, gr_user_id: PositiveInt) -> list[Membership]: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT id, uuid, privilege, owner, created, user_id, team_id FROM common_membership WHERE user_id = %s LIMIT 250 """), - params=(gr_user_id,), - ) - res = c.fetchall() + params=(gr_user_id,), + ) + res = c.fetchall() return [Membership.model_validate(i) for i in res] @@ -154,24 +151,15 @@ class TeamManager(PostgresManagerWithRedis): def get_all(self) -> list[Team]: from generalresearch.models.gr.team import Team - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=sql.SQL(""" SELECT t.id, t.uuid, t.name FROM common_team AS t """)) - res = c.fetchall() + res = c.fetchall() return [Team.model_validate(i) for i in res] - def create_dummy( - self, uuid: UUIDStr | None = None, name: str | None = None - ) -> Team: - uuid = uuid or uuid4().hex - name = name or f"name-{uuid4().hex[:12]}" - - return self.create(uuid=uuid, name=name) - def create( self, name: str, @@ -228,19 +216,18 @@ class TeamManager(PostgresManagerWithRedis): def get_by_uuid(self, team_uuid: UUIDStr) -> Team | None: from generalresearch.models.gr.team import Team - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" SELECT t.* FROM common_team AS t WHERE t.uuid = %s LIMIT 1; """, - params=(team_uuid,), - ) + params=(team_uuid,), + ) - res = c.fetchone() + res = c.fetchone() if not isinstance(res, dict): return None @@ -250,19 +237,18 @@ class TeamManager(PostgresManagerWithRedis): def get_by_id(self, team_id: PositiveInt) -> Team | None: from generalresearch.models.gr.team import Team - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT t.id, t.uuid, t.name FROM common_team AS t WHERE t.id = %s LIMIT 1; """), - params=(team_id,), - ) + params=(team_id,), + ) - res = c.fetchone() + res = c.fetchone() if not isinstance(res, dict): return None @@ -272,19 +258,18 @@ class TeamManager(PostgresManagerWithRedis): def get_by_user(self, gr_user: GRUser) -> list[Team]: from generalresearch.models.gr.team import Team - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=sql.SQL(""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=sql.SQL(""" SELECT team.* FROM common_team AS team INNER JOIN common_membership AS mem ON mem.team_id = team.id WHERE mem.user_id = %s """), - params=(gr_user.id,), - ) + params=(gr_user.id,), + ) - res = c.fetchall() + res = c.fetchall() return [Team.model_validate(item) for item in res] diff --git a/generalresearch/managers/leaderboard/__init__.py b/generalresearch/managers/leaderboard/__init__.py index aae4a05..d5138cd 100644 --- a/generalresearch/managers/leaderboard/__init__.py +++ b/generalresearch/managers/leaderboard/__init__.py @@ -1,4 +1,4 @@ -from typing import Dict +from __future__ import annotations import pytz from cachetools import LRUCache, cached @@ -6,7 +6,7 @@ from zoneinfo import ZoneInfo @cached(cache=LRUCache(maxsize=1)) -def country_timezone() -> Dict[str, ZoneInfo]: +def country_timezone() -> dict[str, ZoneInfo]: """ Most countries only have 1 tz. I am picking the most populous for the rest. A timezone is unique for a country, as in America/New_York and America/Toronto diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index 510dc63..e1143c2 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -3,7 +3,6 @@ from __future__ import annotations import ipaddress from collections.abc import Collection from decimal import Decimal -from random import randint import faker import pymysql @@ -33,38 +32,6 @@ fake = faker.Faker() class IPGeonameManager(PostgresManager): - def create_dummy( - self, - geoname_id: PositiveInt | None = None, - continent_code: str | None = None, - continent_name: str | None = None, - country_iso: str | None = None, - country_name: str | None = None, - subdivision_1_iso: str | None = None, - subdivision_1_name: str | None = None, - subdivision_2_iso: str | None = None, - subdivision_2_name: str | None = None, - city_name: str | None = None, - metro_code: int | None = None, - time_zone: str | None = None, - is_in_european_union: bool | None = None, - ) -> IPGeoname: - return self.create( - geoname_id=geoname_id or randint(1, 999_999_999), - continent_code=continent_code or "na", - continent_name=continent_name or "North America", - country_iso=country_iso or "us", - country_name=country_name or "United States", - subdivision_1_iso=subdivision_1_iso or "fl", - subdivision_1_name=subdivision_1_name or "Florida", - subdivision_2_iso=subdivision_2_iso, - subdivision_2_name=subdivision_2_name, - city_name=city_name, - metro_code=metro_code, - time_zone=time_zone, - is_in_european_union=is_in_european_union, - ) - def create_basic( self, geoname_id: PositiveInt, @@ -230,60 +197,6 @@ class IPGeonameManager(PostgresManager): class IPInformationManager(PostgresManager): - def create_dummy( - self, - ip: IPvAnyAddressStr | None = None, - geoname_id: PositiveInt | None = None, - country_iso: str | None = None, - registered_country_iso: str | None = None, - is_anonymous: bool | None = None, - is_anonymous_vpn: bool | None = None, - is_hosting_provider: bool | None = None, - is_public_proxy: bool | None = None, - is_tor_exit_node: bool | None = None, - is_residential_proxy: bool | None = None, - autonomous_system_number: PositiveInt | None = None, - autonomous_system_organization: str | None = None, - domain: str | None = None, - isp: str | None = None, - mobile_country_code: str | None = None, - mobile_network_code: str | None = None, - network: str | None = None, - organization: str | None = None, - static_ip_score: float | None = None, - user_type: UserType | None = None, - postal_code: str | None = None, - latitude: Decimal | None = None, - longitude: Decimal | None = None, - accuracy_radius: int | None = None, - ) -> IPInformation: - return self.create( - ip=ip or fake.ipv4_public(), - geoname_id=geoname_id, - country_iso=country_iso or fake.country_code(), - registered_country_iso=registered_country_iso, - is_anonymous=is_anonymous, - is_anonymous_vpn=is_anonymous_vpn, - is_hosting_provider=is_hosting_provider, - is_public_proxy=is_public_proxy, - is_tor_exit_node=is_tor_exit_node, - is_residential_proxy=is_residential_proxy, - autonomous_system_number=autonomous_system_number, - autonomous_system_organization=autonomous_system_organization, - domain=domain, - isp=isp, - mobile_country_code=mobile_country_code, - mobile_network_code=mobile_network_code, - network=network, - organization=organization, - static_ip_score=static_ip_score, - user_type=user_type, - postal_code=postal_code, - latitude=latitude, - longitude=longitude, - accuracy_radius=accuracy_radius, - ) - def create_basic( self, ip: IPvAnyAddressStr, @@ -440,16 +353,15 @@ class IPInformationManager(PostgresManager): if len(filter_ips) == 0: return [] - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - res = [] - for chunk in chunked(filter_ips, 500): - res.extend( - self.fetch_ip_information_( - c=c, - filter_ips=chunk, - ) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + res = [] + for chunk in chunked(filter_ips, 500): + res.extend( + self.fetch_ip_information_( + c=c, + filter_ips=chunk, ) + ) return res def fetch_ip_information_( @@ -518,8 +430,6 @@ class IPInformationManager(PostgresManager): # percent_empty = numerator / (denominator or 1) # TODO: Post to telegraf / grafana - return None - class GeoIpInfoManager(PostgresManagerWithRedis): diff --git a/pyproject.toml b/pyproject.toml index dd2e649..183e271 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,6 +9,7 @@ description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.8" dependencies = [ + "fastapi", "Faker", "PyMySQL", "psycopg", diff --git a/test_utils/managers/gr/__init__.py b/test_utils/managers/gr/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py new file mode 100644 index 0000000..3e3a4ad --- /dev/null +++ b/test_utils/managers/gr/conftest.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +from typing import Callable +from uuid import uuid4 + +import pytest +from pydantic import PositiveInt +from pydantic_extra_types.phone_numbers import PhoneNumber + +from generalresearch.managers.gr.authentication import GRUserManager +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import TeamManager +from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + BusinessType, + TransferMethod, +) +from generalresearch.models.gr.team import Team + + +@pytest.fixture +def gr_user_factory(gr_um: GRUserManager) -> Callable[..., GRUser]: + + def _inner( + sub: str | None = None, + is_superuser: bool = False, + ) -> GRUser: + sub = sub or f"{uuid4().hex}-{uuid4().hex}" + + return gr_um.create( + sub=sub, + is_superuser=is_superuser, + ) + + return _inner + + +@pytest.fixture +def gr_business_bank_account_factory( + gr_bbam: BusinessBankAccountManager, +) -> Callable[..., BusinessBankAccount]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + transfer_method: TransferMethod | None = None, + account_number: str | None = None, + routing_number: str | None = None, + iban: str | None = None, + swift: str | None = None, + ): + from generalresearch.models.gr.business import TransferMethod + + return gr_bbam.create( + business_id=business_id, + uuid=uuid or uuid4().hex, + transfer_method=transfer_method or TransferMethod.ACH, + account_number=account_number or uuid4().hex[:6], + routing_number=routing_number or uuid4().hex[:6], + iban=iban or uuid4().hex[:6], + swift=swift or uuid4().hex[:6], + ) + + return _inner + + +@pytest.fixture +def gr_business_address_factory( + gr_bam: BusinessAddressManager, +) -> Callable[..., BusinessAddress]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + line_1: str | None = None, + line_2: str | None = None, + city: str | None = None, + state: str | None = None, + postal_code: str | None = None, + phone_number: PhoneNumber | None = None, + country: str | None = None, + ): + uuid = uuid or uuid4().hex + line_1 = line_1 or "abc" + line_2 = line_2 or "bczx" + city = city or "Downingtown" + state = state or "CA" + postal_code = postal_code or "94041" + phone_number = None + country = country or "US" + + return gr_bam.create( + business_id=business_id, + uuid=uuid, + line_1=line_1, + line_2=line_2, + city=city, + state=state, + postal_code=postal_code, + phone_number=phone_number, + country=country, + ) + + return _inner + + +@pytest.fixture +def gr_business_factory( + gr_bm: BusinessManager, +) -> Callable[..., Business]: + + def _inner( + uuid: UUIDStr | None = None, + name: str | None = None, + team: Team | None = None, + kind: BusinessType | None = None, + tax_number: str | None = None, + ) -> Business: + from random import randint + + uuid = uuid or uuid4().hex + name = name or "< Unknown >" + tax_number = tax_number or str(randint(1, 999_999_999)) + + return gr_bm.create( + uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number + ) + + return _inner + + +@pytest.fixture +def gr_team( + gr_tm: TeamManager, +) -> Callable[..., Team]: + + def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team: + uuid = uuid or uuid4().hex + name = name or f"name-{uuid4().hex[:12]}" + + return gr_tm.create(uuid=uuid, name=name) + + return _inner diff --git a/test_utils/managers/grliq/__init__.py b/test_utils/managers/grliq/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/managers/grliq/conftest.py b/test_utils/managers/grliq/conftest.py new file mode 100644 index 0000000..525d8c8 --- /dev/null +++ b/test_utils/managers/grliq/conftest.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from typing import Callable +from uuid import uuid4 + +import pytest + +from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA +from generalresearch.grliq.managers.forensic_data import GrlIqDataManager +from generalresearch.grliq.models.forensic_data import GrlIqData + + +@pytest.fixture +def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]: + + def _inner( + is_attempt_allowed: bool = True, + product_id: str | None = None, + product_user_id: str | None = None, + uuid: str | None = None, + mid: str | None = None, + created_at: datetime | None = None, + ) -> GrlIqData: + """ + Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), + and GrlIqForensicCategoryResult (category_results) + :param is_attempt_allowed: Whether the attempt is allowed. + :param product_id: product_id of user + :param product_user_id: product_user_id of user + :param uuid: uuid for the grliq data record + :param mid: the thl_session:uuid / mid for the attempt. + :return: + """ + import copy + + res: GrlIqData = copy.deepcopy(DUMMY_GRLIQ_DATA[int(is_attempt_allowed)]) + + product_id = product_id or uuid4().hex + product_user_id = product_user_id or uuid4().hex + uuid = uuid or uuid4().hex + mid = mid or uuid4().hex + created_at = created_at or datetime.now(tz=timezone.utc) + + res["data"].product_id = product_id + res["data"].product_user_id = product_user_id + res["data"].uuid = uuid + res["data"].mid = mid + res["data"].created_at = created_at + res["result_data"].uuid = uuid + res["category_result"].uuid = uuid + + return grliq_dm.create( + iq_data=res["data"], + result_data=res["result_data"], + category_result=res["category_result"], + fraud_score=res["category_result"].fraud_score, + is_attempt_allowed=res["category_result"].is_attempt_allowed(), + ) + + return _inner diff --git a/test_utils/managers/thl/__init__.py b/test_utils/managers/thl/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py new file mode 100644 index 0000000..76e4226 --- /dev/null +++ b/test_utils/managers/thl/conftest.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +from decimal import Decimal +from random import randint +from typing import Callable + +import faker +from pydantic import PositiveInt + +from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager +from generalresearch.models.custom_types import IPvAnyAddressStr +from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType + +fake = faker.Faker() + + +def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: + + def _inner( + geoname_id: PositiveInt | None = None, + continent_code: str | None = None, + continent_name: str | None = None, + country_iso: str | None = None, + country_name: str | None = None, + subdivision_1_iso: str | None = None, + subdivision_1_name: str | None = None, + subdivision_2_iso: str | None = None, + subdivision_2_name: str | None = None, + city_name: str | None = None, + metro_code: int | None = None, + time_zone: str | None = None, + is_in_european_union: bool | None = None, + ) -> IPGeoname: + + return ipgeoname_manager.create( + geoname_id=geoname_id or randint(1, 999_999_999), + continent_code=continent_code or "na", + continent_name=continent_name or "North America", + country_iso=country_iso or "us", + country_name=country_name or "United States", + subdivision_1_iso=subdivision_1_iso or "fl", + subdivision_1_name=subdivision_1_name or "Florida", + subdivision_2_iso=subdivision_2_iso, + subdivision_2_name=subdivision_2_name, + city_name=city_name, + metro_code=metro_code, + time_zone=time_zone, + is_in_european_union=is_in_european_union, + ) + + return _inner + + +def ipinformation_factory( + ipinformation_manager: IPInformationManager, +) -> Callable[..., IPInformation]: + + def _inner( + ip: IPvAnyAddressStr | None = None, + geoname_id: PositiveInt | None = None, + country_iso: str | None = None, + registered_country_iso: str | None = None, + is_anonymous: bool | None = None, + is_anonymous_vpn: bool | None = None, + is_hosting_provider: bool | None = None, + is_public_proxy: bool | None = None, + is_tor_exit_node: bool | None = None, + is_residential_proxy: bool | None = None, + autonomous_system_number: PositiveInt | None = None, + autonomous_system_organization: str | None = None, + domain: str | None = None, + isp: str | None = None, + mobile_country_code: str | None = None, + mobile_network_code: str | None = None, + network: str | None = None, + organization: str | None = None, + static_ip_score: float | None = None, + user_type: UserType | None = None, + postal_code: str | None = None, + latitude: Decimal | None = None, + longitude: Decimal | None = None, + accuracy_radius: int | None = None, + ) -> IPInformation: + + return ipinformation_manager.create( + ip=ip or fake.ipv4_public(), + geoname_id=geoname_id, + country_iso=country_iso or fake.country_code(), + registered_country_iso=registered_country_iso, + is_anonymous=is_anonymous, + is_anonymous_vpn=is_anonymous_vpn, + is_hosting_provider=is_hosting_provider, + is_public_proxy=is_public_proxy, + is_tor_exit_node=is_tor_exit_node, + is_residential_proxy=is_residential_proxy, + autonomous_system_number=autonomous_system_number, + autonomous_system_organization=autonomous_system_organization, + domain=domain, + isp=isp, + mobile_country_code=mobile_country_code, + mobile_network_code=mobile_network_code, + network=network, + organization=organization, + static_ip_score=static_ip_score, + user_type=user_type, + postal_code=postal_code, + latitude=latitude, + longitude=longitude, + accuracy_radius=accuracy_radius, + ) + + return _inner diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 64bdec6..1133d32 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,11 +1,14 @@ +from __future__ import annotations + from datetime import datetime, timedelta, timezone from decimal import Decimal from random import choice as randchoice from random import randint -from typing import TYPE_CHECKING, Callable, Dict, List, Optional +from typing import TYPE_CHECKING, Callable from uuid import uuid4 import pytest +from fastapi import Request from pydantic import AwareDatetime, PositiveInt from generalresearch.models import Source @@ -84,26 +87,27 @@ def user( @pytest.fixture def user_with_wallet( - request, user_factory: Callable[..., "User"], product_user_wallet_yes: "Product" -) -> "User": + user_factory: Callable[..., User], + product_user_wallet_yes: Product, +) -> User: # A user on a product with user wallet enabled, but they have no money return user_factory(product=product_user_wallet_yes) @pytest.fixture def user_with_wallet_amt( - request, user_factory: Callable[..., "User"], product_amt_true: "Product" -) -> "User": + user_factory: Callable[..., User], product_amt_true: Product +) -> User: # A user on a product with user wallet enabled, on AMT, but they have no money return user_factory(product=product_amt_true) @pytest.fixture(scope="function") def user_factory( - user_manager: "UserManager", thl_web_rr: PostgresConfig -) -> Callable[..., "User"]: + user_manager: UserManager, thl_web_rr: PostgresConfig +) -> Callable[..., User]: - def _inner(product: "Product", created: Optional[datetime] = None) -> "User": + def _inner(product: Product, created: datetime | None = None) -> User: u = user_manager.create_dummy(product=product, created=created) u.prefetch_product(pg_config=thl_web_rr) @@ -113,11 +117,11 @@ def user_factory( @pytest.fixture -def wall_factory(wall_manager: "WallManager") -> Callable[..., "Wall"]: +def wall_factory(wall_manager: WallManager) -> Callable[..., Wall]: def _inner( - session: "Session", wall_status: "Status", req_cpi: Optional[Decimal] = None - ) -> "Wall": + session: Session, wall_status: Status, req_cpi: Decimal | None = None + ) -> Wall: assert session.started <= datetime.now( tz=timezone.utc @@ -153,9 +157,7 @@ def wall_factory(wall_manager: "WallManager") -> Callable[..., "Wall"]: @pytest.fixture -def wall( - session: "Session", user: "User", wall_manager: "WallManager" -) -> Optional["Wall"]: +def wall(session: Session, user: User, wall_manager: WallManager) -> Wall | None: from generalresearch.models.thl.task_status import StatusCode1 wall = wall_manager.create_dummy(session_id=session.id, user_id=user.user_id) @@ -170,20 +172,20 @@ def wall( @pytest.fixture def session_factory( - wall_factory: Callable[..., "Wall"], - session_manager: "SessionManager", - wall_manager: "WallManager", + wall_factory: Callable[..., Wall], + session_manager: SessionManager, + wall_manager: WallManager, utc_hour_ago: datetime, -) -> Callable[..., "Session"]: +) -> Callable[..., Session]: from generalresearch.models.thl.session import Source def _inner( - user: "User", + user: User, # Wall details wall_count: int = 5, wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: Optional[List[Decimal]] = None, - wall_statuses: Optional[List[Status]] = None, + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, wall_source: Source = Source.TESTING, # Session details final_status: Status = Status.COMPLETE, @@ -236,24 +238,24 @@ def session_factory( @pytest.fixture(scope="function") def finished_session_factory( - session_factory: Callable[..., "Session"], - session_manager: "SessionManager", + session_factory: Callable[..., Session], + session_manager: SessionManager, utc_hour_ago: datetime, -) -> Callable[..., "Session"]: +) -> Callable[..., Session]: from generalresearch.models.thl.session import Source def _inner( - user: "User", + user: User, # Wall details wall_count: int = 5, wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: Optional[List[Decimal]] = None, - wall_statuses: Optional[List[Status]] = None, + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, wall_source: Source = Source.TESTING, # Session details final_status: Status = Status.COMPLETE, started: datetime = utc_hour_ago, - ) -> "Session": + ) -> Session: s: Session = session_factory( user=user, wall_count=wall_count, @@ -281,9 +283,8 @@ def finished_session_factory( @pytest.fixture def session( - user: "User", session_manager: "SessionManager", wall_manager: "WallManager" -) -> "Session": - from generalresearch.models.thl.session import Session, Wall + user: User, session_manager: SessionManager, wall_manager: WallManager +) -> Session: session: Session = session_manager.create_dummy(user=user, country_iso="us") wall: Wall = wall_manager.create_dummy( @@ -297,7 +298,7 @@ def session( @pytest.fixture -def product(request, product_manager: "ProductManager") -> "Product": +def product(request: Request, product_manager: ProductManager) -> Product: team = getattr(request, "team", None) business = getattr(request, "business", None) @@ -309,13 +310,13 @@ def product(request, product_manager: "ProductManager") -> "Product": @pytest.fixture -def product_factory(product_manager: "ProductManager") -> Callable[..., "Product"]: +def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( - team: Optional["Team"] = None, - business: Optional["Business"] = None, + team: Team | None = None, + business: Business | None = None, commission_pct: Decimal = Decimal("0.05"), - ) -> "Product": + ) -> Product: return product_manager.create_dummy( team_id=team.uuid if team else None, business_id=business.uuid if business else None, @@ -326,7 +327,7 @@ def product_factory(product_manager: "ProductManager") -> Callable[..., "Product @pytest.fixture -def payout_config(request) -> "PayoutConfig": +def payout_config(request: Request) -> PayoutConfig: from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, @@ -348,8 +349,8 @@ def payout_config(request) -> "PayoutConfig": @pytest.fixture def product_user_wallet_yes( - payout_config: "PayoutConfig", product_manager: "ProductManager" -) -> "Product": + payout_config: PayoutConfig, product_manager: ProductManager +) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -358,7 +359,7 @@ def product_user_wallet_yes( @pytest.fixture -def product_user_wallet_no(product_manager: "ProductManager") -> "Product": +def product_user_wallet_no(product_manager: ProductManager) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -368,8 +369,8 @@ def product_user_wallet_no(product_manager: "ProductManager") -> "Product": @pytest.fixture def product_amt_true( - product_manager: "ProductManager", payout_config: "PayoutConfig" -) -> "Product": + product_manager: ProductManager, payout_config: PayoutConfig +) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -380,19 +381,19 @@ def product_amt_true( @pytest.fixture def bp_payout_factory( - thl_lm: "ThlLedgerManager", - product_manager: "ProductManager", - business_payout_event_manager: "BusinessPayoutEventManager", -) -> Callable[..., "BrokerageProductPayoutEvent"]: + thl_lm: ThlLedgerManager, + product_manager: ProductManager, + business_payout_event_manager: BusinessPayoutEventManager, +) -> Callable[..., BrokerageProductPayoutEvent]: def _inner( - product: Optional["Product"] = None, - amount: Optional["USDCent"] = None, - ext_ref_id: Optional[str] = None, - created: Optional[AwareDatetime] = None, + product: Product | None = None, + amount: USDCent | None = None, + ext_ref_id: str | None = None, + created: AwareDatetime | None = None, skip_wallet_balance_check: bool = False, skip_one_per_day_check: bool = False, - ) -> "BrokerageProductPayoutEvent": + ) -> BrokerageProductPayoutEvent: from generalresearch.currency import USDCent product = product or product_manager.create_dummy() @@ -415,43 +416,43 @@ def bp_payout_factory( @pytest.fixture -def business(request, business_manager: "BusinessManager") -> "Business": +def business(request, business_manager: BusinessManager) -> Business: return business_manager.create_dummy() @pytest.fixture def business_address( - request, business: "Business", business_address_manager: "BusinessAddressManager" -) -> "BusinessAddress": + request, business: "Business", business_address_manager: BusinessAddressManager +) -> BusinessAddress: return business_address_manager.create_dummy(business_id=business.id) @pytest.fixture def business_bank_account( request, - business: "Business", - business_bank_account_manager: "BusinessBankAccountManager", -) -> "BusinessBankAccount": + business: Business, + business_bank_account_manager: BusinessBankAccountManager, +) -> BusinessBankAccount: return business_bank_account_manager.create_dummy(business_id=business.id) @pytest.fixture -def team(request, team_manager: "TeamManager") -> "Team": +def team(request, team_manager: TeamManager) -> Team: return team_manager.create_dummy() @pytest.fixture -def gr_user(gr_um: "GRUserManager") -> "GRUser": +def gr_user(gr_um: GRUserManager) -> GRUser: return gr_um.create_dummy() @pytest.fixture def gr_user_cache( - gr_user: "GRUser", + gr_user: GRUser, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, -) -> "GRUser": +) -> GRUser: gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) @@ -459,7 +460,7 @@ def gr_user_cache( @pytest.fixture -def gr_user_factory(gr_um: "GRUserManager") -> Callable[..., "GRUser"]: +def gr_user_factory(gr_um: GRUserManager) -> Callable[..., GRUser]: def _inner(): return gr_um.create_dummy() @@ -469,8 +470,8 @@ def gr_user_factory(gr_um: "GRUserManager") -> Callable[..., "GRUser"]: @pytest.fixture() def gr_user_token( - gr_user: "GRUser", gr_tm: "GRTokenManager", gr_db: PostgresConfig -) -> "GRToken": + gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig +) -> GRToken: gr_tm.create(user_id=gr_user.id) gr_user.prefetch_token(pg_config=gr_db) @@ -480,14 +481,14 @@ def gr_user_token( @pytest.fixture() -def gr_user_token_header(gr_user_token: "GRToken") -> Dict[str, str]: +def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: return gr_user_token.auth_header @pytest.fixture(scope="function") def membership( - request, team: "Team", gr_user: "GRUser", team_manager: "TeamManager" -) -> "Membership": + request, team: Team, gr_user: GRUser, team_manager: TeamManager +) -> Membership: assert team.id, "Team must be saved" assert gr_user.id, "GRUser must be saved" return team_manager.add_user(team=team, gr_user=gr_user) @@ -495,14 +496,14 @@ def membership( @pytest.fixture(scope="function") def membership_factory( - team: "Team", - gr_user: "GRUser", - membership_manager: "MembershipManager", - team_manager: "TeamManager", - gr_um: "GRUserManager", -) -> Callable[..., "Membership"]: - - def _inner(**kwargs) -> "Membership": + team: Team, + gr_user: GRUser, + membership_manager: MembershipManager, + team_manager: TeamManager, + gr_um: GRUserManager, +) -> Callable[..., Membership]: + + def _inner(**kwargs) -> Membership: _team = kwargs.get("team", team_manager.create_dummy()) _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) @@ -512,23 +513,23 @@ def membership_factory( @pytest.fixture -def audit_log(audit_log_manager: "AuditLogManager", user: "User") -> "AuditLog": +def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog: return audit_log_manager.create_dummy(user_id=user.user_id) @pytest.fixture def audit_log_factory( - audit_log_manager: "AuditLogManager", -) -> Callable[..., "AuditLog"]: + audit_log_manager: AuditLogManager, +) -> Callable[..., AuditLog]: def _inner( user_id: PositiveInt, - level: Optional["AuditLogLevel"] = None, - event_type: Optional[str] = None, - event_msg: Optional[str] = None, - event_value: Optional[float] = None, - ) -> "AuditLog": + level: AuditLogLevel | None = None, + event_type: str | None = None, + event_msg: str | None = None, + event_value: float | None = None, + ) -> AuditLog: return audit_log_manager.create_dummy( user_id=user_id, level=level, @@ -541,14 +542,14 @@ def audit_log_factory( @pytest.fixture -def ip_geoname(ip_geoname_manager: "IPGeonameManager") -> "IPGeoname": +def ip_geoname(ip_geoname_manager: IPGeonameManager) -> IPGeoname: return ip_geoname_manager.create_dummy() @pytest.fixture def ip_information( - ip_information_manager: "IPInformationManager", ip_geoname: "IPGeoname" -) -> "IPInformation": + ip_information_manager: IPInformationManager, ip_geoname: IPGeoname +) -> IPInformation: return ip_information_manager.create_dummy( geoname_id=ip_geoname.geoname_id, country_iso=ip_geoname.country_iso ) @@ -556,10 +557,10 @@ def ip_information( @pytest.fixture def ip_information_factory( - ip_information_manager: "IPInformationManager", -) -> Callable[..., "IPInformation"]: + ip_information_manager: IPInformationManager, +) -> Callable[..., IPInformation]: - def _inner(ip: str, geoname: "IPGeoname", **kwargs) -> "IPInformation": + def _inner(ip: str, geoname: IPGeoname, **kwargs) -> IPInformation: return ip_information_manager.create_dummy( ip=ip, geoname_id=geoname.geoname_id, @@ -572,25 +573,25 @@ def ip_information_factory( @pytest.fixture def ip_record( - ip_record_manager: "IPRecordManager", ip_geoname: "IPGeoname", user: "User" -) -> "IPRecord": + ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User +) -> IPRecord: return ip_record_manager.create_dummy(user_id=user.user_id) @pytest.fixture def ip_record_factory( - ip_record_manager: "IPRecordManager", user: "User" -) -> Callable[..., "IPRecord"]: + ip_record_manager: IPRecordManager, user: User +) -> Callable[..., IPRecord]: - def _inner(user_id: PositiveInt, ip: Optional[str] = None) -> "IPRecord": + def _inner(user_id: PositiveInt, ip: str | None = None) -> IPRecord: return ip_record_manager.create_dummy(user_id=user_id, ip=ip) return _inner @pytest.fixture(scope="session") -def buyer(buyer_manager: "BuyerManager") -> "Buyer": +def buyer(buyer_manager: BuyerManager) -> Buyer: buyer_code = uuid4().hex buyer_manager.bulk_get_or_create(source=Source.TESTING, codes=[buyer_code]) b = Buyer( @@ -601,7 +602,7 @@ def buyer(buyer_manager: "BuyerManager") -> "Buyer": @pytest.fixture(scope="session") -def buyer_factory(buyer_manager: "BuyerManager") -> Callable[..., "Buyer"]: +def buyer_factory(buyer_manager: BuyerManager) -> Callable[..., Buyer]: def _inner() -> Buyer: return buyer_manager.bulk_get_or_create( @@ -612,7 +613,7 @@ def buyer_factory(buyer_manager: "BuyerManager") -> Callable[..., "Buyer"]: @pytest.fixture(scope="session") -def survey(survey_manager: "SurveyManager", buyer: "Buyer") -> "Survey": +def survey(survey_manager: SurveyManager, buyer: Buyer) -> "Survey": s = Survey(source=Source.TESTING, survey_id=uuid4().hex, buyer_code=buyer.code) survey_manager.create_bulk([s]) return s @@ -620,10 +621,10 @@ def survey(survey_manager: "SurveyManager", buyer: "Buyer") -> "Survey": @pytest.fixture(scope="session") def survey_factory( - survey_manager: "SurveyManager", buyer_factory: Callable[..., "Buyer"] -) -> Callable[..., "Survey"]: + survey_manager: SurveyManager, buyer_factory: Callable[..., Buyer] +) -> Callable[..., Survey]: - def _inner(buyer: Optional[Buyer] = None) -> "Survey": + def _inner(buyer: Buyer | None = None) -> Survey: buyer = buyer or buyer_factory() s = Survey( source=Source.TESTING, diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 497e5ab..7e6605f 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -1,5 +1,6 @@ -from datetime import datetime, timezone, timedelta -from os.path import exists as pexists, join as pjoin +from datetime import datetime, timedelta, timezone +from os.path import exists as pexists +from os.path import join as pjoin from pathlib import Path from uuid import uuid4 @@ -244,6 +245,7 @@ class TestCollectionBaseMethodsCleanup: class TestCollectionBaseMethodsCleanup: + @pytest.mark.skip def test_cleanup_partials(self, mnt_filepath): instance = CollectionBase(archive_path=mnt_filepath.data_src) diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 14d1343..530142e 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -1,6 +1,7 @@ +from __future__ import annotations + import logging from datetime import datetime, timezone -from typing import Optional import pytest import pytz @@ -12,7 +13,7 @@ logger = logging.getLogger() class AwareDatetimeISOModel(BaseModel): - dt_optional: Optional[AwareDatetimeISO] = Field(default=None) + dt_optional: AwareDatetimeISO | None = Field(default=None) dt: AwareDatetimeISO diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 5e9b249..39469dc 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -1,8 +1,10 @@ +from __future__ import annotations + import os import shutil from datetime import datetime, timedelta, timezone from decimal import Decimal -from typing import Callable, Optional +from typing import Callable from uuid import uuid4 import pytest @@ -10,7 +12,7 @@ from dask.distributed import Client as DaskClient from pydantic import ValidationError from generalresearch.currency import USDCent -from generalresearch.incite import GRLDatasets +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, @@ -25,6 +27,7 @@ from generalresearch.models.thl.product import ( IntegrationMode, PayoutConfig, PayoutTransformation, + PayoutTransformationPercentArgs, Product, ProfilingConfig, SourceConfig, @@ -137,6 +140,12 @@ class TestProduct: redirect_url="https://www.google.com/hey", ) + assert isinstance(p.payout_config.payout_transformation, PayoutTransformation) + assert isinstance( + p.payout_config.payout_transformation.kwargs, + PayoutTransformationPercentArgs, + ) + p.payout_config.payout_transformation = PayoutTransformation.model_validate( { "f": "payout_transformation_percent", @@ -576,7 +585,7 @@ class TestGlobalProductConfigFor: class TestProductFinancials: @pytest.fixture - def start(self) -> "datetime": + def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) @pytest.fixture @@ -584,7 +593,7 @@ class TestProductFinancials: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None def test_balance( @@ -759,7 +768,7 @@ class TestProductFinancials: class TestProductBalance: @pytest.fixture - def start(self) -> "datetime": + def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) @pytest.fixture @@ -767,7 +776,7 @@ class TestProductBalance: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None def test_inconsistent( @@ -783,7 +792,7 @@ class TestProductBalance: user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], pop_ledger_merge, - start, + start: datetime, bp_payout_factory, payout_event_manager, ): @@ -792,8 +801,6 @@ class TestProductBalance: create_main_accounts() delete_df_collection(coll=ledger_collection) - from generalresearch.models.thl.user import User - u1: User = user_factory(product=product) # 1. Complete and Build Parquets 1st time @@ -827,7 +834,7 @@ class TestProductBalance: def test_not_inconsistent( self, product: Product, - mnt_filepath, + mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db, @@ -852,8 +859,6 @@ class TestProductBalance: create_main_accounts() delete_df_collection(coll=ledger_collection) - from generalresearch.models.thl.user import User - u1: User = user_factory(product=product) # 1. Complete and Build Parquets 1st time @@ -886,7 +891,7 @@ class TestProductBalance: class TestProductPOPFinancial: @pytest.fixture - def start(self) -> "datetime": + def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) @pytest.fixture @@ -894,15 +899,15 @@ class TestProductPOPFinancial: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None def test_base( self, - product, - mnt_filepath, + product: Product, + mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, - client_no_amm, + client_no_amm: DaskClient, delete_ledger_db, create_main_accounts, delete_df_collection, @@ -923,8 +928,6 @@ class TestProductPOPFinancial: create_main_accounts() delete_df_collection(coll=ledger_collection) - from generalresearch.models.thl.user import User - u1: User = user_factory(product=product) # 1. Complete and Build Parquets 1st time @@ -961,7 +964,7 @@ class TestProductPOPFinancial: class TestProductCache: @pytest.fixture - def start(self) -> "datetime": + def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc) @pytest.fixture @@ -969,7 +972,7 @@ class TestProductCache: return "30d" @pytest.fixture - def duration(self) -> Optional["timedelta"]: + def duration(self) -> timedelta | None: return None def test_basic( @@ -1008,7 +1011,6 @@ class TestProductCache: ) from generalresearch.models.thl.product import Product - from generalresearch.models.thl.user import User u1: User = user_factory(product=product) @@ -1047,7 +1049,7 @@ class TestProductCache: def test_neg_balance_cache( self, product: Product, - mnt_filepath, + mnt_filepath: GRLDatasets, thl_lm, client_no_amm: DaskClient, thl_redis_config, @@ -1070,7 +1072,6 @@ class TestProductCache: delete_df_collection(coll=ledger_collection) from generalresearch.models.thl.product import Product - from generalresearch.models.thl.user import User u1: User = user_factory(product=product) @@ -1112,10 +1113,12 @@ class TestProductCache: # Fetch from cache and assert the instance loaded from redis rc = thl_redis_config.create_redis_client() - res: Optional[str] = rc.get(product.cache_key) + res: str | None = rc.get(product.cache_key) assert isinstance(res, str) p1: Product = Product.model_validate_json(res) + assert p1.balance + assert p1.balance.product_id == product.uuid assert p1.balance.payout_usd_str == "$0.71" assert p1.balance.adjustment == -71 -- cgit v1.2.3 From 05102628a7dc85a5a19a32415a7ec41ea7a91812 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 21 Aug 2026 01:38:27 -0700 Subject: moving create_dummy, organizing test_utils, working on django config for any repos --- generalresearch/config.py | 2 + generalresearch/managers/thl/delete_request.py | 178 ----- generalresearch/managers/thl/payout.py | 51 +- generalresearch/managers/thl/product.py | 42 +- generalresearch/managers/thl/session.py | 40 +- .../managers/thl/user_manager/user_manager.py | 20 - generalresearch/managers/thl/userhealth.py | 66 +- generalresearch/managers/thl/wall.py | 51 +- generalresearch/models/gr/business.py | 2 +- generalresearch/models/network/nmap/parser.py | 2 +- generalresearch/models/network/rdns/parser.py | 2 +- test_utils/conftest.py | 228 ++---- test_utils/grliq/conftest.py | 124 +++- test_utils/grliq/managers/__init__.py | 0 test_utils/grliq/managers/conftest.py | 0 test_utils/grliq/models/__init__.py | 0 test_utils/grliq/models/conftest.py | 0 test_utils/incite/collections/conftest.py | 38 +- test_utils/incite/conftest.py | 34 +- test_utils/incite/mergers/conftest.py | 105 ++- test_utils/managers/conftest.py | 578 ++------------- test_utils/managers/contest/conftest.py | 294 +------- test_utils/managers/gr/conftest.py | 197 +++--- test_utils/managers/grliq/__init__.py | 0 test_utils/managers/grliq/conftest.py | 61 -- test_utils/managers/ledger/conftest.py | 778 ++------------------- test_utils/managers/network/conftest.py | 143 ---- test_utils/managers/thl/conftest.py | 362 +++++++--- test_utils/managers/upk/conftest.py | 188 ++--- .../managers/upk/marketplace_category.csv.gz | Bin 100990 -> 0 bytes test_utils/managers/upk/marketplace_item.csv.gz | Bin 3225 -> 0 bytes .../managers/upk/marketplace_property.csv.gz | Bin 3315 -> 0 bytes .../marketplace_propertycategoryassociation.csv.gz | Bin 2079 -> 0 bytes .../upk/marketplace_propertycountry.csv.gz | Bin 71359 -> 0 bytes .../upk/marketplace_propertyitemrange.csv.gz | Bin 65389 -> 0 bytes ...rketplace_propertymarketplaceassociation.csv.gz | Bin 4272 -> 0 bytes .../managers/upk/marketplace_question.csv.gz | Bin 283465 -> 0 bytes test_utils/models/conftest.py | 86 +-- test_utils/models/contest/__init__.py | 0 test_utils/models/contest/conftest.py | 292 ++++++++ test_utils/models/gr/__init__.py | 0 test_utils/models/gr/conftest.py | 213 ++++++ test_utils/models/ledger/__init__.py | 0 test_utils/models/ledger/conftest.py | 724 +++++++++++++++++++ test_utils/models/network/__init__.py | 0 test_utils/models/network/conftest.py | 144 ++++ test_utils/models/thl/__init__.py | 0 test_utils/models/thl/conftest.py | 434 ++++++++++++ test_utils/models/upk/__init__.py | 0 test_utils/models/upk/conftest.py | 178 +++++ test_utils/models/upk/marketplace_category.csv.gz | Bin 0 -> 100990 bytes test_utils/models/upk/marketplace_item.csv.gz | Bin 0 -> 3225 bytes test_utils/models/upk/marketplace_property.csv.gz | Bin 0 -> 3315 bytes .../marketplace_propertycategoryassociation.csv.gz | Bin 0 -> 2079 bytes .../models/upk/marketplace_propertycountry.csv.gz | Bin 0 -> 71359 bytes .../upk/marketplace_propertyitemrange.csv.gz | Bin 0 -> 65389 bytes ...rketplace_propertymarketplaceassociation.csv.gz | Bin 0 -> 4272 bytes test_utils/models/upk/marketplace_question.csv.gz | Bin 0 -> 283465 bytes tests/conftest.py | 10 +- tests/grliq/managers/test_forensic_data.py | 35 +- tests/grliq/managers/test_forensic_results.py | 4 +- tests/models/admin/test_report_request.py | 23 +- tests/models/custom_types/test_dsn.py | 9 +- tests/models/custom_types/test_uuid_str.py | 7 +- tests/models/dynata/test_eligbility.py | 10 +- tests/models/gr/test_authentication.py | 36 +- tests/models/gr/test_base.py | 25 + tests/models/gr/test_business.py | 75 +- 68 files changed, 2929 insertions(+), 2962 deletions(-) delete mode 100644 generalresearch/managers/thl/delete_request.py delete mode 100644 test_utils/grliq/managers/__init__.py delete mode 100644 test_utils/grliq/managers/conftest.py delete mode 100644 test_utils/grliq/models/__init__.py delete mode 100644 test_utils/grliq/models/conftest.py delete mode 100644 test_utils/managers/grliq/__init__.py delete mode 100644 test_utils/managers/grliq/conftest.py delete mode 100644 test_utils/managers/upk/marketplace_category.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_item.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_property.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_propertycountry.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_propertyitemrange.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz delete mode 100644 test_utils/managers/upk/marketplace_question.csv.gz create mode 100644 test_utils/models/contest/__init__.py create mode 100644 test_utils/models/contest/conftest.py create mode 100644 test_utils/models/gr/__init__.py create mode 100644 test_utils/models/gr/conftest.py create mode 100644 test_utils/models/ledger/__init__.py create mode 100644 test_utils/models/ledger/conftest.py create mode 100644 test_utils/models/network/__init__.py create mode 100644 test_utils/models/network/conftest.py create mode 100644 test_utils/models/thl/__init__.py create mode 100644 test_utils/models/thl/conftest.py create mode 100644 test_utils/models/upk/__init__.py create mode 100644 test_utils/models/upk/conftest.py create mode 100644 test_utils/models/upk/marketplace_category.csv.gz create mode 100644 test_utils/models/upk/marketplace_item.csv.gz create mode 100644 test_utils/models/upk/marketplace_property.csv.gz create mode 100644 test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz create mode 100644 test_utils/models/upk/marketplace_propertycountry.csv.gz create mode 100644 test_utils/models/upk/marketplace_propertyitemrange.csv.gz create mode 100644 test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz create mode 100644 test_utils/models/upk/marketplace_question.csv.gz create mode 100644 tests/models/gr/test_base.py (limited to 'tests') diff --git a/generalresearch/config.py b/generalresearch/config.py index af80069..44f3db7 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -53,6 +53,8 @@ class GRLBaseSettings(BaseSettings): testing_postgres_user: str | None = Field(default=None) testing_postgres_pass: str | None = Field(default=None) + git_creds: str | None = Field(default=None) + # --- redis: RedisDsn | None = Field(default=None) diff --git a/generalresearch/managers/thl/delete_request.py b/generalresearch/managers/thl/delete_request.py deleted file mode 100644 index 963cb7b..0000000 --- a/generalresearch/managers/thl/delete_request.py +++ /dev/null @@ -1,178 +0,0 @@ -# from datetime import datetime, timezone -# from typing import Optional -# -# from generalresearch.managers.gr.authentication import GRUserManager -# from generalresearch.managers.thl.user_manager.user_manager import UserManager -# from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -# from pydantic import BaseModel, Field, PositiveInt, model_validator -# from pydantic.json_schema import SkipJsonSchema -# -# from api.decorators import THL_WEB_RR, GR_DB -# -# GR_UM = GRUserManager(sql_helper=GR_DB) -# UM = UserManager(sql_helper_rr=THL_WEB_RR) -# - -# @pytest.mark.skip(reason="moving to pyutils 2.5.1") -# class TestUserDeleteRequestManager: -# -# def test_delete_request(self, gr_user, user, product, user_manager, gr_um): -# from api.models.product_user import DeleteRequest -# from api.managers.product_user import UserDeletionRequestManager -# -# # A valid Respondent and GR Admin account need to exist in the test -# # database for any of this to work -# user = user_manager.create_dummy( -# product_id=product.id, -# product_user_id=f"test-{uuid4().hex[:6]}", -# ) -# -# instance = DeleteRequest( -# product_id=user.product_id, -# product_user_id=user.product_user_id, -# created_by_user_id=gr_user.id, -# ) -# -# start: int = UserDeletionRequestManager().get_count_by_product_id( -# product_id=user.product_id -# ) -# -# UserDeletionRequestManager.save(deletion_request=instance) -# -# finish: int = UserDeletionRequestManager().get_count_by_product_id( -# product_id=user.product_id -# ) -# -# assert finish == start + 1 - - -# @pytest.mark.skip(reason="Moving to generalresearch in 2.5.1") -# class TestProductUserDeleteRequest: -# -# def test_no_user_provided(self, product, business, team, gr_user): -# from api.models.product_user import DeleteRequest -# -# # product_id and product_user_id is required -# with pytest.raises(expected_exception=ValueError) as cm: -# DeleteRequest(created_by_user_id=gr_user.id) -# -# assert "2 validation errors" in str(cm.value) -# -# def test_no_user_exists(self, gr_user, product): -# from api.models.product_user import DeleteRequest -# -# with pytest.raises(expected_exception=ValueError) as cm: -# DeleteRequest( -# product_id=product.id, -# product_user_id=f"test-user-{uuid4().hex[:12]}", -# created_by_user_id=gr_user.id, -# ) -# -# assert "Could not find Worker" in str(cm.value) -# -# def test_no_create_by_user(self, user, product): -# from api.models.product_user import DeleteRequest -# -# with pytest.raises(expected_exception=ValueError) as cm: -# DeleteRequest( -# product_id=user.product_id, -# product_user_id=user.product_user_id, -# created_by_user_id=randint(a=999_999, b=999_999_999), -# ) -# assert "GRUser not found" in str(cm.value) - - -# -# class DeleteRequest(BaseModel): -# id: SkipJsonSchema[Optional[PositiveInt]] = Field(default=None, exclude=True) -# uuid: UUIDStr = Field(examples=[uuid4().hex], default_factory=lambda: uuid4().hex) -# -# product_id: UUIDStr = Field(examples=["00e96773d4ae47f8812488a976a080c8"]) -# product_user_id: str = Field( -# min_length=3, max_length=128, examples=["bpuid-68d989"] -# ) -# -# created: AwareDatetimeISO = Field( -# default=datetime.now(tz=timezone.utc), -# description="When the DeleteRequest was created, this is the UTC time " -# "that a Worker / Respondent's Profiling Questions were " -# "deleted.", -# ) -# created_by_user_id: SkipJsonSchema[PositiveInt] = Field(exclude=True) -# -# @model_validator(mode="after") -# def check_valid_worker(self) -> "DeleteRequest": -# """ Raise an error if the User that the GRUser is attempting to delete -# does not actually exist in the system. We can check the production -# thl-web user table here for real time users -# """ -# user = UM.get_user_if_exists( -# product_id=self.product_id, product_user_id=self.product_user_id -# ) -# -# if not user: -# raise ValueError("Could not find Worker") -# -# return self -# -# @model_validator(mode="after") -# def check_valid_owner(self) -> "DeleteRequest": -# """ Ensure we can track which GRUser made a deletion request so we can -# track the chain of command for who took what action. -# -# """ -# gr_user = GR_UM.get_by_id(gr_user_id=self.created_by_user_id) -# -# if not gr_user: -# raise ValueError("Could not find General Research account") -# -# return self - - -# @staticmethod -# def save(deletion_request: DeleteRequest) -> bool: -# with GR_DB.make_connection() as conn: -# with conn.cursor(row_factory=dict_row) as c: -# c: Cursor -# -# c.execute( -# query=f""" -# INSERT INTO product_user_deleterequest -# (uuid, product_id, product_user_id, created, -# created_by_user_id) -# VALUES (%s, %s, %s, %s, %s) -# """, -# params=[ -# deletion_request.uuid, -# deletion_request.product_id, -# deletion_request.product_user_id, -# deletion_request.created, -# deletion_request.created_by_user_id, -# ], -# ) -# -# conn.commit() -# -# return True -# -# -# @staticmethod -# def get_count_by_product_id(product_id: UUIDStr) -> NonNegativeInt: -# with GR_DB.make_connection() as conn: -# with conn.cursor(row_factory=dict_row) as c: -# c: Cursor -# -# c.execute( -# query=f""" -# SELECT COUNT(1) as cnt -# FROM product_user_deleterequest AS dr -# WHERE dr.product_id = %s -# """, -# params=[ -# product_id, -# ], -# ) -# res = c.fetchall() -# -# assert len(res) == 1, "invalid query" -# return int(res[0]["cnt"]) diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 364e1aa..1f3742e 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -3,8 +3,6 @@ from __future__ import annotations from collections import defaultdict from collections.abc import Collection from datetime import datetime, timedelta, timezone -from random import choice as rand_choice -from random import randint from time import sleep from typing import Any from uuid import UUID, uuid4 @@ -339,53 +337,6 @@ class UserPayoutEventManager(PayoutEventManager): return payout_event - def create_dummy( - self, - uuid: UUIDStr | None = None, - debit_account_uuid: UUIDStr | None = None, - account_reference_type: str | None = None, - account_reference_uuid: UUIDStr | None = None, - cashout_method_uuid: UUIDStr | None = None, - description: str | None = None, - created: AwareDatetimeISO | None = None, - amount: PositiveInt | None = None, - status: PayoutStatus | None = None, - ext_ref_id: str | None = None, - payout_type: PayoutType | None = None, - request_data: dict[str, Any] | None = None, - order_data: dict[str, Any] | CashMailOrderData | None = None, - ) -> UserPayoutEvent: - - debit_account_uuid = debit_account_uuid or uuid4().hex - cashout_method_uuid = cashout_method_uuid or uuid4().hex - # account_reference_type = account_reference_type or f"acct-ref-{uuid4().hex}" - # account_reference_uuid = account_reference_uuid or uuid4().hex - # cashout_method_uuid = cashout_method_uuid or uuid4().hex - amount = amount or randint(a=99, b=9_999) - status = status or rand_choice(list(PayoutStatus)) - - description = description or f"desc-{uuid4().hex[:12]}" - # ext_ref_id = ext_ref_id or f"ext-ref-{uuid4().hex[:8]}" - payout_type = payout_type or rand_choice(list(PayoutType)) - request_data = request_data or {} - # order_data = order_data or None - - return self.create( - uuid=uuid, - debit_account_uuid=debit_account_uuid, - account_reference_type=account_reference_type, - account_reference_uuid=account_reference_uuid, - cashout_method_uuid=cashout_method_uuid, - description=description, - created=created, - amount=amount, - status=status, - ext_ref_id=ext_ref_id, - payout_type=payout_type, - request_data=request_data, - order_data=order_data, - ) - class BrokerageProductPayoutEventManager(PayoutEventManager): # This is what makes a PayoutEvent a Brokerage Product Payout @@ -1231,7 +1182,7 @@ class BusinessPayoutEventManager(BrokerageProductPayoutEventManager): amount=USDCent(item["issue_amount"]), created=created + timedelta(milliseconds=idx + 1), ext_ref_id=transaction_id, - skip_wallet_balance_check=True + skip_wallet_balance_check=True, ) assert bp_pe.status == PayoutStatus.COMPLETE diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 00bf032..46280b8 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -8,7 +8,7 @@ from datetime import datetime, timezone from decimal import Decimal from threading import Lock from typing import TYPE_CHECKING -from uuid import UUID, uuid4 +from uuid import UUID from cachetools import TTLCache, cachedmethod, keys from more_itertools import chunked @@ -264,46 +264,6 @@ class ProductManager(PostgresManager): raise e return r - def create_dummy( - self, - product_id: UUIDStr | None = None, - team_id: UUIDStr | None = None, - business_id: UUIDStr | None = None, - name: str | None = None, - redirect_url: str | None = None, - harmonizer_domain: str | None = None, - commission_pct: Decimal = Decimal("0.05000"), - sources_config: SourcesConfig | SupplyConfigs | None = None, - payout_config: PayoutConfig | None = None, - session_config: SessionConfig | None = None, - profiling_config: ProfilingConfig | None = None, - user_wallet_config: UserWalletConfig | None = None, - user_create_config: UserCreateConfig | None = None, - user_health_config: UserHealthConfig | None = None, - ) -> Product: - """To be used in tests, where we don't care about certain fields""" - product_id = product_id if product_id else uuid4().hex - team_id = team_id if team_id else uuid4().hex - name = name if name else f"name-{product_id[:12]}" - redirect_url = redirect_url if redirect_url else "https://www.example.com/" - - return self.create( - product_id=product_id, - team_id=team_id, - business_id=business_id, - name=name, - redirect_url=redirect_url, - harmonizer_domain=harmonizer_domain, - commission_pct=commission_pct, - sources_config=sources_config, - payout_config=payout_config, - session_config=session_config, - profiling_config=profiling_config, - user_wallet_config=user_wallet_config, - user_create_config=user_create_config, - user_health_config=user_health_config, - ) - def create( self, product_id: UUIDStr, diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index bb467a2..746a518 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -91,40 +91,6 @@ class SessionManager(PostgresManager): conn.commit() return session - def create_dummy( - self, - # -- Create Dummy "optional" -- # - started: datetime | None = None, - user: User | None = None, - # -- Optional -- # - country_iso: str | None = None, - device_type: DeviceType | None = None, - ip: str | None = None, - bucket: Bucket | None = None, - url_metadata: dict[str, str] | None = None, - uuid_id: str | None = None, - ) -> Session: - """To be used in tests, where we don't care about certain fields""" - started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), - end_date=datetime(year=2000, month=1, day=1, tzinfo=timezone.utc), - tzinfo=timezone.utc, - ) - user = user or User( - user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex - ) - - return self.create( - started=started, - user=user, - country_iso=country_iso, - device_type=device_type, - ip=ip, - bucket=bucket, - url_metadata=url_metadata, - uuid_id=uuid_id, - ) - def get_from_uuid(self, session_uuid: UUIDStr) -> Session: query = """ SELECT @@ -171,7 +137,7 @@ class SessionManager(PostgresManager): assert len(res) == 1 return self.session_from_mysql(res[0]) - def session_from_mysql(self, d: dict) -> Session: + def session_from_mysql(self, d: dict[str, Any]) -> Session: d["id"] = d.pop("session_id") d["uuid"] = UUID(d.pop("session_uuid")).hex d["user"] = User( @@ -283,8 +249,6 @@ class SessionManager(PostgresManager): params=d, ) - return None - def filter_paginated( self, user_id: PositiveInt | None = None, @@ -538,7 +502,7 @@ class SessionManager(PostgresManager): # We need to include the cases where status is NULL as ABANDON. We'll handle the distinction # between TIMEOUT (no status, older than 90 min) and UNKNOWN (no status, newer than 90 min) later. params["status"] = status.value - filters.append(f"COALESCE(status, 'a') = %(status)s") + filters.append("COALESCE(status, 'a') = %(status)s") if extra_filters: filters.append(extra_filters) diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index a7bbd7e..3794020 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -4,7 +4,6 @@ import logging from collections.abc import Collection from datetime import datetime from functools import lru_cache -from uuid import uuid4 from pydantic import RedisDsn @@ -294,25 +293,6 @@ class UserManager: return user - def create_dummy( - self, - # --- Create dummy "optional" --- # - product_user_id: str | None = None, - # --- Optional --- # - product_id: UUIDStr | None = None, - product: Product | None = None, - created: datetime | None = None, - ) -> User: - - product_user_id = product_user_id or uuid4().hex - - return self.create_user( - product_user_id=product_user_id, - product_id=product_id, - product=product, - created=created, - ) - def product_id_exists(self, product_id: str) -> bool: mysql_user_manager = self.mysql_user_manager_rr or self.mysql_user_manager return mysql_user_manager.product_id_exists(product_id) diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 8b951c0..fe2163f 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -4,8 +4,6 @@ import ipaddress from collections.abc import Collection from datetime import datetime, timedelta, timezone from itertools import zip_longest -from random import choice as rchoice -from random import random from typing import Any import faker @@ -184,30 +182,6 @@ class IPRecordManager(PostgresManagerWithRedis): permissions=self.permissions, ) - def create_dummy( - self, - user_id: PositiveInt, - ip: IPvAnyAddressStr | None = None, - forwarded_ip1: IPvAnyAddressStr | None = None, - forwarded_ip2: IPvAnyAddressStr | None = None, - forwarded_ip3: IPvAnyAddressStr | None = None, - forwarded_ip4: IPvAnyAddressStr | None = None, - forwarded_ip5: IPvAnyAddressStr | None = None, - forwarded_ip6: IPvAnyAddressStr | None = None, - ) -> IPRecord: - return self.create( - user_id=user_id, - ip=ip or fake.ipv4_public(), - forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), - forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None), - forwarded_ip3=( - forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None - ), - forwarded_ip4=forwarded_ip4, - forwarded_ip5=forwarded_ip5, - forwarded_ip6=forwarded_ip6, - ) - def create_unpack( self, user_id: PositiveInt, @@ -348,29 +322,6 @@ class IPRecordManager(PostgresManagerWithRedis): class AuditLogManager(PostgresManager): - def create_dummy( - self, - user_id: PositiveInt, - level: AuditLogLevel | None = None, - event_type: str | None = None, - event_msg: str | None = None, - event_value: float | None = None, - ) -> AuditLog: - - event_types = { - "offerwall-enter.blocked", - "offerwall-enter.rate-limited", - "offerwall-enter.url-modified", - } - - return self.create( - user_id=user_id, - level=level or rchoice(list(AuditLogLevel)), - event_type=event_type or rchoice(list(event_types)), - event_msg=event_msg, - event_value=event_value, - ) - def create( self, user_id: PositiveInt, @@ -392,10 +343,9 @@ class AuditLogManager(PostgresManager): } ) - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute( - query=""" + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute( + query=""" INSERT INTO userhealth_auditlog (user_id, created, level, event_type, event_msg, event_value) @@ -403,10 +353,10 @@ class AuditLogManager(PostgresManager): %(event_type)s, %(event_msg)s, %(event_value)s) RETURNING id; """, - params=al.model_dump_mysql(), - ) - pk = c.fetchone()["id"] # type: ignore - conn.commit() + params=al.model_dump_mysql(), + ) + pk = c.fetchone()["id"] # type: ignore + conn.commit() al.id = pk return al @@ -536,7 +486,7 @@ class AuditLogManager(PostgresManager): level_ge: int | None = None, event_type: str | None = None, event_type_like: str | None = None, - event_msg: str | Nond = None, + event_msg: str | None = None, created_after: datetime | None = None, ) -> tuple[str, dict[str, Any]]: assert user_ids, "must pass at least 1 user_id" diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index de7b599..c2eb821 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -4,9 +4,8 @@ import logging from collections import defaultdict from collections.abc import Collection from datetime import datetime, timedelta, timezone -from decimal import ROUND_DOWN, Decimal +from decimal import Decimal from functools import cached_property -from random import choice as rchoice from uuid import uuid4 from faker import Faker @@ -94,51 +93,6 @@ class WallManager(PostgresManager): self.pg_config.execute_write(query=query, params=d) return wall - def create_dummy( - self, - session_id: int | None = None, - user_id: int | None = None, - started: datetime | None = None, - source: Source | None = None, - req_survey_id: str | None = None, - req_cpi: Decimal | None = None, - buyer_id: str | None = None, - uuid_id: str | None = None, - ): - """To be used in tests, where we don't care about certain fields""" - - user_id = user_id or fake.random_int(min=1, max=2_147_483_648) - started = started or fake.date_time_between( - start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), - end_date=datetime.now(tz=timezone.utc), - tzinfo=timezone.utc, - ) - - if session_id is None: - from generalresearch.managers.thl.session import SessionManager - - session = SessionManager(pg_config=self.pg_config).create_dummy( - started=started - ) - session_id = session.id - - source = source or rchoice(list(Source)) - req_survey_id = req_survey_id or uuid4().hex - req_cpi = req_cpi or Decimal(fake.random_int(min=1, max=150) / 100).quantize( - Decimal(".01"), rounding=ROUND_DOWN - ) - - return self.create( - session_id=session_id, - user_id=user_id, - started=started, - source=source, - req_survey_id=req_survey_id, - req_cpi=req_cpi, - buyer_id=buyer_id, - uuid_id=uuid_id, - ) - def get_from_uuid(self, wall_uuid: UUIDStr) -> Wall: query = """ SELECT @@ -230,8 +184,6 @@ class WallManager(PostgresManager): assert c.rowcount == 1 conn.commit() - return None - def get_wall_events( self, session_id: PositiveInt | None = None, @@ -365,7 +317,6 @@ class WallManager(PostgresManager): c.execute(query=query, params=params) assert c.rowcount == 1 conn.commit() - return None def filter_count_attempted_live(self, user_id: int) -> int: """ diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index cbb4bcb..70aafc6 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -30,7 +30,7 @@ from generalresearch.models.custom_types import ( UUIDStr, UUIDStrCoerce, ) -from generalresearch.models.thl.finance import POPFinancial, BusinessBalances +from generalresearch.models.thl.finance import BusinessBalances, POPFinancial from generalresearch.models.thl.ledger import LedgerAccount, OrderBy from generalresearch.models.thl.payout import BusinessPayoutEvent from generalresearch.pg_helper import PostgresConfig diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index 967f208..e946e5f 100644 --- a/generalresearch/models/network/nmap/parser.py +++ b/generalresearch/models/network/nmap/parser.py @@ -410,5 +410,5 @@ class NmapXmlParser: ) -def parse_nmap_xml(raw): +def parse_nmap_xml(raw) -> NmapResult: return NmapXmlParser.parse_xml(raw) diff --git a/generalresearch/models/network/rdns/parser.py b/generalresearch/models/network/rdns/parser.py index e1cf023..31a5ed6 100644 --- a/generalresearch/models/network/rdns/parser.py +++ b/generalresearch/models/network/rdns/parser.py @@ -7,7 +7,7 @@ from generalresearch.models.network.rdns.result import RDNSResult PTR_RE = re.compile(r"\sPTR\s+([^\s]+)\.") -def parse_rdns_output(ip: IPvAnyAddressStr, raw: str): +def parse_rdns_output(ip: IPvAnyAddressStr, raw: str) -> RDNSResult: hostnames: list[str] = [] for line in raw.splitlines(): diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 9c80065..cd7f282 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -1,29 +1,27 @@ +from __future__ import annotations + import os import shutil +import stat +import subprocess +import tempfile from datetime import datetime, timedelta, timezone from os.path import join as pjoin from pathlib import Path -from typing import TYPE_CHECKING, Callable, Generator +from typing import Callable, Generator from uuid import uuid4 import pytest -import redis from _pytest.config import Config from dotenv import load_dotenv from pydantic import MariaDBDsn, PostgresDsn, TypeAdapter -from pydantic_core import MultiHostHost -from redis import Redis +from generalresearch.config import GRLBaseSettings +from generalresearch.currency import USDCent from generalresearch.models.custom_types import InternalHostname, PostgresDict from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig from generalresearch.sql_helper import SqlHelper -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - from generalresearch.currency import USDCent - from generalresearch.models.thl.session import Status - @pytest.fixture(scope="session") def env_file_path(pytestconfig: Config) -> Path: @@ -44,7 +42,7 @@ def env_file_path(pytestconfig: Config) -> Path: @pytest.fixture(scope="session") -def settings(env_file_path: Path) -> "GRLBaseSettings": +def settings(env_file_path: Path) -> GRLBaseSettings: from generalresearch.config import GRLBaseSettings s = GRLBaseSettings() @@ -64,7 +62,7 @@ def settings(env_file_path: Path) -> "GRLBaseSettings": @pytest.fixture(scope="session") -def postgres_instance(settings: "GRLBaseSettings") -> Generator[PostgresDsn]: +def postgres_instance(settings: GRLBaseSettings) -> Generator[PostgresDsn]: """Create a ephemeral postgresql instance for us to use during pytest. This does not create any tables, or schema definitions within the instance. @@ -153,13 +151,61 @@ def postgres_instance_host( yield value +# @pytest.fixture(scope="session") +# def git_key_path(settings: GRLBaseSettings) -> Path: +# return Path('/tmp/') + + +@pytest.fixture(scope="session") +def git_key_path( + settings: GRLBaseSettings, +) -> Generator[Path]: + + assert settings.git_creds + with tempfile.NamedTemporaryFile(mode="w", delete=False, suffix="_id_rsa") as f: + f.write(settings.git_creds) + key_path = f.name + + os.chmod(key_path, stat.S_IRUSR | stat.S_IWUSR) + + yield Path(key_path) + + os.unlink(key_path) + + +@pytest.fixture(scope="session") +def gr_models(git_key_path: Path) -> Callable[..., Path]: + repo_url = "ssh://code.g-r-l.com/general-research/gr-carer.git" + repo_path = Path("/tmp/gr-carer") + + def _inner() -> Path: + ssh_cmd = ( + f"ssh -i {git_key_path} " + "-o IdentitiesOnly=yes " + "-o StrictHostKeyChecking=no " # or accept-new, see note below + ) + env = {"GIT_SSH_COMMAND": ssh_cmd} + + if repo_path.exists(): + subprocess.run(["git", "-C", str(repo_path), "pull"], check=True, env=env) + else: + subprocess.run( + ["git", "clone", "--depth", "1", repo_url, str(repo_path)], + check=True, + env=env, + ) + + return repo_path + + return _inner + + @pytest.fixture(scope="session") def django_db_factory( postgres_instance: PostgresDsn, postgres_instance_dict: PostgresDict ) -> Callable[..., PostgresDsn]: import django - from django.apps import apps from django.conf import settings as django_settings from django.core.management import call_command @@ -199,46 +245,7 @@ def django_db_factory( @pytest.fixture(scope="session") -def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: - - return PostgresConfig( - dsn=django_db_factory("generalresearch.thl_django"), - connect_timeout=1, - statement_timeout=5, - ) - - -@pytest.fixture(scope="session") -def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: - return thl_web_rr - - -@pytest.fixture(scope="session") -def gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: - - return PostgresConfig( - dsn=django_db_factory("gr_carer"), - connect_timeout=1, - statement_timeout=5, - ) - - -@pytest.fixture(scope="session") -def grliq_db(postgres_instance: PostgresDsn) -> PostgresConfig: - - # test_words = {"localhost", "127.0.0.1", "unittest", "grliq-test"} - # assert any(w in str(postgres_config.dsn) for w in test_words), "check grliq postgres_config" - # assert "grliqdeceezpocymo" not in str(postgres_config.dsn), "check grliq postgres_config" - - return PostgresConfig( - dsn=postgres_instance, - connect_timeout=1, - statement_timeout=5, - ) - - -@pytest.fixture(scope="session") -def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper: +def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: dsn = settings.spectrum_rw_db assert dsn assert dsn.path @@ -254,93 +261,14 @@ def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper: ) -@pytest.fixture(scope="session") -def thl_redis(settings: "GRLBaseSettings") -> "Redis": - # todo: this should get replaced with redisconfig (in most places) - # I'm not sure where this would be? in the domain name? - assert "unittest" in str(settings.thl_redis) or "127.0.0.1" in str( - settings.thl_redis - ) - - return redis.Redis.from_url( - **{ - "url": str(settings.thl_redis), - "decode_responses": True, - "socket_timeout": settings.redis_timeout, - "socket_connect_timeout": settings.redis_timeout, - } - ) - - -@pytest.fixture(scope="session") -def thl_redis_config(settings: "GRLBaseSettings") -> RedisConfig: - assert "unittest" in str(settings.thl_redis) or "127.0.0.1" in str( - settings.thl_redis - ) - return RedisConfig( - dsn=settings.thl_redis, - decode_responses=True, - socket_timeout=settings.redis_timeout, - socket_connect_timeout=settings.redis_timeout, - ) - - -@pytest.fixture(scope="session") -def gr_redis_config(settings: "GRLBaseSettings") -> "RedisConfig": - assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) - - return RedisConfig( - dsn=settings.gr_redis, - decode_responses=True, - socket_timeout=settings.redis_timeout, - socket_connect_timeout=settings.redis_timeout, - ) - - -@pytest.fixture(scope="session") -def gr_redis(settings: "GRLBaseSettings") -> "Redis": - assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) - return redis.Redis.from_url( - **{ - "url": str(settings.gr_redis), - "decode_responses": True, - "socket_timeout": settings.redis_timeout, - "socket_connect_timeout": settings.redis_timeout, - } - ) - - -@pytest.fixture -def gr_redis_async(settings: "GRLBaseSettings"): - assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) - - import redis.asyncio as redis_async - - return redis_async.Redis.from_url( - str(settings.gr_redis), - decode_responses=True, - socket_timeout=0.20, - socket_connect_timeout=0.20, - ) - - # === Random helpers === @pytest.fixture -def start() -> "datetime": - from datetime import datetime, timezone - +def start() -> datetime: return datetime(year=1900, month=1, day=1, tzinfo=timezone.utc) -@pytest.fixture -def wall_status(request) -> "Status": - from generalresearch.models.thl.session import Status - - return request.param if hasattr(request, "wall_status") else Status.COMPLETE - - @pytest.fixture def utc_now() -> datetime: return datetime.now(tz=timezone.utc) @@ -352,30 +280,22 @@ def utc_hour_ago() -> datetime: @pytest.fixture -def utc_day_ago() -> "datetime": - from datetime import datetime, timedelta, timezone - +def utc_day_ago() -> datetime: return datetime.now(tz=timezone.utc) - timedelta(hours=24) @pytest.fixture -def utc_90days_ago() -> "datetime": - from datetime import datetime, timedelta, timezone - +def utc_90days_ago() -> datetime: return datetime.now(tz=timezone.utc) - timedelta(days=90) @pytest.fixture -def utc_60days_ago() -> "datetime": - from datetime import datetime, timedelta, timezone - +def utc_60days_ago() -> datetime: return datetime.now(tz=timezone.utc) - timedelta(days=60) @pytest.fixture -def utc_30days_ago() -> "datetime": - from datetime import datetime, timedelta, timezone - +def utc_30days_ago() -> datetime: return datetime.now(tz=timezone.utc) - timedelta(days=30) @@ -425,6 +345,8 @@ def delete_df_collection( ) case _: + assert coll.data_type + thl_web_rw.execute_write( query=f"DELETE FROM {coll.data_type.value};", ) @@ -436,23 +358,23 @@ def delete_df_collection( @pytest.fixture(scope="function") -def amount_1(request) -> "USDCent": - from generalresearch.currency import USDCent - +def amount_1() -> USDCent: return USDCent(1) @pytest.fixture(scope="function") -def amount_100(request) -> "USDCent": - from generalresearch.currency import USDCent - +def amount_100() -> USDCent: return USDCent(100) -def clear_directory(path: Path): - for entry in os.listdir(path): +def clear_directory(path: Path | str): + dir_path = Path(path) + + for entry in os.listdir(dir_path): + full_path = os.path.join(path, entry) if os.path.isfile(full_path) or os.path.islink(full_path): os.unlink(full_path) # remove file or symlink + elif os.path.isdir(full_path): shutil.rmtree(full_path) # remove folder diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index edd777e..e8175a5 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -1,23 +1,83 @@ +from __future__ import annotations + from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Optional +from typing import Callable from uuid import uuid4 import pytest +from pydantic import PostgresDsn + +from generalresearch.config import GRLBaseSettings +from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA +from generalresearch.grliq.managers.forensic_data import ( + GrlIqDataManager, +) +from generalresearch.grliq.managers.forensic_events import ( + GrlIqEventManager, +) +from generalresearch.grliq.managers.forensic_results import ( + GrlIqCategoryResultsReader, +) +from generalresearch.grliq.models.forensic_data import GrlIqData +from generalresearch.pg_helper import PostgresConfig -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - from generalresearch.grliq.models.forensic_data import GrlIqData +# === Miscellaneous === @pytest.fixture(scope="function") -def mnt_grliq_archive_dir(settings: "GRLBaseSettings") -> Optional[str]: +def mnt_grliq_archive_dir(settings: GRLBaseSettings) -> str | None: return settings.mnt_grliq_archive_dir +@pytest.fixture(scope="session") +def grliq_db(postgres_instance: PostgresDsn) -> PostgresConfig: + # TODO: This will need to specificy a different DATABASE on the + # Postgres SERVER. That selection process will also need to + # selectively migrate only the tables from grliq + + return PostgresConfig( + dsn=postgres_instance, + connect_timeout=1, + statement_timeout=5, + ) + + +# === Managers === + + +@pytest.fixture(scope="session") +def grliq_dm(grliq_db: PostgresConfig) -> GrlIqDataManager: + assert grliq_db.dsn.path + assert "/unittest-" in grliq_db.dsn.path + return GrlIqDataManager(postgres_config=grliq_db) + + +@pytest.fixture(scope="session") +def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager: + assert grliq_db.dsn.path + assert "/unittest-" in grliq_db.dsn.path + + from generalresearch.grliq.managers.forensic_events import ( + GrlIqEventManager, + ) + + return GrlIqEventManager(postgres_config=grliq_db) + + +@pytest.fixture(scope="session") +def grliq_crr(grliq_db: PostgresConfig) -> GrlIqCategoryResultsReader: + assert grliq_db.dsn.path + assert "/unittest-" in grliq_db.dsn.path + + return GrlIqCategoryResultsReader(postgres_config=grliq_db) + + +# === Models === + + @pytest.fixture(scope="function") -def grliq_data() -> "GrlIqData": +def grliq_data() -> GrlIqData: from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA - from generalresearch.grliq.models.forensic_data import GrlIqData g: GrlIqData = DUMMY_GRLIQ_DATA[1]["data"] @@ -26,3 +86,53 @@ def grliq_data() -> "GrlIqData": g.created_at = datetime.now(tz=timezone.utc) g.timestamp = g.created_at - timedelta(seconds=10) return g + + +@pytest.fixture +def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]: + + def _inner( + is_attempt_allowed: bool = True, + product_id: str | None = None, + product_user_id: str | None = None, + uuid: str | None = None, + mid: str | None = None, + created_at: datetime | None = None, + ) -> GrlIqData: + """ + Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), + and GrlIqForensicCategoryResult (category_results) + :param is_attempt_allowed: Whether the attempt is allowed. + :param product_id: product_id of user + :param product_user_id: product_user_id of user + :param uuid: uuid for the grliq data record + :param mid: the thl_session:uuid / mid for the attempt. + :return: + """ + import copy + + res: GrlIqData = copy.deepcopy(DUMMY_GRLIQ_DATA[int(is_attempt_allowed)]) + + product_id = product_id or uuid4().hex + product_user_id = product_user_id or uuid4().hex + uuid = uuid or uuid4().hex + mid = mid or uuid4().hex + created_at = created_at or datetime.now(tz=timezone.utc) + + res["data"].product_id = product_id + res["data"].product_user_id = product_user_id + res["data"].uuid = uuid + res["data"].mid = mid + res["data"].created_at = created_at + res["result_data"].uuid = uuid + res["category_result"].uuid = uuid + + return grliq_dm.create( + iq_data=res["data"], + result_data=res["result_data"], + category_result=res["category_result"], + fraud_score=res["category_result"].fraud_score, + is_attempt_allowed=res["category_result"].is_attempt_allowed(), + ) + + return _inner diff --git a/test_utils/grliq/managers/__init__.py b/test_utils/grliq/managers/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/grliq/managers/conftest.py b/test_utils/grliq/managers/conftest.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/grliq/models/__init__.py b/test_utils/grliq/models/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/grliq/models/conftest.py b/test_utils/grliq/models/conftest.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py index 74e4081..88eef72 100644 --- a/test_utils/incite/collections/conftest.py +++ b/test_utils/incite/collections/conftest.py @@ -1,5 +1,7 @@ +from __future__ import annotations + from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Callable, Optional +from typing import TYPE_CHECKING, Callable import pytest @@ -21,12 +23,12 @@ if TYPE_CHECKING: @pytest.fixture def user_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, thl_web_rr: PostgresConfig, -) -> "UserDFCollection": +) -> UserDFCollection: from generalresearch.incite.collections.thl_web import ( DFCollectionType, UserDFCollection, @@ -43,12 +45,12 @@ def user_collection( @pytest.fixture def wall_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, thl_web_rr: PostgresConfig, -) -> "WallDFCollection": +) -> WallDFCollection: from generalresearch.incite.collections.thl_web import ( DFCollectionType, WallDFCollection, @@ -65,12 +67,12 @@ def wall_collection( @pytest.fixture def session_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, thl_web_rr: PostgresConfig, -) -> "SessionDFCollection": +) -> SessionDFCollection: from generalresearch.incite.collections.thl_web import ( DFCollectionType, SessionDFCollection, @@ -103,12 +105,12 @@ def session_collection( @pytest.fixture def task_adj_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, - duration: Optional[timedelta], + duration: timedelta | None, start: datetime, thl_web_rr: PostgresConfig, -) -> "TaskAdjustmentDFCollection": +) -> TaskAdjustmentDFCollection: from generalresearch.incite.collections.thl_web import ( DFCollectionType, TaskAdjustmentDFCollection, @@ -127,12 +129,12 @@ def task_adj_collection( @pytest.fixture def auditlog_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, thl_web_rr: PostgresConfig, -) -> "AuditLogDFCollection": +) -> AuditLogDFCollection: from generalresearch.incite.collections.thl_web import ( AuditLogDFCollection, DFCollectionType, @@ -149,12 +151,12 @@ def auditlog_collection( @pytest.fixture def ledger_collection( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, thl_web_rr: PostgresConfig, -) -> "LedgerDFCollection": +) -> LedgerDFCollection: from generalresearch.incite.collections.thl_web import ( DFCollectionType, LedgerDFCollection, @@ -171,7 +173,7 @@ def ledger_collection( @pytest.fixture def rm_ledger_collection( - ledger_collection: "LedgerDFCollection", + ledger_collection: LedgerDFCollection, ) -> Callable[..., None]: def _inner(): @@ -187,13 +189,13 @@ def rm_ledger_collection( @pytest.fixture def df_collection( - mnt_filepath: "GRLDatasets", - df_collection_data_type: "DFCollectionType", + mnt_filepath: GRLDatasets, + df_collection_data_type: DFCollectionType, offset: str, duration: timedelta, utc_90days_ago: datetime, thl_web_rr: PostgresConfig, -) -> "DFCollection": +) -> DFCollection: from generalresearch.incite.collections import DFCollection start = utc_90days_ago.replace(microsecond=0) diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 0e2f7bd..12e57c5 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -1,9 +1,11 @@ +from __future__ import annotations + from datetime import datetime, timedelta, timezone from os.path import join as pjoin from pathlib import Path from random import choice as randchoice from shutil import rmtree -from typing import TYPE_CHECKING, Callable, Optional +from typing import TYPE_CHECKING, Callable from uuid import uuid4 import pytest @@ -30,7 +32,7 @@ fake = Faker() @pytest.fixture -def mnt_gr_api_dir(request: SubRequest, settings: "GRLBaseSettings") -> Path: +def mnt_gr_api_dir(request: SubRequest, settings: GRLBaseSettings) -> Path: p = Path(settings.mnt_gr_api_dir) p.mkdir(parents=True, exist_ok=True) @@ -53,7 +55,7 @@ def mnt_gr_api_dir(request: SubRequest, settings: "GRLBaseSettings") -> Path: @pytest.fixture -def event_report_request(utc_hour_ago: datetime, start: datetime) -> "ReportRequest": +def event_report_request(utc_hour_ago: datetime, start: datetime) -> ReportRequest: from generalresearch.models.admin.request import ( ReportRequest, ReportType, @@ -69,7 +71,7 @@ def event_report_request(utc_hour_ago: datetime, start: datetime) -> "ReportRequ @pytest.fixture -def session_report_request(utc_hour_ago: datetime, start: datetime) -> "ReportRequest": +def session_report_request(utc_hour_ago: datetime, start: datetime) -> ReportRequest: from generalresearch.models.admin.request import ( ReportRequest, ReportType, @@ -85,7 +87,7 @@ def session_report_request(utc_hour_ago: datetime, start: datetime) -> "ReportRe @pytest.fixture -def mnt_filepath(request: SubRequest) -> "GRLDatasets": +def mnt_filepath(request: SubRequest) -> GRLDatasets: """ Creates a temporary file path for all DFCollections & Mergers parquet files. @@ -111,7 +113,7 @@ def mnt_filepath(request: SubRequest) -> "GRLDatasets": @pytest.fixture -def start(utc_90days_ago: datetime) -> "datetime": +def start(utc_90days_ago: datetime) -> datetime: s = utc_90days_ago.replace(microsecond=0) return s @@ -122,19 +124,19 @@ def offset() -> str: @pytest.fixture -def duration() -> Optional["timedelta"]: +def duration() -> timedelta | None: return timedelta(hours=1) @pytest.fixture -def df_collection_data_type() -> "DFCollectionType": +def df_collection_data_type() -> DFCollectionType: from generalresearch.incite.collections import DFCollectionType return DFCollectionType.TEST @pytest.fixture -def merge_type() -> "MergeType": +def merge_type() -> MergeType: from generalresearch.incite.mergers import MergeType return MergeType.TEST @@ -142,16 +144,16 @@ def merge_type() -> "MergeType": @pytest.fixture def incite_item_factory( - session_factory: Callable[..., "Session"], - product: "Product", - user_factory: Callable[..., "User"], - session_with_tx_factory: Callable[..., "Session"], + session_factory: Callable[..., Session], + product: Product, + user_factory: Callable[..., User], + session_with_tx_factory: Callable[..., Session], ) -> Callable[..., None]: def _inner( - item: "DFCollectionItem", + item: DFCollectionItem, observations: int = 3, - user: Optional["User"] = None, + user: User | None = None, ): from generalresearch.incite.collections import ( DFCollection, @@ -201,6 +203,4 @@ def incite_item_factory( case _: raise ValueError("Unsupported DFCollectionItem") - return None - return _inner diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py index d094b84..e9970c2 100644 --- a/test_utils/incite/mergers/conftest.py +++ b/test_utils/incite/mergers/conftest.py @@ -1,46 +1,45 @@ +from __future__ import annotations + from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Callable +from typing import Callable import pytest +from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.mergers import MergeType +from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, +) +from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( + EnrichedTaskAdjustMerge, +) +from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, +) +from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, +) +from generalresearch.incite.mergers.pop_ledger import ( + PopLedgerMerge, + PopLedgerMergeItem, +) +from generalresearch.incite.mergers.ym_survey_wall import ( + YMSurveyWallMerge, + YMSurveyWallMergeCollectionItem, +) +from generalresearch.incite.mergers.ym_wall_summary import ( + YMWallSummaryMerge, + YMWallSummaryMergeItem, +) from test_utils.conftest import clear_directory -if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets - from generalresearch.incite.mergers import MergeType - from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, - ) - from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( - EnrichedTaskAdjustMerge, - ) - from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, - ) - from generalresearch.incite.mergers.foundations.user_id_product import ( - UserIdProductMerge, - ) - from generalresearch.incite.mergers.pop_ledger import ( - PopLedgerMerge, - PopLedgerMergeItem, - ) - from generalresearch.incite.mergers.ym_survey_wall import ( - YMSurveyWallMerge, - YMSurveyWallMergeCollectionItem, - ) - from generalresearch.incite.mergers.ym_wall_summary import ( - YMWallSummaryMerge, - YMWallSummaryMergeItem, - ) - - # -------------------------- # Merges # -------------------------- @pytest.fixture -def rm_pop_ledger_merge(pop_ledger_merge: "PopLedgerMerge") -> Callable[..., None]: +def rm_pop_ledger_merge(pop_ledger_merge: PopLedgerMerge) -> Callable[..., None]: def _inner(): clear_directory(pop_ledger_merge.archive_path) @@ -50,11 +49,11 @@ def rm_pop_ledger_merge(pop_ledger_merge: "PopLedgerMerge") -> Callable[..., Non @pytest.fixture def pop_ledger_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, start: datetime, duration: timedelta, -) -> "PopLedgerMerge": +) -> PopLedgerMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge @@ -70,8 +69,8 @@ def pop_ledger_merge( @pytest.fixture def pop_ledger_merge_item( start: datetime, - pop_ledger_merge: "PopLedgerMerge", -) -> "PopLedgerMergeItem": + pop_ledger_merge: PopLedgerMerge, +) -> PopLedgerMergeItem: from generalresearch.incite.mergers.pop_ledger import PopLedgerMergeItem @@ -83,9 +82,9 @@ def pop_ledger_merge_item( @pytest.fixture def ym_survey_wall_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, start: datetime, -) -> "YMSurveyWallMerge": +) -> YMSurveyWallMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge @@ -98,8 +97,8 @@ def ym_survey_wall_merge( @pytest.fixture def ym_survey_wall_merge_item( - start: datetime, ym_survey_wall_merge: "YMSurveyWallMerge" -) -> "YMSurveyWallMergeCollectionItem": + start: datetime, ym_survey_wall_merge: YMSurveyWallMerge +) -> YMSurveyWallMergeCollectionItem: from generalresearch.incite.mergers.ym_survey_wall import ( YMSurveyWallMergeCollectionItem, ) @@ -112,11 +111,11 @@ def ym_survey_wall_merge_item( @pytest.fixture def ym_wall_summary_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, -) -> "YMWallSummaryMerge": +) -> YMWallSummaryMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.ym_wall_summary import YMWallSummaryMerge @@ -129,8 +128,8 @@ def ym_wall_summary_merge( def ym_wall_summary_merge_item( - start: datetime, ym_wall_summary_merge: "YMWallSummaryMerge" -) -> "YMWallSummaryMergeItem": + start: datetime, ym_wall_summary_merge: YMWallSummaryMerge +) -> YMWallSummaryMergeItem: from generalresearch.incite.mergers.ym_wall_summary import ( YMWallSummaryMergeItem, ) @@ -148,11 +147,11 @@ def ym_wall_summary_merge_item( @pytest.fixture def enriched_session_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, -) -> "EnrichedSessionMerge": +) -> EnrichedSessionMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.foundations.enriched_session import ( EnrichedSessionMerge, @@ -168,11 +167,11 @@ def enriched_session_merge( @pytest.fixture def enriched_task_adjust_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, -) -> "EnrichedTaskAdjustMerge": +) -> EnrichedTaskAdjustMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( EnrichedTaskAdjustMerge, @@ -190,11 +189,11 @@ def enriched_task_adjust_merge( @pytest.fixture def enriched_wall_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, offset: str, duration: timedelta, start: datetime, -) -> "EnrichedWallMerge": +) -> EnrichedWallMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMerge, @@ -210,11 +209,11 @@ def enriched_wall_merge( @pytest.fixture def user_id_product_merge( - mnt_filepath: "GRLDatasets", + mnt_filepath: GRLDatasets, duration: timedelta, offset: str, start: datetime, -) -> "UserIdProductMerge": +) -> UserIdProductMerge: from generalresearch.incite.mergers import MergeType from generalresearch.incite.mergers.foundations.user_id_product import ( UserIdProductMerge, @@ -235,8 +234,8 @@ def user_id_product_merge( @pytest.fixture def merge_collection( - mnt_filepath: "GRLDatasets", - merge_type: "MergeType", + mnt_filepath: GRLDatasets, + merge_type: MergeType, offset: str, duration: timedelta, start: datetime, diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index c8a6e2f..d2e5d20 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -1,9 +1,33 @@ -from typing import TYPE_CHECKING, Callable +from __future__ import annotations + +from typing import Callable import pytest -from generalresearch.managers.base import Permission +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import ( + MembershipManager, + TeamManager, +) +from generalresearch.managers.spectrum.survey import SpectrumSurveyManager +from generalresearch.managers.thl.buyer import BuyerManager +from generalresearch.managers.thl.ipinfo import ( + GeoIpInfoManager, + IPGeonameManager, + IPInformationManager, +) +from generalresearch.managers.thl.profiling.uqa import UQAManager +from generalresearch.managers.thl.userhealth import ( + AuditLogManager, + IPRecordManager, + UserIpHistoryManager, +) from generalresearch.models import Source +from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig from generalresearch.sql_helper import SqlHelper @@ -11,437 +35,11 @@ from test_utils.managers.cashout_methods import ( EXAMPLE_TANGO_CASHOUT_METHODS, ) -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - from generalresearch.grliq.managers.forensic_data import ( - GrlIqDataManager, - ) - from generalresearch.grliq.managers.forensic_events import ( - GrlIqEventManager, - ) - from generalresearch.grliq.managers.forensic_results import ( - GrlIqCategoryResultsReader, - ) - from generalresearch.managers.gr.authentication import ( - GRTokenManager, - GRUserManager, - ) - from generalresearch.managers.gr.business import ( - BusinessAddressManager, - BusinessBankAccountManager, - BusinessManager, - ) - from generalresearch.managers.gr.team import ( - MembershipManager, - TeamManager, - ) - from generalresearch.managers.thl.buyer import BuyerManager - from generalresearch.managers.thl.category import CategoryManager - from generalresearch.managers.thl.contest_manager import ContestManager - from generalresearch.managers.thl.ipinfo import ( - GeoIpInfoManager, - IPGeonameManager, - IPInformationManager, - ) - from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerAccountManager, - LedgerManager, - LedgerTransactionManager, - ) - from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, - ) - from generalresearch.managers.thl.maxmind import MaxmindManager - from generalresearch.managers.thl.maxmind.basic import ( - MaxmindBasicManager, - ) - from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - BusinessPayoutEventManager, - PayoutEventManager, - UserPayoutEventManager, - ) - from generalresearch.managers.thl.product import ProductManager - from generalresearch.managers.thl.session import SessionManager - from generalresearch.managers.thl.task_adjustment import ( - TaskAdjustmentManager, - ) - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) - from generalresearch.managers.thl.user_manager.user_metadata_manager import ( - UserMetadataManager, - ) - from generalresearch.managers.thl.userhealth import ( - AuditLogManager, - IPRecordManager, - UserIpHistoryManager, - ) - from generalresearch.managers.thl.wall import ( - WallCacheManager, - WallManager, - ) - from generalresearch.models.thl.user import User - - # === THL === @pytest.fixture(scope="session") -def ltxm( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "LedgerTransactionManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerTransactionManager, - ) - - return LedgerTransactionManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - testing=True, - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def lam( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "LedgerAccountManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerAccountManager, - ) - - return LedgerAccountManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - testing=True, - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def lm(thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig) -> "LedgerManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ledger_manager.ledger import ( - LedgerManager, - ) - - return LedgerManager( - pg_config=thl_web_rw, - permissions=[ - Permission.CREATE, - Permission.READ, - Permission.UPDATE, - Permission.DELETE, - ], - testing=True, - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def thl_lm( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "ThlLedgerManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, - ) - - return ThlLedgerManager( - pg_config=thl_web_rw, - permissions=[ - Permission.CREATE, - Permission.READ, - Permission.UPDATE, - Permission.DELETE, - ], - testing=True, - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def payout_event_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "PayoutEventManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.payout import PayoutEventManager - - return PayoutEventManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def user_payout_event_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "UserPayoutEventManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.payout import UserPayoutEventManager - - return UserPayoutEventManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def brokerage_product_payout_event_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "BrokerageProductPayoutEventManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - ) - - return BrokerageProductPayoutEventManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def business_payout_event_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "BusinessPayoutEventManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.payout import ( - BusinessPayoutEventManager, - ) - - return BusinessPayoutEventManager( - pg_config=thl_web_rw, - permissions=[Permission.CREATE, Permission.READ], - redis_config=thl_redis_config, - ) - - -@pytest.fixture(scope="session") -def product_manager(thl_web_rw: PostgresConfig) -> "ProductManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.product import ProductManager - - return ProductManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def user_manager( - settings: "GRLBaseSettings", thl_web_rw: PostgresConfig, thl_web_rr: PostgresConfig -) -> "UserManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert thl_web_rr.dsn - assert thl_web_rr.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rr.dsn.path - - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) - - return UserManager( - pg_config=thl_web_rw, - pg_config_rr=thl_web_rr, - redis=settings.redis, - ) - - -@pytest.fixture(scope="session") -def user_metadata_manager(thl_web_rw: PostgresConfig) -> "UserMetadataManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.user_manager.user_metadata_manager import ( - UserMetadataManager, - ) - - return UserMetadataManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def session_manager(thl_web_rw: PostgresConfig) -> "SessionManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.session import SessionManager - - return SessionManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def wall_manager(thl_web_rw: PostgresConfig) -> "WallManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.wall import WallManager - - return WallManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def wall_cache_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "WallCacheManager": - # assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.wall import WallCacheManager - - return WallCacheManager(pg_config=thl_web_rw, redis_config=thl_redis_config) - - -@pytest.fixture(scope="session") -def task_adjustment_manager(thl_web_rw: PostgresConfig) -> "TaskAdjustmentManager": - # assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.task_adjustment import ( - TaskAdjustmentManager, - ) - - return TaskAdjustmentManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def contest_manager(thl_web_rw: PostgresConfig) -> "ContestManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.contest_manager import ContestManager - - return ContestManager( - pg_config=thl_web_rw, - permissions=[ - Permission.CREATE, - Permission.READ, - Permission.UPDATE, - Permission.DELETE, - ], - ) - - -@pytest.fixture(scope="session") -def category_manager(thl_web_rw: PostgresConfig) -> "CategoryManager": - assert thl_web_rw.dsn - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.category import CategoryManager - - return CategoryManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def buyer_manager(thl_web_rw: PostgresConfig) -> "BuyerManager": - # assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.buyer import BuyerManager - - return BuyerManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def survey_manager(thl_web_rw: PostgresConfig): - # assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.survey import SurveyManager - - return SurveyManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def surveystat_manager(thl_web_rw: PostgresConfig): - # assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.survey import SurveyStatManager - - return SurveyStatManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def surveypenalty_manager(thl_redis_config: RedisConfig): - from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager - - return SurveyPenaltyManager(redis_config=thl_redis_config) - - -@pytest.fixture(scope="session") -def upk_schema_manager(thl_web_rw: PostgresConfig): - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.profiling.schema import ( - UpkSchemaManager, - ) - - return UpkSchemaManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def user_upk_manager(thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig): - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.profiling.user_upk import ( - UserUpkManager, - ) - - return UserUpkManager(pg_config=thl_web_rw, redis_config=thl_redis_config) - - -@pytest.fixture(scope="session") -def question_manager(thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig): - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.profiling.question import ( - QuestionManager, - ) - - return QuestionManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def uqa_manager(thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig): - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.profiling.uqa import UQAManager - - return UQAManager(redis_config=thl_redis_config, pg_config=thl_web_rw) - - -@pytest.fixture(scope="function") -def uqa_manager_clear_cache(uqa_manager, user: "User"): - # On successive py-test/jenkins runs, the cache may contain - # the previous run's info (keyed under the same user_id) - uqa_manager.clear_cache(user) - yield - uqa_manager.clear_cache(user) - - -@pytest.fixture(scope="session") -def audit_log_manager(thl_web_rw: PostgresConfig) -> "AuditLogManager": +def audit_log_manager(thl_web_rw: PostgresConfig) -> AuditLogManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -451,7 +49,7 @@ def audit_log_manager(thl_web_rw: PostgresConfig) -> "AuditLogManager": @pytest.fixture(scope="session") -def ip_geoname_manager(thl_web_rw: PostgresConfig) -> "IPGeonameManager": +def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -461,7 +59,7 @@ def ip_geoname_manager(thl_web_rw: PostgresConfig) -> "IPGeonameManager": @pytest.fixture(scope="session") -def ip_information_manager(thl_web_rw: PostgresConfig) -> "IPInformationManager": +def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -473,7 +71,7 @@ def ip_information_manager(thl_web_rw: PostgresConfig) -> "IPInformationManager" @pytest.fixture(scope="session") def ip_record_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "IPRecordManager": +) -> IPRecordManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -485,7 +83,7 @@ def ip_record_manager( @pytest.fixture(scope="session") def user_iphistory_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "UserIpHistoryManager": +) -> UserIpHistoryManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -508,7 +106,7 @@ def user_iphistory_manager_clear_cache(user_iphistory_manager, user): @pytest.fixture(scope="session") def geoipinfo_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> "GeoIpInfoManager": +) -> GeoIpInfoManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path @@ -517,38 +115,6 @@ def geoipinfo_manager( return GeoIpInfoManager(pg_config=thl_web_rw, redis_config=thl_redis_config) -@pytest.fixture(scope="session") -def maxmind_basic_manager(settings: "GRLBaseSettings") -> "MaxmindBasicManager": - from generalresearch.managers.thl.maxmind.basic import ( - MaxmindBasicManager, - ) - - return MaxmindBasicManager( - data_dir="/tmp/", - maxmind_account_id=settings.maxmind_account_id, - maxmind_license_key=settings.maxmind_license_key, - ) - - -@pytest.fixture(scope="session") -def maxmind_manager( - settings: "GRLBaseSettings", - thl_web_rw: PostgresConfig, - thl_redis_config: RedisConfig, -) -> "MaxmindManager": - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.maxmind import MaxmindManager - - return MaxmindManager( - pg_config=thl_web_rw, - redis_config=thl_redis_config, - maxmind_account_id=settings.maxmind_account_id, - maxmind_license_key=settings.maxmind_license_key, - ) - - @pytest.fixture(scope="session") def cashout_method_manager(thl_web_rw: PostgresConfig): assert thl_web_rw.dsn.path @@ -592,7 +158,6 @@ def uqa_db_index(thl_web_rw: PostgresConfig): # except pymysql.OperationalError as e: # if "Duplicate key name 'idx_user_id'" not in str(e): # raise - return None @pytest.fixture(scope="session") @@ -606,9 +171,7 @@ def delete_cashoutmethod_db(thl_web_rw: PostgresConfig) -> Callable[..., None]: @pytest.fixture(scope="session") -def setup_cashoutmethod_db( - settings: "GRLBaseSettings", cashout_method_manager, delete_cashoutmethod_db -): +def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db): delete_cashoutmethod_db() for x in EXAMPLE_TANGO_CASHOUT_METHODS: cashout_method_manager.create(x) @@ -621,14 +184,12 @@ def setup_cashoutmethod_db( # cashout_method_manager.create(AMT_BONUS_CASHOUT_METHOD) raise NotImplementedError("Need to implement setup_cashoutmethod_db") - return None - # === THL: Marketplaces === @pytest.fixture(scope="session") -def spectrum_manager(spectrum_rw: SqlHelper) -> "SpectrumSurveyManager": +def spectrum_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: from generalresearch.managers.spectrum.survey import ( SpectrumSurveyManager, ) @@ -640,7 +201,7 @@ def spectrum_manager(spectrum_rw: SqlHelper) -> "SpectrumSurveyManager": @pytest.fixture(scope="session") def business_manager( gr_db: PostgresConfig, gr_redis_config: RedisConfig -) -> "BusinessManager": +) -> BusinessManager: from generalresearch.redis_helper import RedisConfig assert gr_db.dsn.path @@ -656,7 +217,7 @@ def business_manager( @pytest.fixture(scope="session") -def business_address_manager(gr_db: PostgresConfig) -> "BusinessAddressManager": +def business_address_manager(gr_db: PostgresConfig) -> BusinessAddressManager: assert gr_db.dsn.path assert "/unittest-" in gr_db.dsn.path @@ -668,7 +229,7 @@ def business_address_manager(gr_db: PostgresConfig) -> "BusinessAddressManager": @pytest.fixture(scope="session") def business_bank_account_manager( gr_db: PostgresConfig, -) -> "BusinessBankAccountManager": +) -> BusinessBankAccountManager: assert gr_db.dsn.path assert "/unittest-" in gr_db.dsn.path @@ -680,7 +241,7 @@ def business_bank_account_manager( @pytest.fixture(scope="session") -def team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> "TeamManager": +def team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> TeamManager: assert gr_db.dsn.path assert "/unittest-" in gr_db.dsn.path @@ -690,27 +251,7 @@ def team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> "TeamMa @pytest.fixture(scope="session") -def gr_um(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> "GRUserManager": - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.authentication import GRUserManager - - return GRUserManager(pg_config=gr_db, redis_config=gr_redis_config) - - -@pytest.fixture(scope="session") -def gr_tm(gr_db: PostgresConfig) -> "GRTokenManager": - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.authentication import GRTokenManager - - return GRTokenManager(pg_config=gr_db) - - -@pytest.fixture(scope="session") -def membership_manager(gr_db: PostgresConfig) -> "MembershipManager": +def membership_manager(gr_db: PostgresConfig) -> MembershipManager: assert gr_db.dsn.path assert "/unittest-" in gr_db.dsn.path @@ -719,47 +260,8 @@ def membership_manager(gr_db: PostgresConfig) -> "MembershipManager": return MembershipManager(pg_config=gr_db) -# === GRL IQ === - - -@pytest.fixture(scope="session") -def grliq_dm(grliq_db: PostgresConfig) -> "GrlIqDataManager": - assert grliq_db.dsn.path - assert "/unittest-" in grliq_db.dsn.path - - from generalresearch.grliq.managers.forensic_data import ( - GrlIqDataManager, - ) - - return GrlIqDataManager(postgres_config=grliq_db) - - -@pytest.fixture(scope="session") -def grliq_em(grliq_db: PostgresConfig) -> "GrlIqEventManager": - assert grliq_db.dsn.path - assert "/unittest-" in grliq_db.dsn.path - - from generalresearch.grliq.managers.forensic_events import ( - GrlIqEventManager, - ) - - return GrlIqEventManager(postgres_config=grliq_db) - - -@pytest.fixture(scope="session") -def grliq_crr(grliq_db: PostgresConfig) -> "GrlIqCategoryResultsReader": - assert grliq_db.dsn.path - assert "/unittest-" in grliq_db.dsn.path - - from generalresearch.grliq.managers.forensic_results import ( - GrlIqCategoryResultsReader, - ) - - return GrlIqCategoryResultsReader(postgres_config=grliq_db) - - @pytest.fixture(scope="session") -def delete_buyers_surveys(thl_web_rw: PostgresConfig, buyer_manager: "BuyerManager"): +def delete_buyers_surveys(thl_web_rw: PostgresConfig, buyer_manager: BuyerManager): # assert "/unittest-" in thl_web_rw.dsn.path thl_web_rw.execute_write( """ diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py index fb0b44b..67935e7 100644 --- a/test_utils/managers/contest/conftest.py +++ b/test_utils/managers/contest/conftest.py @@ -1,286 +1,24 @@ -from datetime import datetime, timezone -from decimal import Decimal -from typing import TYPE_CHECKING, Callable -from uuid import uuid4 - import pytest -from generalresearch.currency import USDCent - -if TYPE_CHECKING: - from generalresearch.managers.thl.contest_manager import ContestManager - from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager - from generalresearch.models.thl.contest.contest import Contest - from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContestCreate, - ) - from generalresearch.models.thl.contest.milestone import ( - MilestoneContestCreate, - ) - from generalresearch.models.thl.contest.raffle import ( - RaffleContestCreate, - ) - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.user import User - - -@pytest.fixture -def raffle_contest_create() -> "RaffleContestCreate": - from generalresearch.models.thl.contest import ( - ContestEndCondition, - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.raffle import ( - ContestEntryType, - RaffleContestCreate, - ) - - # This is what we'll get from the fastapi endpoint - return RaffleContestCreate( - name="test", - contest_type=ContestType.RAFFLE, - entry_type=ContestEntryType.CASH, - prizes=[ - ContestPrize( - name="iPod 64GB White", - kind=ContestPrizeKind.PHYSICAL, - estimated_cash_value=USDCent(100), - ) - ], - end_condition=ContestEndCondition(target_entry_amount=USDCent(100)), - ) - - -@pytest.fixture -def raffle_contest_in_db( - product_user_wallet_yes: "Product", - raffle_contest_create: "RaffleContestCreate", - contest_manager: "ContestManager", -) -> "Contest": - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create - ) - - -@pytest.fixture -def raffle_contest( - product_user_wallet_yes: "Product", raffle_contest_create: "RaffleContestCreate" -) -> "Contest": - from generalresearch.models.thl.contest.io import contest_create_to_contest - - return contest_create_to_contest( - product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create - ) - - -@pytest.fixture(scope="function") -def raffle_contest_factory( - product_user_wallet_yes: "Product", - raffle_contest_create: "RaffleContestCreate", - contest_manager: "ContestManager", -) -> Callable[..., "Contest"]: - - def _inner(**kwargs): - raffle_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=raffle_contest_create, - ) +from generalresearch.managers.base import Permission +from generalresearch.managers.thl.contest_manager import ContestManager +from generalresearch.pg_helper import PostgresConfig - return _inner - - -@pytest.fixture -def milestone_contest_create() -> "MilestoneContestCreate": - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, - MilestoneContestCreate, - MilestoneContestEndCondition, - ) - - # This is what we'll get from the fastapi endpoint - return MilestoneContestCreate( - name="Win a 50% bonus for 7 days and a $1 bonus after your first 3 completes!", - description="only valid for the first 5 users", - contest_type=ContestType.MILESTONE, - prizes=[ - ContestPrize( - name="50% for 7 days", - kind=ContestPrizeKind.PROMOTION, - estimated_cash_value=USDCent(0), - ), - ContestPrize( - name="$1 Bonus", - kind=ContestPrizeKind.CASH, - cash_amount=USDCent(1_00), - estimated_cash_value=USDCent(1_00), - ), - ], - end_condition=MilestoneContestEndCondition( - ends_at=datetime(year=2030, month=1, day=1, tzinfo=timezone.utc), - max_winners=5, - ), - entry_trigger=ContestEntryTrigger.TASK_COMPLETE, - target_amount=3, - ) +@pytest.fixture(scope="session") +def contest_manager(thl_web_rw: PostgresConfig) -> ContestManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path -@pytest.fixture -def milestone_contest_in_db( - product_user_wallet_yes: "Product", - milestone_contest_create: "MilestoneContestCreate", - contest_manager: "ContestManager", -) -> "Contest": - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create - ) - - -@pytest.fixture -def milestone_contest( - product_user_wallet_yes: "Product", - milestone_contest_create: "MilestoneContestCreate", -) -> "Contest": - from generalresearch.models.thl.contest.io import contest_create_to_contest - - return contest_create_to_contest( - product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create - ) - - -@pytest.fixture(scope="function") -def milestone_contest_factory( - product_user_wallet_yes: "Product", - milestone_contest_create: "MilestoneContestCreate", - contest_manager: "ContestManager", -) -> Callable[..., "Contest"]: - - def _inner(**kwargs): - milestone_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=milestone_contest_create, - ) - - return _inner - - -@pytest.fixture -def leaderboard_contest_create( - product_user_wallet_yes: "Product", -) -> "LeaderboardContestCreate": - from generalresearch.models.thl.contest import ( - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.leaderboard import ( - LeaderboardContestCreate, - ) + from generalresearch.managers.thl.contest_manager import ContestManager - # This is what we'll get from the fastapi endpoint - return LeaderboardContestCreate( - name="test", - contest_type=ContestType.LEADERBOARD, - prizes=[ - ContestPrize( - name="$15 Cash", - estimated_cash_value=USDCent(15_00), - cash_amount=USDCent(15_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=1, - ), - ContestPrize( - name="$10 Cash", - estimated_cash_value=USDCent(10_00), - cash_amount=USDCent(10_00), - kind=ContestPrizeKind.CASH, - leaderboard_rank=2, - ), + return ContestManager( + pg_config=thl_web_rw, + permissions=[ + Permission.CREATE, + Permission.READ, + Permission.UPDATE, + Permission.DELETE, ], - leaderboard_key=f"leaderboard:{product_user_wallet_yes.uuid}:us:daily:2025-01-01:complete_count", ) - - -@pytest.fixture -def leaderboard_contest_in_db( - product_user_wallet_yes: "Product", - leaderboard_contest_create: "LeaderboardContestCreate", - contest_manager: "ContestManager", -) -> "Contest": - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=leaderboard_contest_create, - ) - - -@pytest.fixture -def leaderboard_contest( - product_user_wallet_yes: "Product", - leaderboard_contest_create: "LeaderboardContestCreate", -): - from generalresearch.models.thl.contest.io import contest_create_to_contest - - return contest_create_to_contest( - product_id=product_user_wallet_yes.uuid, - contest_create=leaderboard_contest_create, - ) - - -@pytest.fixture(scope="function") -def leaderboard_contest_factory( - product_user_wallet_yes: "Product", - leaderboard_contest_create: "LeaderboardContestCreate", - contest_manager: "ContestManager", -) -> Callable[..., "Contest"]: - - def _inner(**kwargs): - leaderboard_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=leaderboard_contest_create, - ) - - return _inner - - -@pytest.fixture -def user_with_money( - request, - user_factory: Callable[..., "User"], - product_user_wallet_yes: "Product", - thl_lm: "ThlLedgerManager", -) -> "User": - from generalresearch.models.thl.user import User - - params = getattr(request, "param", dict()) or {} - min_balance = int(params.get("min_balance", USDCent(1_00))) - - user: User = user_factory(product=product_user_wallet_yes) - wallet = thl_lm.get_account_or_create_user_wallet(user) - balance = thl_lm.get_account_balance(wallet) - todo = min_balance - balance - if todo > 0: - # # Put money in user's wallet - thl_lm.create_tx_user_bonus( - user=user, - ref_uuid=uuid4().hex, - description="bonus", - amount=Decimal(todo) / 100, - ) - print(f"wallet balance: {thl_lm.get_user_wallet_balance(user=user)}") - - return user diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 3e3a4ad..37da164 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -1,151 +1,110 @@ from __future__ import annotations from typing import Callable -from uuid import uuid4 import pytest -from pydantic import PositiveInt -from pydantic_extra_types.phone_numbers import PhoneNumber +import redis.asyncio as redis_async +from pydantic import PostgresDsn +from redis import Redis -from generalresearch.managers.gr.authentication import GRUserManager +from generalresearch.config import GRLBaseSettings +from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, BusinessManager, ) -from generalresearch.managers.gr.team import TeamManager -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.gr.authentication import GRUser -from generalresearch.models.gr.business import ( - Business, - BusinessAddress, - BusinessBankAccount, - BusinessType, - TransferMethod, -) -from generalresearch.models.gr.team import Team +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig + + +# === Msc === +@pytest.fixture(scope="session") +def gr_redis(settings: GRLBaseSettings) -> Redis: + assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) + return Redis.from_url( + url=str(settings.gr_redis), + decode_responses=True, + socket_timeout=settings.redis_timeout, + socket_connect_timeout=settings.redis_timeout, + ) @pytest.fixture -def gr_user_factory(gr_um: GRUserManager) -> Callable[..., GRUser]: +def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis: + assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) - def _inner( - sub: str | None = None, - is_superuser: bool = False, - ) -> GRUser: - sub = sub or f"{uuid4().hex}-{uuid4().hex}" + return redis_async.Redis.from_url( + str(settings.gr_redis), + decode_responses=True, + socket_timeout=0.20, + socket_connect_timeout=0.20, + ) - return gr_um.create( - sub=sub, - is_superuser=is_superuser, - ) - return _inner +@pytest.fixture(scope="session") +def gr_redis_config(settings: GRLBaseSettings) -> RedisConfig: + assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) + return RedisConfig( + dsn=settings.gr_redis, + decode_responses=True, + socket_timeout=settings.redis_timeout, + socket_connect_timeout=settings.redis_timeout, + ) -@pytest.fixture -def gr_business_bank_account_factory( - gr_bbam: BusinessBankAccountManager, -) -> Callable[..., BusinessBankAccount]: - - def _inner( - business_id: PositiveInt, - uuid: UUIDStr | None = None, - transfer_method: TransferMethod | None = None, - account_number: str | None = None, - routing_number: str | None = None, - iban: str | None = None, - swift: str | None = None, - ): - from generalresearch.models.gr.business import TransferMethod - - return gr_bbam.create( - business_id=business_id, - uuid=uuid or uuid4().hex, - transfer_method=transfer_method or TransferMethod.ACH, - account_number=account_number or uuid4().hex[:6], - routing_number=routing_number or uuid4().hex[:6], - iban=iban or uuid4().hex[:6], - swift=swift or uuid4().hex[:6], - ) - - return _inner +@pytest.fixture(scope="session") +def gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: -@pytest.fixture -def gr_business_address_factory( - gr_bam: BusinessAddressManager, -) -> Callable[..., BusinessAddress]: - - def _inner( - business_id: PositiveInt, - uuid: UUIDStr | None = None, - line_1: str | None = None, - line_2: str | None = None, - city: str | None = None, - state: str | None = None, - postal_code: str | None = None, - phone_number: PhoneNumber | None = None, - country: str | None = None, - ): - uuid = uuid or uuid4().hex - line_1 = line_1 or "abc" - line_2 = line_2 or "bczx" - city = city or "Downingtown" - state = state or "CA" - postal_code = postal_code or "94041" - phone_number = None - country = country or "US" - - return gr_bam.create( - business_id=business_id, - uuid=uuid, - line_1=line_1, - line_2=line_2, - city=city, - state=state, - postal_code=postal_code, - phone_number=phone_number, - country=country, - ) - - return _inner + return PostgresConfig( + dsn=django_db_factory("gr_carer"), + connect_timeout=1, + statement_timeout=5, + ) -@pytest.fixture -def gr_business_factory( - gr_bm: BusinessManager, -) -> Callable[..., Business]: +# === Managers === - def _inner( - uuid: UUIDStr | None = None, - name: str | None = None, - team: Team | None = None, - kind: BusinessType | None = None, - tax_number: str | None = None, - ) -> Business: - from random import randint - uuid = uuid or uuid4().hex - name = name or "< Unknown >" - tax_number = tax_number or str(randint(1, 999_999_999)) +@pytest.fixture(scope="session") +def gr_user_manager( + gr_db: PostgresConfig, gr_redis_config: RedisConfig +) -> GRUserManager: + assert gr_db.dsn.path + assert "/unittest-" in gr_db.dsn.path - return gr_bm.create( - uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number - ) + from generalresearch.managers.gr.authentication import GRUserManager - return _inner + return GRUserManager(pg_config=gr_db, redis_config=gr_redis_config) -@pytest.fixture -def gr_team( - gr_tm: TeamManager, -) -> Callable[..., Team]: +@pytest.fixture(scope="session") +def gr_team_manager(gr_db: PostgresConfig) -> GRTokenManager: + assert gr_db.dsn.path + assert "/unittest-" in gr_db.dsn.path + + from generalresearch.managers.gr.authentication import GRTokenManager + + return GRTokenManager(pg_config=gr_db) + + +@pytest.fixture(scope="session") +def gr_business_manager( + gr_db: PostgresConfig, gr_redis_config: RedisConfig +) -> BusinessManager: + return BusinessManager(pg_config=gr_db, redis_config=gr_redis_config) + - def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team: - uuid = uuid or uuid4().hex - name = name or f"name-{uuid4().hex[:12]}" +@pytest.fixture(scope="session") +def gr_business_bank_account_manager( + gr_db: PostgresConfig, +) -> BusinessBankAccountManager: + return BusinessBankAccountManager(pg_config=gr_db) - return gr_tm.create(uuid=uuid, name=name) - return _inner +@pytest.fixture(scope="session") +def gr_business_address_manager( + gr_db: PostgresConfig, +) -> BusinessAddressManager: + return BusinessAddressManager(pg_config=gr_db) diff --git a/test_utils/managers/grliq/__init__.py b/test_utils/managers/grliq/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/managers/grliq/conftest.py b/test_utils/managers/grliq/conftest.py deleted file mode 100644 index 525d8c8..0000000 --- a/test_utils/managers/grliq/conftest.py +++ /dev/null @@ -1,61 +0,0 @@ -from __future__ import annotations - -from datetime import datetime, timezone -from typing import Callable -from uuid import uuid4 - -import pytest - -from generalresearch.grliq.managers import DUMMY_GRLIQ_DATA -from generalresearch.grliq.managers.forensic_data import GrlIqDataManager -from generalresearch.grliq.models.forensic_data import GrlIqData - - -@pytest.fixture -def grliq_data_factory(grliq_dm: GrlIqDataManager) -> Callable[..., GrlIqData]: - - def _inner( - is_attempt_allowed: bool = True, - product_id: str | None = None, - product_user_id: str | None = None, - uuid: str | None = None, - mid: str | None = None, - created_at: datetime | None = None, - ) -> GrlIqData: - """ - Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), - and GrlIqForensicCategoryResult (category_results) - :param is_attempt_allowed: Whether the attempt is allowed. - :param product_id: product_id of user - :param product_user_id: product_user_id of user - :param uuid: uuid for the grliq data record - :param mid: the thl_session:uuid / mid for the attempt. - :return: - """ - import copy - - res: GrlIqData = copy.deepcopy(DUMMY_GRLIQ_DATA[int(is_attempt_allowed)]) - - product_id = product_id or uuid4().hex - product_user_id = product_user_id or uuid4().hex - uuid = uuid or uuid4().hex - mid = mid or uuid4().hex - created_at = created_at or datetime.now(tz=timezone.utc) - - res["data"].product_id = product_id - res["data"].product_user_id = product_user_id - res["data"].uuid = uuid - res["data"].mid = mid - res["data"].created_at = created_at - res["result_data"].uuid = uuid - res["category_result"].uuid = uuid - - return grliq_dm.create( - iq_data=res["data"], - result_data=res["result_data"], - category_result=res["category_result"], - fraud_score=res["category_result"].fraud_score, - is_attempt_allowed=res["category_result"].is_attempt_allowed(), - ) - - return _inner diff --git a/test_utils/managers/ledger/conftest.py b/test_utils/managers/ledger/conftest.py index 0aa6cb3..ce8348e 100644 --- a/test_utils/managers/ledger/conftest.py +++ b/test_utils/managers/ledger/conftest.py @@ -1,740 +1,94 @@ -from datetime import datetime -from decimal import Decimal -from random import randint -from typing import TYPE_CHECKING, Callable, Dict, Optional -from uuid import uuid4 +from __future__ import annotations import pytest -from generalresearch.currency import USDCent -from generalresearch.managers.base import PostgresManager -from test_utils.models.conftest import ( - payout_config, - product_amt_true, - product_user_wallet_no, - product_user_wallet_yes, - session, - session_factory, - user_factory, - wall, - wall_factory, +from generalresearch.managers.base import Permission +from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerAccountManager, + LedgerManager, + LedgerTransactionManager, ) - -_ = ( - user_factory, - product_user_wallet_no, - wall, - product_amt_true, - product_user_wallet_yes, - session_factory, - session, - wall_factory, - payout_config, +from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, ) +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig -if TYPE_CHECKING: - - from generalresearch.currency import LedgerCurrency - from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager - from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, - ) - from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - BusinessPayoutEventManager, - ) - from generalresearch.managers.thl.session import SessionManager - from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.thl.ledger import ( - LedgerAccount, - LedgerTransaction, - ) - from generalresearch.models.thl.payout import ( - BrokerageProductPayoutEvent, - ) - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.session import Session - from generalresearch.models.thl.user import User - - -@pytest.fixture -def ledger_account( - request, lm: "LedgerManager", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) +# --- Ledger --- - account_type = getattr(request, "account_type", AccountType.CASH) - direction = getattr(request, "direction", Direction.CREDIT) - acct_uuid = uuid4().hex - qn = ":".join([currency, account_type, acct_uuid]) +@pytest.fixture(scope="session") +def ledger_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> LedgerManager: - acct_model = LedgerAccount( - uuid=acct_uuid, - display_name=f"test-{acct_uuid}", - currency=currency, - qualified_name=qn, - account_type=account_type, - normal_balance=direction, + return LedgerManager( + pg_config=thl_web_rw, + permissions=[ + Permission.CREATE, + Permission.READ, + Permission.UPDATE, + Permission.DELETE, + ], + testing=True, + redis_config=thl_redis_config, ) - return lm.create_account(account=acct_model) - - -@pytest.fixture -def ledger_account_factory( - request, thl_lm: "ThlLedgerManager", lm: "LedgerManager", currency: "LedgerCurrency" -) -> Callable[..., "LedgerAccount"]: - - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - def _inner( - product: "Product", - account_type: AccountType = AccountType.CASH, - direction: Direction = Direction.CREDIT, - ) -> "LedgerAccount": - thl_lm.get_account_or_create_bp_wallet(product=product) - acct_uuid = uuid4().hex - qn = ":".join([currency, account_type, acct_uuid]) - - acct_model = LedgerAccount( - uuid=acct_uuid, - display_name=f"test-{acct_uuid}", - currency=currency, - qualified_name=qn, - account_type=account_type, - normal_balance=direction, - ) - return lm.create_account(account=acct_model) - - return _inner - -@pytest.fixture -def ledger_account_credit( - request, lm: "LedgerManager", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import AccountType, Direction - account_type = AccountType.REVENUE - acct_uuid = uuid4().hex +@pytest.fixture(scope="session") +def ledger_tx_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> LedgerTransactionManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path - qn = ":".join([currency, account_type, acct_uuid]) - from generalresearch.models.thl.ledger import LedgerAccount - - acct_model = LedgerAccount( - uuid=acct_uuid, - display_name=f"test-{acct_uuid}", - currency=currency, - qualified_name=qn, - account_type=account_type, - normal_balance=Direction.CREDIT, + from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerTransactionManager, ) - return lm.create_account(account=acct_model) - - -@pytest.fixture -def ledger_account_debit( - request, lm: "LedgerManager", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import AccountType, Direction - - account_type = AccountType.EXPENSE - acct_uuid = uuid4().hex - - qn = ":".join([currency, account_type, acct_uuid]) - from generalresearch.models.thl.ledger import LedgerAccount - acct_model = LedgerAccount( - uuid=acct_uuid, - display_name=f"test-{acct_uuid}", - currency=currency, - qualified_name=qn, - account_type=account_type, - normal_balance=Direction.DEBIT, + return LedgerTransactionManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + testing=True, + redis_config=thl_redis_config, ) - return lm.create_account(account=acct_model) -@pytest.fixture -def tag(request, lm: "LedgerManager") -> str: - from generalresearch.currency import LedgerCurrency +@pytest.fixture(scope="session") +def ledger_account_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> LedgerAccountManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path - return ( - request.param - if hasattr(request, "tag") - else f"{LedgerCurrency.TEST}:{uuid4().hex}" + from generalresearch.managers.thl.ledger_manager.ledger import ( + LedgerAccountManager, ) - -@pytest.fixture -def usd_cent(request) -> USDCent: - amount = randint(99, 9_999) - return request.param if hasattr(request, "usd_cent") else USDCent(amount) - - -@pytest.fixture -def bp_payout_event( - product: "Product", - usd_cent: "USDCent", - business_payout_event_manager: "BusinessPayoutEventManager", - thl_lm: "ThlLedgerManager", -) -> "BrokerageProductPayoutEvent": - - return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, - product=product, - amount=usd_cent, - skip_wallet_balance_check=True, - skip_one_per_day_check=True, + return LedgerAccountManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + testing=True, + redis_config=thl_redis_config, ) -@pytest.fixture -def bp_payout_event_factory( - brokerage_product_payout_event_manager: "BrokerageProductPayoutEventManager", - thl_lm: "ThlLedgerManager", -) -> Callable[..., "BrokerageProductPayoutEvent"]: +# --- THL Ledger --- - from generalresearch.currency import USDCent - from generalresearch.models.thl.product import Product - - def _inner( - product: Product, usd_cent: USDCent, ext_ref_id: Optional[str] = None - ) -> "BrokerageProductPayoutEvent": - - return brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, - product=product, - amount=usd_cent, - ext_ref_id=ext_ref_id, - skip_wallet_balance_check=True, - skip_one_per_day_check=True, - ) - - return _inner - - -@pytest.fixture -def currency(lm: "LedgerManager") -> "LedgerCurrency": - # return request.param if hasattr(request, "currency") else LedgerCurrency.TEST - assert lm.currency, "LedgerManager must have a currency specified for these tests" - return lm.currency - - -@pytest.fixture -def tx_metadata(request) -> Optional[Dict[str, str]]: - return ( - request.param - if hasattr(request, "tx_metadata") - else {f"key-{uuid4().hex[:10]}": uuid4().hex} - ) - - -@pytest.fixture -def ledger_tx( - request, - ledger_account_credit: "LedgerAccount", - ledger_account_debit: "LedgerAccount", - tag: str, - currency: "LedgerCurrency", - tx_metadata: Optional[Dict[str, str]], - lm: "LedgerManager", -) -> "LedgerTransaction": - from generalresearch.models.thl.ledger import Direction, LedgerEntry - - amount = int(Decimal("1.00") * 100) - - entries = [ - LedgerEntry( - direction=Direction.CREDIT, - account_uuid=ledger_account_credit.uuid, - amount=amount, - ), - LedgerEntry( - direction=Direction.DEBIT, - account_uuid=ledger_account_debit.uuid, - amount=amount, - ), - ] - - return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata) - - -@pytest.fixture -def create_main_accounts( - lm: "LedgerManager", currency: "LedgerCurrency" -) -> Callable[..., None]: - - def _inner() -> None: - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - account = LedgerAccount( - display_name="Cash flow task complete", - qualified_name=f"{currency.value}:revenue:task_complete", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - currency=lm.currency, - ) - lm.get_account_or_create(account=account) - - account = LedgerAccount( - display_name="Operating Cash Account", - qualified_name=f"{currency.value}:cash", - normal_balance=Direction.DEBIT, - account_type=AccountType.CASH, - currency=currency, - ) - - lm.get_account_or_create(account=account) - - return None - - return _inner - - -@pytest.fixture -def delete_ledger_db(thl_web_rw: "PostgresManager") -> Callable[..., None]: - - def _inner(): - for table in [ - "ledger_transactionmetadata", - "ledger_entry", - "ledger_transaction", - "ledger_account", - ]: - thl_web_rw.execute_write( - query=f"DELETE FROM {table};", - ) - - return _inner - - -@pytest.fixture -def wipe_main_accounts( - thl_web_rw: "PostgresManager", lm: "LedgerManager", currency: "LedgerCurrency" -) -> Callable[..., None]: - - def _inner() -> None: - db_table = thl_web_rw.db_name - qual_names = [ - f"{currency.value}:revenue:task_complete", - f"{currency.value}:cash", - ] - - res = thl_web_rw.execute_sql_query( - query=f""" - SELECT lt.id as ltid, le.id as leid, tmd.id as tmdid, la.uuid as lauuid - FROM `{db_table}`.`ledger_transaction` AS lt - LEFT JOIN `{db_table}`.ledger_entry le - ON lt.id = le.transaction_id - LEFT JOIN `{db_table}`.ledger_account la - ON la.uuid = le.account_id - LEFT JOIN `{db_table}`.ledger_transactionmetadata tmd - ON lt.id = tmd.transaction_id - WHERE la.qualified_name IN %s - """, - params=[qual_names], - ) - - lt = {x["ltid"] for x in res if x["ltid"]} - le = {x["leid"] for x in res if x["leid"]} - tmd = {x["tmdid"] for x in res if x["tmdid"]} - la = {x["lauuid"] for x in res if x["lauuid"]} - - thl_web_rw.execute_sql_query( - query=f""" - DELETE FROM `{db_table}`.`ledger_transactionmetadata` - WHERE id IN %s - """, - params=[tmd], - commit=True, - ) - - thl_web_rw.execute_sql_query( - query=f""" - DELETE FROM `{db_table}`.`ledger_entry` - WHERE id IN %s - """, - params=[le], - commit=True, - ) - - thl_web_rw.execute_sql_query( - query=f""" - DELETE FROM `{db_table}`.`ledger_transaction` - WHERE id IN %s - """, - params=[lt], - commit=True, - ) - - thl_web_rw.execute_sql_query( - query=f""" - DELETE FROM `{db_table}`.`ledger_account` - WHERE uuid IN %s - """, - params=[la], - commit=True, - ) - - return None - - return _inner - - -@pytest.fixture -def account_cash(lm: "LedgerManager", currency: "LedgerCurrency") -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - account = LedgerAccount( - display_name="Operating Cash Account", - qualified_name=f"{currency.value}:cash", - normal_balance=Direction.DEBIT, - account_type=AccountType.CASH, - currency=currency, - ) - return lm.get_account_or_create(account=account) +@pytest.fixture(scope="session") +def thl_ledger_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> ThlLedgerManager: -@pytest.fixture -def account_revenue_task_complete( - lm: "LedgerManager", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, + return ThlLedgerManager( + pg_config=thl_web_rw, + permissions=[ + Permission.CREATE, + Permission.READ, + Permission.UPDATE, + Permission.DELETE, + ], + testing=True, + redis_config=thl_redis_config, ) - - account = LedgerAccount( - display_name="Cash flow task complete", - qualified_name=f"{currency.value}:revenue:task_complete", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - currency=currency, - ) - return lm.get_account_or_create(account=account) - - -@pytest.fixture -def account_expense_tango( - lm: "LedgerManager", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - account = LedgerAccount( - display_name="Tango Fee", - qualified_name=f"{currency.value}:expense:tango_fee", - normal_balance=Direction.DEBIT, - account_type=AccountType.EXPENSE, - currency=currency, - ) - return lm.get_account_or_create(account=account) - - -@pytest.fixture -def user_account_user_wallet( - lm: "LedgerManager", user: "User", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - account = LedgerAccount( - display_name=f"{user.uuid} Wallet", - qualified_name=f"{currency.value}:user_wallet:{user.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.USER_WALLET, - reference_type="user", - reference_uuid=user.uuid, - currency=currency, - ) - return lm.get_account_or_create(account=account) - - -@pytest.fixture -def product_account_bp_wallet( - lm: "LedgerManager", product: "Product", currency: "LedgerCurrency" -) -> "LedgerAccount": - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - account = LedgerAccount.model_validate( - dict( - display_name=f"{product.name} Wallet", - qualified_name=f"{currency.value}:bp_wallet:{product.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.BP_WALLET, - reference_type="bp", - reference_uuid=product.uuid, - currency=currency, - ) - ) - return lm.get_account_or_create(account=account) - - -@pytest.fixture -def setup_accounts( - product_factory: Callable[..., "Product"], - lm: "LedgerManager", - user: "User", - currency: "LedgerCurrency", -) -> None: - from generalresearch.models.thl.ledger import ( - AccountType, - Direction, - LedgerAccount, - ) - - # BP's wallet and a revenue from their commissions account. - p1 = product_factory() - - account = LedgerAccount( - display_name=f"Revenue from {p1.name} commission", - qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - reference_type="bp", - reference_uuid=p1.uuid, - currency=currency, - ) - lm.get_account_or_create(account=account) - - account = LedgerAccount.model_validate( - dict( - display_name=f"{p1.name} Wallet", - qualified_name=f"{currency.value}:bp_wallet:{p1.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.BP_WALLET, - reference_type="bp", - reference_uuid=p1.uuid, - currency=currency, - ) - ) - lm.get_account_or_create(account=account) - - # BP's wallet, user's wallet, and a revenue from their commissions account. - p2 = product_factory() - account = LedgerAccount( - display_name=f"Revenue from {p2.name} commission", - qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.REVENUE, - reference_type="bp", - reference_uuid=p2.uuid, - currency=currency, - ) - lm.get_account_or_create(account) - - account = LedgerAccount( - display_name=f"{p2.name} Wallet", - qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.BP_WALLET, - reference_type="bp", - reference_uuid=p2.uuid, - currency=currency, - ) - lm.get_account_or_create(account) - - account = LedgerAccount( - display_name=f"{user.uuid} Wallet", - qualified_name=f"{currency.value}:user_wallet:{user.uuid}", - normal_balance=Direction.CREDIT, - account_type=AccountType.USER_WALLET, - reference_type="user", - reference_uuid=user.uuid, - currency="test", - ) - lm.get_account_or_create(account=account) - - -@pytest.fixture -def session_with_tx_factory( - user_factory: Callable[..., "User"], - product: "Product", - session_factory: Callable[..., "Session"], - session_manager: "SessionManager", - wall_manager: "WallManager", - utc_hour_ago: datetime, - thl_lm: "ThlLedgerManager", -) -> Callable[..., "Session"]: - - from generalresearch.models.thl.session import ( - Status, - StatusCode1, - ) - from generalresearch.models.thl.user import User - - def _inner( - user: User, - final_status: Status = Status.COMPLETE, - wall_req_cpi: Decimal = Decimal(".50"), - started: datetime = utc_hour_ago, - ) -> Session: - s: Session = session_factory( - user=user, - wall_count=2, - final_status=final_status, - wall_req_cpi=wall_req_cpi, - started=started, - ) - last_wall = s.wall_events[-1] - - wall_manager.finish( - wall=last_wall, - status=Status.COMPLETE, - status_code_1=StatusCode1.COMPLETE, - finished=last_wall.finished, - ) - - status, status_code_1 = s.determine_session_status() - _, _, bp_pay, user_pay = s.determine_payments() - session_manager.finish_with_status( - session=s, - finished=last_wall.finished, - payout=bp_pay, - user_payout=user_pay, - status=status, - status_code_1=status_code_1, - ) - - thl_lm.create_tx_task_complete( - wall=last_wall, - user=user, - created=last_wall.finished, - force=True, - ) - - thl_lm.create_tx_bp_payment(session=s, created=last_wall.finished, force=True) - - return s - - return _inner - - -@pytest.fixture -def adj_to_fail_with_tx_factory( - session_manager: "SessionManager", - wall_manager: "WallManager", - thl_lm: "ThlLedgerManager", -) -> Callable[..., None]: - from datetime import datetime, timedelta - - from generalresearch.models.thl.definitions import WallAdjustedStatus - from generalresearch.models.thl.session import ( - Session, - ) - - def _inner( - session: Session, - created: datetime, - ) -> None: - w1 = wall_manager.get_wall_events(session_id=session.id)[-1] - - # This is defined in `thl-grpc/thl/user_quality_history/recons.py:150` - # so we can't use it as part of this test anyway to add rows to the - # thl_taskadjustment table anyway.. until we created a - # TaskAdjustment Manager to put into generalresearch! - - # create_task_adjustment_event( - # wall, - # user, - # adjusted_status, - # amount_usd=amount_usd, - # alert_time=alert_time, - # ext_status_code=ext_status_code, - # ) - - wall_manager.adjust_status( - wall=w1, - adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, - adjusted_cpi=Decimal("0.00"), - adjusted_timestamp=created, - ) - - thl_lm.create_tx_task_adjustment( - wall=w1, - user=session.user, - created=created + timedelta(milliseconds=1), - ) - - session.wall_events = wall_manager.get_wall_events(session_id=session.id) - session_manager.adjust_status(session=session) - - thl_lm.create_tx_bp_adjustment( - session=session, created=created + timedelta(milliseconds=2) - ) - - return None - - return _inner - - -@pytest.fixture -def adj_to_complete_with_tx_factory( - session_manager: "SessionManager", - wall_manager: "WallManager", - thl_lm: "ThlLedgerManager", -) -> Callable[..., None]: - from datetime import timedelta - - from generalresearch.models.thl.definitions import WallAdjustedStatus - from generalresearch.models.thl.session import ( - Session, - ) - - def _inner( - session: Session, - created: datetime, - ) -> None: - w1 = wall_manager.get_wall_events(session_id=session.id)[-1] - - wall_manager.adjust_status( - wall=w1, - adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE, - adjusted_cpi=w1.req_cpi, - adjusted_timestamp=created, - ) - - thl_lm.create_tx_task_adjustment( - wall=w1, - user=session.user, - created=created + timedelta(milliseconds=1), - ) - - session.wall_events = wall_manager.get_wall_events(session_id=session.id) - session_manager.adjust_status(session=session) - - thl_lm.create_tx_bp_adjustment( - session=session, created=created + timedelta(milliseconds=2) - ) - - return None - - return _inner diff --git a/test_utils/managers/network/conftest.py b/test_utils/managers/network/conftest.py index 6c5ea23..e69de29 100644 --- a/test_utils/managers/network/conftest.py +++ b/test_utils/managers/network/conftest.py @@ -1,143 +0,0 @@ -import os -from datetime import datetime, timedelta, timezone -from uuid import uuid4 - -import pytest - -from generalresearch.managers.network.label import IPLabelManager -from generalresearch.managers.network.tool_run import ToolRunManager -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.parser import parse_mtr_output -from generalresearch.models.network.nmap.parser import parse_nmap_xml -from generalresearch.models.network.rdns.parser import parse_rdns_output -from generalresearch.models.network.tool_run import MTRRun, NmapRun, RDNSRun, Status -from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - MTRRunCommandOptions, - NmapRunCommand, - NmapRunCommandOptions, - RDNSRunCommand, - RDNSRunCommandOptions, -) - - -@pytest.fixture(scope="session") -def scan_group_id(): - return uuid4().hex - - -@pytest.fixture(scope="session") -def iplabel_manager(thl_web_rw) -> IPLabelManager: - assert "/unittest-" in thl_web_rw.dsn.path - - return IPLabelManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def toolrun_manager(thl_web_rw) -> ToolRunManager: - assert "/unittest-" in thl_web_rw.dsn.path - - return ToolRunManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def nmap_raw_output(request) -> str: - fp = os.path.join(request.config.rootpath, "data/nmaprun1.xml") - with open(fp) as f: - data = f.read() - return data - - -@pytest.fixture(scope="session") -def nmap_result(nmap_raw_output): - return parse_nmap_xml(nmap_raw_output) - - -@pytest.fixture(scope="session") -def nmap_run(nmap_result, scan_group_id): - r = nmap_result - config = NmapRunCommand( - command="nmap", - options=NmapRunCommandOptions( - ip=r.target_ip, ports="22-1000,11000,1100,3389,61232", top_ports=None - ), - ) - return NmapRun( - tool_version=r.version, - status=Status.SUCCESS, - ip=r.target_ip, - started_at=r.started_at, - finished_at=r.finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - ) - - -@pytest.fixture(scope="session") -def dig_raw_output(): - return "156.32.33.45.in-addr.arpa. 300 IN PTR scanme.nmap.org." - - -@pytest.fixture(scope="session") -def rdns_result(dig_raw_output): - return parse_rdns_output(ip="45.33.32.156", raw=dig_raw_output) - - -@pytest.fixture(scope="session") -def rdns_run(rdns_result, scan_group_id): - r = rdns_result - ip = "45.33.32.156" - utc_now = datetime.now(tz=timezone.utc) - config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip)) - return RDNSRun( - tool_version="1.2.3", - status=Status.SUCCESS, - ip=ip, - started_at=utc_now, - finished_at=utc_now + timedelta(seconds=1), - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - ) - - -@pytest.fixture(scope="session") -def mtr_raw_output(request): - fp = os.path.join(request.config.rootpath, "data/mtr_fatbeam.json") - with open(fp) as f: - data = f.read() - return data - - -@pytest.fixture(scope="session") -def mtr_result(mtr_raw_output): - return parse_mtr_output(mtr_raw_output, port=443, protocol=IPProtocol.TCP) - - -@pytest.fixture(scope="session") -def mtr_run(mtr_result, scan_group_id): - r = mtr_result - utc_now = datetime.now(tz=timezone.utc) - config = MTRRunCommand( - command="mtr", - options=MTRRunCommandOptions( - ip=r.destination, protocol=IPProtocol.TCP, port=443 - ), - ) - - return MTRRun( - tool_version="1.2.3", - status=Status.SUCCESS, - ip=r.destination, - started_at=utc_now, - finished_at=utc_now + timedelta(seconds=1), - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - facility_id=1, - source_ip="1.2.3.4", - ) diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 76e4226..5b70961 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -1,112 +1,258 @@ from __future__ import annotations -from decimal import Decimal -from random import randint from typing import Callable -import faker -from pydantic import PositiveInt - -from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager -from generalresearch.models.custom_types import IPvAnyAddressStr -from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType - -fake = faker.Faker() - - -def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: - - def _inner( - geoname_id: PositiveInt | None = None, - continent_code: str | None = None, - continent_name: str | None = None, - country_iso: str | None = None, - country_name: str | None = None, - subdivision_1_iso: str | None = None, - subdivision_1_name: str | None = None, - subdivision_2_iso: str | None = None, - subdivision_2_name: str | None = None, - city_name: str | None = None, - metro_code: int | None = None, - time_zone: str | None = None, - is_in_european_union: bool | None = None, - ) -> IPGeoname: - - return ipgeoname_manager.create( - geoname_id=geoname_id or randint(1, 999_999_999), - continent_code=continent_code or "na", - continent_name=continent_name or "North America", - country_iso=country_iso or "us", - country_name=country_name or "United States", - subdivision_1_iso=subdivision_1_iso or "fl", - subdivision_1_name=subdivision_1_name or "Florida", - subdivision_2_iso=subdivision_2_iso, - subdivision_2_name=subdivision_2_name, - city_name=city_name, - metro_code=metro_code, - time_zone=time_zone, - is_in_european_union=is_in_european_union, - ) - - return _inner - - -def ipinformation_factory( - ipinformation_manager: IPInformationManager, -) -> Callable[..., IPInformation]: - - def _inner( - ip: IPvAnyAddressStr | None = None, - geoname_id: PositiveInt | None = None, - country_iso: str | None = None, - registered_country_iso: str | None = None, - is_anonymous: bool | None = None, - is_anonymous_vpn: bool | None = None, - is_hosting_provider: bool | None = None, - is_public_proxy: bool | None = None, - is_tor_exit_node: bool | None = None, - is_residential_proxy: bool | None = None, - autonomous_system_number: PositiveInt | None = None, - autonomous_system_organization: str | None = None, - domain: str | None = None, - isp: str | None = None, - mobile_country_code: str | None = None, - mobile_network_code: str | None = None, - network: str | None = None, - organization: str | None = None, - static_ip_score: float | None = None, - user_type: UserType | None = None, - postal_code: str | None = None, - latitude: Decimal | None = None, - longitude: Decimal | None = None, - accuracy_radius: int | None = None, - ) -> IPInformation: - - return ipinformation_manager.create( - ip=ip or fake.ipv4_public(), - geoname_id=geoname_id, - country_iso=country_iso or fake.country_code(), - registered_country_iso=registered_country_iso, - is_anonymous=is_anonymous, - is_anonymous_vpn=is_anonymous_vpn, - is_hosting_provider=is_hosting_provider, - is_public_proxy=is_public_proxy, - is_tor_exit_node=is_tor_exit_node, - is_residential_proxy=is_residential_proxy, - autonomous_system_number=autonomous_system_number, - autonomous_system_organization=autonomous_system_organization, - domain=domain, - isp=isp, - mobile_country_code=mobile_country_code, - mobile_network_code=mobile_network_code, - network=network, - organization=organization, - static_ip_score=static_ip_score, - user_type=user_type, - postal_code=postal_code, - latitude=latitude, - longitude=longitude, - accuracy_radius=accuracy_radius, - ) - - return _inner +import pytest +from pydantic import PostgresDsn + +from generalresearch.config import GRLBaseSettings +from generalresearch.managers.base import Permission +from generalresearch.managers.thl.buyer import BuyerManager +from generalresearch.managers.thl.category import CategoryManager +from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + PayoutEventManager, + UserPayoutEventManager, +) +from generalresearch.managers.thl.product import ProductManager +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.task_adjustment import ( + TaskAdjustmentManager, +) +from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, +) +from generalresearch.managers.thl.user_manager.user_metadata_manager import ( + UserMetadataManager, +) +from generalresearch.managers.thl.wall import ( + WallCacheManager, + WallManager, +) +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig + + +@pytest.fixture(scope="session") +def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: + + return PostgresConfig( + dsn=django_db_factory("generalresearch.thl_django"), + connect_timeout=1, + statement_timeout=5, + ) + + +@pytest.fixture(scope="session") +def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: + return thl_web_rr + + +@pytest.fixture(scope="session") +def thl_redis_config(settings: GRLBaseSettings) -> RedisConfig: + return RedisConfig( + dsn=settings.thl_redis, + decode_responses=True, + socket_timeout=settings.redis_timeout, + socket_connect_timeout=settings.redis_timeout, + ) + + +@pytest.fixture(scope="session") +def payout_event_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> PayoutEventManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.payout import PayoutEventManager + + return PayoutEventManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + redis_config=thl_redis_config, + ) + + +@pytest.fixture(scope="session") +def user_payout_event_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> UserPayoutEventManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.payout import UserPayoutEventManager + + return UserPayoutEventManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + redis_config=thl_redis_config, + ) + + +@pytest.fixture(scope="session") +def brokerage_product_payout_event_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> BrokerageProductPayoutEventManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + ) + + return BrokerageProductPayoutEventManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + redis_config=thl_redis_config, + ) + + +@pytest.fixture(scope="session") +def business_payout_event_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> BusinessPayoutEventManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.payout import ( + BusinessPayoutEventManager, + ) + + return BusinessPayoutEventManager( + pg_config=thl_web_rw, + permissions=[Permission.CREATE, Permission.READ], + redis_config=thl_redis_config, + ) + + +@pytest.fixture(scope="session") +def product_manager(thl_web_rw: PostgresConfig) -> ProductManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.product import ProductManager + + return ProductManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def user_manager( + settings: GRLBaseSettings, thl_web_rw: PostgresConfig, thl_web_rr: PostgresConfig +) -> UserManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert thl_web_rr.dsn + assert thl_web_rr.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rr.dsn.path + + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + + return UserManager( + pg_config=thl_web_rw, + pg_config_rr=thl_web_rr, + redis=settings.redis, + ) + + +@pytest.fixture(scope="session") +def user_metadata_manager(thl_web_rw: PostgresConfig) -> UserMetadataManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.user_manager.user_metadata_manager import ( + UserMetadataManager, + ) + + return UserMetadataManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def session_manager(thl_web_rw: PostgresConfig) -> SessionManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.session import SessionManager + + return SessionManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def wall_manager(thl_web_rw: PostgresConfig) -> WallManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.wall import WallManager + + return WallManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def wall_cache_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> WallCacheManager: + # assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.wall import WallCacheManager + + return WallCacheManager(pg_config=thl_web_rw, redis_config=thl_redis_config) + + +@pytest.fixture(scope="session") +def task_adjustment_manager(thl_web_rw: PostgresConfig) -> TaskAdjustmentManager: + # assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.task_adjustment import ( + TaskAdjustmentManager, + ) + + return TaskAdjustmentManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def category_manager(thl_web_rw: PostgresConfig) -> CategoryManager: + assert thl_web_rw.dsn + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + from generalresearch.managers.thl.category import CategoryManager + + return CategoryManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def buyer_manager(thl_web_rw: PostgresConfig) -> BuyerManager: + # assert "/unittest-" in thl_web_rw.dsn.path + from generalresearch.managers.thl.buyer import BuyerManager + + return BuyerManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def survey_manager(thl_web_rw: PostgresConfig): + # assert "/unittest-" in thl_web_rw.dsn.path + from generalresearch.managers.thl.survey import SurveyManager + + return SurveyManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def surveystat_manager(thl_web_rw: PostgresConfig): + # assert "/unittest-" in thl_web_rw.dsn.path + from generalresearch.managers.thl.survey import SurveyStatManager + + return SurveyStatManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def surveypenalty_manager(thl_redis_config: RedisConfig): + from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager + + return SurveyPenaltyManager(redis_config=thl_redis_config) diff --git a/test_utils/managers/upk/conftest.py b/test_utils/managers/upk/conftest.py index e28d085..d8f956c 100644 --- a/test_utils/managers/upk/conftest.py +++ b/test_utils/managers/upk/conftest.py @@ -1,173 +1,69 @@ -import os -import time -from typing import TYPE_CHECKING, Optional -from uuid import UUID +from typing import Callable, Generator -import pandas as pd import pytest +from generalresearch.managers.thl.profiling.question import ( + QuestionManager, +) +from generalresearch.managers.thl.profiling.schema import ( + UpkSchemaManager, +) +from generalresearch.managers.thl.profiling.uqa import UQAManager +from generalresearch.managers.thl.profiling.user_upk import ( + UserUpkManager, +) +from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig - -if TYPE_CHECKING: - from generalresearch.managers.thl.category import CategoryManager - - -def insert_data_from_csv( - thl_web_rw: PostgresConfig, - table_name: str, - fp: Optional[str] = None, - disable_fk_checks: bool = False, - df: Optional[pd.DataFrame] = None, -): - assert fp is not None or df is not None and not (fp is not None and df is not None) - if fp: - df = pd.read_csv(fp, dtype=str) - df = df.where(pd.notnull(df), None) - cols = list(df.columns) - col_str = ", ".join(cols) - values_str = ", ".join(["%s"] * len(cols)) - if "id" in df.columns and len(df["id"].iloc[0]) == 36: - df["id"] = df["id"].map(lambda x: UUID(x).hex) - args = df.to_dict("tight")["data"] - - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - if disable_fk_checks: - c.execute("SET CONSTRAINTS ALL DEFERRED") - c.executemany( - f"INSERT INTO {table_name} ({col_str}) VALUES ({values_str})", - params_seq=args, - ) - conn.commit() +from generalresearch.redis_helper import RedisConfig @pytest.fixture(scope="session") -def category_data( - thl_web_rw: PostgresConfig, category_manager: "CategoryManager" -) -> None: - fp = os.path.join(os.path.dirname(__file__), "marketplace_category.csv.gz") - insert_data_from_csv( - thl_web_rw, - fp=fp, - table_name="marketplace_category", - disable_fk_checks=True, - ) - # Don't strictly need to do this, but probably we should - category_manager.populate_caches() - cats = category_manager.categories.values() - path_id = {c.path: c.id for c in cats} - data = [ - {"id": c.id, "parent_id": path_id[c.parent_path]} for c in cats if c.parent_path - ] - query = """ - UPDATE marketplace_category - SET parent_id = %(parent_id)s - WHERE id = %(id)s; - """ - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - c.executemany(query=query, params_seq=data) - conn.commit() +def upk_schema_manager(thl_web_rw: PostgresConfig) -> UpkSchemaManager: + return UpkSchemaManager(pg_config=thl_web_rw) @pytest.fixture(scope="session") -def property_data(thl_web_rw: PostgresConfig) -> None: - fp = os.path.join(os.path.dirname(__file__), "marketplace_property.csv.gz") - insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_property") - +def user_upk_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> UserUpkManager: -@pytest.fixture(scope="session") -def item_data(thl_web_rw: PostgresConfig) -> None: - fp = os.path.join(os.path.dirname(__file__), "marketplace_item.csv.gz") - insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_item") + return UserUpkManager(pg_config=thl_web_rw, redis_config=thl_redis_config) @pytest.fixture(scope="session") -def propertycategoryassociation_data( +def question_manager( thl_web_rw: PostgresConfig, - category_data, - property_data, - category_manager: "CategoryManager", -) -> None: - table_name = "marketplace_propertycategoryassociation" - fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") - # Need to lookup category pk from uuid - category_manager.populate_caches() - df = pd.read_csv(fp, dtype=str) - df["category_id"] = df["category_id"].map( - lambda x: category_manager.categories[x].id - ) - insert_data_from_csv(thl_web_rw, df=df, table_name=table_name) +) -> QuestionManager: + return QuestionManager(pg_config=thl_web_rw) @pytest.fixture(scope="session") -def propertycountry_data(thl_web_rw: PostgresConfig, property_data) -> None: - fp = os.path.join(os.path.dirname(__file__), "marketplace_propertycountry.csv.gz") - insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_propertycountry") +def uqa_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> UQAManager: + return UQAManager(redis_config=thl_redis_config, pg_config=thl_web_rw) -@pytest.fixture(scope="session") -def propertymarketplaceassociation_data( - thl_web_rw: PostgresConfig, property_data -) -> None: - table_name = "marketplace_propertymarketplaceassociation" - fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") - insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) +@pytest.fixture(scope="function") +def uqa_manager_clear_cache_factory( + uqa_manager: UQAManager, +) -> Callable[..., Generator[None]]: -@pytest.fixture(scope="session") -def propertyitemrange_data( - thl_web_rw: PostgresConfig, property_data, item_data -) -> None: - table_name = "marketplace_propertyitemrange" - fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") - insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) + def _inner(user: User) -> Generator[None]: + # On successive py-test/jenkins runs, the cache may contain + # the previous run's info (keyed under the same user_id) + uqa_manager.clear_cache(user) + yield -@pytest.fixture(scope="session") -def question_data(thl_web_rw: PostgresConfig) -> None: - table_name = "marketplace_question" - fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") - insert_data_from_csv( - thl_web_rw, fp=fp, table_name=table_name, disable_fk_checks=True - ) + uqa_manager.clear_cache(user) + return _inner -@pytest.fixture(scope="session") -def clear_upk_tables(thl_web_rw: PostgresConfig): - tables = [ - "marketplace_propertyitemrange", - "marketplace_propertymarketplaceassociation", - "marketplace_propertycategoryassociation", - "marketplace_category", - "marketplace_item", - "marketplace_property", - "marketplace_propertycountry", - "marketplace_question", - ] - table_str = ", ".join(tables) - - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - c.execute(f"TRUNCATE {table_str} RESTART IDENTITY CASCADE;") - conn.commit() - -@pytest.fixture(scope="session") -def upk_data( - clear_upk_tables, - category_data, - property_data, - item_data, - propertycategoryassociation_data, - propertycountry_data, - propertymarketplaceassociation_data, - propertyitemrange_data, - question_data, -) -> None: - # Wait a second to make sure the HarmonizerCache refresh loop pulls these in - time.sleep(2) - - -def test_fixtures(upk_data): - pass +@pytest.fixture(scope="function") +def uqa_manager_clear_cache( + uqa_manager_clear_cache_factory: Callable[..., None], user: User +): + uqa_manager_clear_cache_factory(user=user) diff --git a/test_utils/managers/upk/marketplace_category.csv.gz b/test_utils/managers/upk/marketplace_category.csv.gz deleted file mode 100644 index 0f8ec1c..0000000 Binary files a/test_utils/managers/upk/marketplace_category.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_item.csv.gz b/test_utils/managers/upk/marketplace_item.csv.gz deleted file mode 100644 index c12c5d8..0000000 Binary files a/test_utils/managers/upk/marketplace_item.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_property.csv.gz b/test_utils/managers/upk/marketplace_property.csv.gz deleted file mode 100644 index a781d1d..0000000 Binary files a/test_utils/managers/upk/marketplace_property.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz b/test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz deleted file mode 100644 index 5b4ea19..0000000 Binary files a/test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_propertycountry.csv.gz b/test_utils/managers/upk/marketplace_propertycountry.csv.gz deleted file mode 100644 index 5d2a637..0000000 Binary files a/test_utils/managers/upk/marketplace_propertycountry.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_propertyitemrange.csv.gz b/test_utils/managers/upk/marketplace_propertyitemrange.csv.gz deleted file mode 100644 index 84f4f0e..0000000 Binary files a/test_utils/managers/upk/marketplace_propertyitemrange.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz b/test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz deleted file mode 100644 index 6b9fd1c..0000000 Binary files a/test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz and /dev/null differ diff --git a/test_utils/managers/upk/marketplace_question.csv.gz b/test_utils/managers/upk/marketplace_question.csv.gz deleted file mode 100644 index bcfc3ad..0000000 Binary files a/test_utils/managers/upk/marketplace_question.csv.gz and /dev/null differ diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 1133d32..9925a9e 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -16,6 +16,7 @@ from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, Status, ) +from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig @@ -59,7 +60,6 @@ if TYPE_CHECKING: Product, ) from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.models.thl.user import User from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel @@ -70,10 +70,10 @@ if TYPE_CHECKING: @pytest.fixture def user( request, - product_manager: "ProductManager", - user_manager: "UserManager", + product_manager: ProductManager, + user_manager: UserManager, thl_web_rr: PostgresConfig, -) -> "User": +) -> User: product = getattr(request, "product", None) if product is None: @@ -172,7 +172,6 @@ def wall(session: Session, user: User, wall_manager: WallManager) -> Wall | None @pytest.fixture def session_factory( - wall_factory: Callable[..., Wall], session_manager: SessionManager, wall_manager: WallManager, utc_hour_ago: datetime, @@ -190,7 +189,7 @@ def session_factory( # Session details final_status: Status = Status.COMPLETE, started: datetime = utc_hour_ago, - ) -> "Session": + ) -> Session: if wall_req_cpis: assert len(wall_req_cpis) == wall_count if wall_statuses: @@ -422,7 +421,7 @@ def business(request, business_manager: BusinessManager) -> Business: @pytest.fixture def business_address( - request, business: "Business", business_address_manager: BusinessAddressManager + request, business: Business, business_address_manager: BusinessAddressManager ) -> BusinessAddress: return business_address_manager.create_dummy(business_id=business.id) @@ -441,77 +440,6 @@ def team(request, team_manager: TeamManager) -> Team: return team_manager.create_dummy() -@pytest.fixture -def gr_user(gr_um: GRUserManager) -> GRUser: - return gr_um.create_dummy() - - -@pytest.fixture -def gr_user_cache( - gr_user: GRUser, - gr_db: PostgresConfig, - thl_web_rr: PostgresConfig, - gr_redis_config: RedisConfig, -) -> GRUser: - gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config - ) - return gr_user - - -@pytest.fixture -def gr_user_factory(gr_um: GRUserManager) -> Callable[..., GRUser]: - - def _inner(): - return gr_um.create_dummy() - - return _inner - - -@pytest.fixture() -def gr_user_token( - gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig -) -> GRToken: - gr_tm.create(user_id=gr_user.id) - gr_user.prefetch_token(pg_config=gr_db) - - res = gr_user.token - assert res is not None, "GRToken should exist after creation and prefetching" - return res - - -@pytest.fixture() -def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: - return gr_user_token.auth_header - - -@pytest.fixture(scope="function") -def membership( - request, team: Team, gr_user: GRUser, team_manager: TeamManager -) -> Membership: - assert team.id, "Team must be saved" - assert gr_user.id, "GRUser must be saved" - return team_manager.add_user(team=team, gr_user=gr_user) - - -@pytest.fixture(scope="function") -def membership_factory( - team: Team, - gr_user: GRUser, - membership_manager: MembershipManager, - team_manager: TeamManager, - gr_um: GRUserManager, -) -> Callable[..., Membership]: - - def _inner(**kwargs) -> Membership: - _team = kwargs.get("team", team_manager.create_dummy()) - _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) - - return membership_manager.create(team=_team, gr_user=_gr_user) - - return _inner - - @pytest.fixture def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog: @@ -613,7 +541,7 @@ def buyer_factory(buyer_manager: BuyerManager) -> Callable[..., Buyer]: @pytest.fixture(scope="session") -def survey(survey_manager: SurveyManager, buyer: Buyer) -> "Survey": +def survey(survey_manager: SurveyManager, buyer: Buyer) -> Survey: s = Survey(source=Source.TESTING, survey_id=uuid4().hex, buyer_code=buyer.code) survey_manager.create_bulk([s]) return s diff --git a/test_utils/models/contest/__init__.py b/test_utils/models/contest/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py new file mode 100644 index 0000000..e750076 --- /dev/null +++ b/test_utils/models/contest/conftest.py @@ -0,0 +1,292 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from decimal import Decimal +from typing import Callable +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.currency import USDCent +from generalresearch.managers.thl.contest_manager import ContestManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.models.thl.contest.contest import Contest +from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, +) +from generalresearch.models.thl.contest.milestone import ( + MilestoneContestCreate, +) +from generalresearch.models.thl.contest.raffle import ( + RaffleContestCreate, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User + +# === Miscellaneous === + +# === Managers === + +# === Models === + + +@pytest.fixture +def raffle_contest_create() -> RaffleContestCreate: + from generalresearch.models.thl.contest import ( + ContestEndCondition, + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.raffle import ( + ContestEntryType, + RaffleContestCreate, + ) + + # This is what we'll get from the fastapi endpoint + return RaffleContestCreate( + name="test", + contest_type=ContestType.RAFFLE, + entry_type=ContestEntryType.CASH, + prizes=[ + ContestPrize( + name="iPod 64GB White", + kind=ContestPrizeKind.PHYSICAL, + estimated_cash_value=USDCent(100), + ) + ], + end_condition=ContestEndCondition(target_entry_amount=USDCent(100)), + ) + + +@pytest.fixture +def raffle_contest_in_db( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create + ) + + +@pytest.fixture +def raffle_contest( + product_user_wallet_yes: Product, raffle_contest_create: RaffleContestCreate +) -> Contest: + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create + ) + + +@pytest.fixture(scope="function") +def raffle_contest_factory( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + raffle_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=raffle_contest_create, + ) + + return _inner + + +@pytest.fixture +def milestone_contest_create() -> MilestoneContestCreate: + from generalresearch.models.thl.contest import ( + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.milestone import ( + ContestEntryTrigger, + MilestoneContestCreate, + MilestoneContestEndCondition, + ) + + # This is what we'll get from the fastapi endpoint + return MilestoneContestCreate( + name="Win a 50% bonus for 7 days and a $1 bonus after your first 3 completes!", + description="only valid for the first 5 users", + contest_type=ContestType.MILESTONE, + prizes=[ + ContestPrize( + name="50% for 7 days", + kind=ContestPrizeKind.PROMOTION, + estimated_cash_value=USDCent(0), + ), + ContestPrize( + name="$1 Bonus", + kind=ContestPrizeKind.CASH, + cash_amount=USDCent(1_00), + estimated_cash_value=USDCent(1_00), + ), + ], + end_condition=MilestoneContestEndCondition( + ends_at=datetime(year=2030, month=1, day=1, tzinfo=timezone.utc), + max_winners=5, + ), + entry_trigger=ContestEntryTrigger.TASK_COMPLETE, + target_amount=3, + ) + + +@pytest.fixture +def milestone_contest_in_db( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create + ) + + +@pytest.fixture +def milestone_contest( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, +) -> Contest: + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create + ) + + +@pytest.fixture(scope="function") +def milestone_contest_factory( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + milestone_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=milestone_contest_create, + ) + + return _inner + + +@pytest.fixture +def leaderboard_contest_create( + product_user_wallet_yes: Product, +) -> LeaderboardContestCreate: + from generalresearch.models.thl.contest import ( + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, + ) + + # This is what we'll get from the fastapi endpoint + return LeaderboardContestCreate( + name="test", + contest_type=ContestType.LEADERBOARD, + prizes=[ + ContestPrize( + name="$15 Cash", + estimated_cash_value=USDCent(15_00), + cash_amount=USDCent(15_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=1, + ), + ContestPrize( + name="$10 Cash", + estimated_cash_value=USDCent(10_00), + cash_amount=USDCent(10_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=2, + ), + ], + leaderboard_key=f"leaderboard:{product_user_wallet_yes.uuid}:us:daily:2025-01-01:complete_count", + ) + + +@pytest.fixture +def leaderboard_contest_in_db( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + +@pytest.fixture +def leaderboard_contest( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, +): + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + +@pytest.fixture(scope="function") +def leaderboard_contest_factory( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + leaderboard_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + return _inner + + +@pytest.fixture +def user_with_money( + request: Request, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + thl_lm: ThlLedgerManager, +) -> User: + + params = getattr(request, "param", {}) or {} + min_balance = int(params.get("min_balance", USDCent(1_00))) + + user: User = user_factory(product=product_user_wallet_yes) + wallet = thl_lm.get_account_or_create_user_wallet(user) + balance = thl_lm.get_account_balance(wallet) + todo = min_balance - balance + if todo > 0: + # # Put money in user's wallet + thl_lm.create_tx_user_bonus( + user=user, + ref_uuid=uuid4().hex, + description="bonus", + amount=Decimal(todo) / 100, + ) + print(f"wallet balance: {thl_lm.get_user_wallet_balance(user=user)}") + + return user diff --git a/test_utils/models/gr/__init__.py b/test_utils/models/gr/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py new file mode 100644 index 0000000..df97306 --- /dev/null +++ b/test_utils/models/gr/conftest.py @@ -0,0 +1,213 @@ +from __future__ import annotations + +from typing import Callable +from uuid import uuid4 + +import pytest +from pydantic import PositiveInt +from pydantic_extra_types.phone_numbers import PhoneNumber + +from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.gr.authentication import GRToken, GRUser +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + BusinessType, + TransferMethod, +) +from generalresearch.models.gr.team import Membership, Team +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig + +# --- Static --- + + +# --- Factory / Database --- + + +@pytest.fixture +def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]: + + def _inner( + sub: str | None = None, + is_superuser: bool = False, + ) -> GRUser: + sub = sub or f"{uuid4().hex}-{uuid4().hex}" + + return gr_user_manager.create( + sub=sub, + is_superuser=is_superuser, + ) + + return _inner + + +@pytest.fixture +def gr_user_cache( + gr_user: GRUser, + gr_db: PostgresConfig, + thl_web_rr: PostgresConfig, + gr_redis_config: RedisConfig, +) -> GRUser: + gr_user.set_cache( + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + ) + return gr_user + + +@pytest.fixture +def gr_business_bank_account_factory( + gr_bbam: BusinessBankAccountManager, +) -> Callable[..., BusinessBankAccount]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + transfer_method: TransferMethod | None = None, + account_number: str | None = None, + routing_number: str | None = None, + iban: str | None = None, + swift: str | None = None, + ): + from generalresearch.models.gr.business import TransferMethod + + return gr_bbam.create( + business_id=business_id, + uuid=uuid or uuid4().hex, + transfer_method=transfer_method or TransferMethod.ACH, + account_number=account_number or uuid4().hex[:6], + routing_number=routing_number or uuid4().hex[:6], + iban=iban or uuid4().hex[:6], + swift=swift or uuid4().hex[:6], + ) + + return _inner + + +@pytest.fixture +def gr_business_address_factory( + gr_bam: BusinessAddressManager, +) -> Callable[..., BusinessAddress]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + line_1: str | None = None, + line_2: str | None = None, + city: str | None = None, + state: str | None = None, + postal_code: str | None = None, + phone_number: PhoneNumber | None = None, + country: str | None = None, + ): + uuid = uuid or uuid4().hex + line_1 = line_1 or "abc" + line_2 = line_2 or "bczx" + city = city or "Downingtown" + state = state or "CA" + postal_code = postal_code or "94041" + phone_number = None + country = country or "US" + + return gr_bam.create( + business_id=business_id, + uuid=uuid, + line_1=line_1, + line_2=line_2, + city=city, + state=state, + postal_code=postal_code, + phone_number=phone_number, + country=country, + ) + + return _inner + + +@pytest.fixture +def gr_business_factory( + gr_bm: BusinessManager, +) -> Callable[..., Business]: + + def _inner( + uuid: UUIDStr | None = None, + name: str | None = None, + team: Team | None = None, + kind: BusinessType | None = None, + tax_number: str | None = None, + ) -> Business: + from random import randint + + uuid = uuid or uuid4().hex + name = name or "< Unknown >" + tax_number = tax_number or str(randint(1, 999_999_999)) + + return gr_bm.create( + uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number + ) + + return _inner + + +@pytest.fixture +def gr_team( + gr_tm: TeamManager, +) -> Callable[..., Team]: + + def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team: + uuid = uuid or uuid4().hex + name = name or f"name-{uuid4().hex[:12]}" + + return gr_tm.create(uuid=uuid, name=name) + + return _inner + + +@pytest.fixture() +def gr_user_token( + gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig +) -> GRToken: + gr_tm.create(user_id=gr_user.id) + gr_user.prefetch_token(pg_config=gr_db) + + res = gr_user.token + assert res is not None, "GRToken should exist after creation and prefetching" + return res + + +@pytest.fixture() +def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: + return gr_user_token.auth_header + + +@pytest.fixture(scope="function") +def membership(team: Team, gr_user: GRUser, team_manager: TeamManager) -> Membership: + assert team.id, "Team must be saved" + assert gr_user.id, "GRUser must be saved" + return team_manager.add_user(team=team, gr_user=gr_user) + + +@pytest.fixture(scope="function") +def membership_factory( + team: Team, + gr_user: GRUser, + membership_manager: MembershipManager, + team_manager: TeamManager, + gr_um: GRUserManager, +) -> Callable[..., Membership]: + + def _inner(**kwargs) -> Membership: + _team = kwargs.get("team", team_manager.create_dummy()) + _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) + + return membership_manager.create(team=_team, gr_user=_gr_user) + + return _inner diff --git a/test_utils/models/ledger/__init__.py b/test_utils/models/ledger/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py new file mode 100644 index 0000000..5bef113 --- /dev/null +++ b/test_utils/models/ledger/conftest.py @@ -0,0 +1,724 @@ +from __future__ import annotations + +from datetime import datetime +from decimal import Decimal +from random import randint +from typing import TYPE_CHECKING, Callable +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.currency import USDCent +from generalresearch.managers.base import PostgresManager +from test_utils.models.conftest import ( + payout_config, + product_amt_true, + product_user_wallet_no, + product_user_wallet_yes, + session, + session_factory, + user_factory, + wall, + wall_factory, +) + +_ = ( + user_factory, + product_user_wallet_no, + wall, + product_amt_true, + product_user_wallet_yes, + session_factory, + session, + wall_factory, + payout_config, +) + +if TYPE_CHECKING: + + from generalresearch.currency import LedgerCurrency + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + ) + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.thl.ledger import ( + LedgerAccount, + LedgerTransaction, + ) + from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User + + +@pytest.fixture +def ledger_account( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account_type = getattr(request, "account_type", AccountType.CASH) + direction = getattr(request, "direction", Direction.CREDIT) + + acct_uuid = uuid4().hex + qn = f"{currency}:{account_type}:{acct_uuid}" + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=direction, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def ledger_account_factory( + request: Request, + thl_lm: ThlLedgerManager, + lm: LedgerManager, + currency: LedgerCurrency, +) -> Callable[..., LedgerAccount]: + + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + def _inner( + product: Product, + account_type: AccountType = AccountType.CASH, + direction: Direction = Direction.CREDIT, + ) -> LedgerAccount: + thl_lm.get_account_or_create_bp_wallet(product=product) + acct_uuid = uuid4().hex + qn = f"{currency}:{account_type}:{acct_uuid}" + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=direction, + ) + return lm.create_account(account=acct_model) + + return _inner + + +@pytest.fixture +def ledger_account_credit( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import AccountType, Direction + + account_type = AccountType.REVENUE + acct_uuid = uuid4().hex + + qn = f"{currency}:{account_type}:{acct_uuid}" + from generalresearch.models.thl.ledger import LedgerAccount + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=Direction.CREDIT, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def ledger_account_debit( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import AccountType, Direction + + account_type = AccountType.EXPENSE + acct_uuid = uuid4().hex + + qn = f"{currency}:{account_type}:{acct_uuid}" + from generalresearch.models.thl.ledger import LedgerAccount + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=Direction.DEBIT, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def tag(request: Request, lm: LedgerManager) -> str: + from generalresearch.currency import LedgerCurrency + + return ( + request.param + if hasattr(request, "tag") + else f"{LedgerCurrency.TEST}:{uuid4().hex}" + ) + + +@pytest.fixture +def usd_cent(request: Request) -> USDCent: + amount = randint(99, 9_999) + return request.param if hasattr(request, "usd_cent") else USDCent(amount) + + +@pytest.fixture +def bp_payout_event( + product: Product, + usd_cent: USDCent, + business_payout_event_manager: BusinessPayoutEventManager, + thl_lm: ThlLedgerManager, +) -> BrokerageProductPayoutEvent: + + return business_payout_event_manager.create_bp_payout_event( + thl_ledger_manager=thl_lm, + product=product, + amount=usd_cent, + skip_wallet_balance_check=True, + skip_one_per_day_check=True, + ) + + +@pytest.fixture +def bp_payout_event_factory( + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., BrokerageProductPayoutEvent]: + + def _inner( + product: Product, usd_cent: USDCent, ext_ref_id: str | None = None + ) -> BrokerageProductPayoutEvent: + + return brokerage_product_payout_event_manager.create_bp_payout_event( + thl_ledger_manager=thl_lm, + product=product, + amount=usd_cent, + ext_ref_id=ext_ref_id, + skip_wallet_balance_check=True, + skip_one_per_day_check=True, + ) + + return _inner + + +@pytest.fixture +def currency(lm: LedgerManager) -> LedgerCurrency: + # return request.param if hasattr(request, "currency") else LedgerCurrency.TEST + assert lm.currency, "LedgerManager must have a currency specified for these tests" + return lm.currency + + +@pytest.fixture +def tx_metadata(request: Request) -> dict[str, str] | None: + return ( + request.param + if hasattr(request, "tx_metadata") + else {f"key-{uuid4().hex[:10]}": uuid4().hex} + ) + + +@pytest.fixture +def ledger_tx( + request: Request, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + tag: str, + currency: LedgerCurrency, + tx_metadata: dict[str, str] | None, + lm: LedgerManager, +) -> LedgerTransaction: + from generalresearch.models.thl.ledger import Direction, LedgerEntry + + amount = int(Decimal("1.00") * 100) + + entries = [ + LedgerEntry( + direction=Direction.CREDIT, + account_uuid=ledger_account_credit.uuid, + amount=amount, + ), + LedgerEntry( + direction=Direction.DEBIT, + account_uuid=ledger_account_debit.uuid, + amount=amount, + ), + ] + + return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata) + + +@pytest.fixture +def create_main_accounts( + lm: LedgerManager, currency: LedgerCurrency +) -> Callable[..., None]: + + def _inner() -> None: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Cash flow task complete", + qualified_name=f"{currency.value}:revenue:task_complete", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + currency=lm.currency, + ) + lm.get_account_or_create(account=account) + + account = LedgerAccount( + display_name="Operating Cash Account", + qualified_name=f"{currency.value}:cash", + normal_balance=Direction.DEBIT, + account_type=AccountType.CASH, + currency=currency, + ) + + lm.get_account_or_create(account=account) + + return _inner + + +@pytest.fixture +def delete_ledger_db(thl_web_rw: PostgresManager) -> Callable[..., None]: + + def _inner(): + for table in [ + "ledger_transactionmetadata", + "ledger_entry", + "ledger_transaction", + "ledger_account", + ]: + thl_web_rw.execute_write( + query=f"DELETE FROM {table};", + ) + + return _inner + + +@pytest.fixture +def wipe_main_accounts( + thl_web_rw: PostgresManager, lm: LedgerManager, currency: LedgerCurrency +) -> Callable[..., None]: + + def _inner() -> None: + db_table = thl_web_rw.db_name + qual_names = [ + f"{currency.value}:revenue:task_complete", + f"{currency.value}:cash", + ] + + res = thl_web_rw.execute_sql_query( + query=f""" + SELECT lt.id as ltid, le.id as leid, tmd.id as tmdid, la.uuid as lauuid + FROM `{db_table}`.`ledger_transaction` AS lt + LEFT JOIN `{db_table}`.ledger_entry le + ON lt.id = le.transaction_id + LEFT JOIN `{db_table}`.ledger_account la + ON la.uuid = le.account_id + LEFT JOIN `{db_table}`.ledger_transactionmetadata tmd + ON lt.id = tmd.transaction_id + WHERE la.qualified_name IN %s + """, + params=[qual_names], + ) + + lt = {x["ltid"] for x in res if x["ltid"]} + le = {x["leid"] for x in res if x["leid"]} + tmd = {x["tmdid"] for x in res if x["tmdid"]} + la = {x["lauuid"] for x in res if x["lauuid"]} + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_transactionmetadata` + WHERE id IN %s + """, + params=[tmd], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_entry` + WHERE id IN %s + """, + params=[le], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_transaction` + WHERE id IN %s + """, + params=[lt], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_account` + WHERE uuid IN %s + """, + params=[la], + commit=True, + ) + + return _inner + + +@pytest.fixture +def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Operating Cash Account", + qualified_name=f"{currency.value}:cash", + normal_balance=Direction.DEBIT, + account_type=AccountType.CASH, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def account_revenue_task_complete( + lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Cash flow task complete", + qualified_name=f"{currency.value}:revenue:task_complete", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Tango Fee", + qualified_name=f"{currency.value}:expense:tango_fee", + normal_balance=Direction.DEBIT, + account_type=AccountType.EXPENSE, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def user_account_user_wallet( + lm: LedgerManager, user: User, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name=f"{user.uuid} Wallet", + qualified_name=f"{currency.value}:user_wallet:{user.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.USER_WALLET, + reference_type="user", + reference_uuid=user.uuid, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def product_account_bp_wallet( + lm: LedgerManager, product: Product, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount.model_validate( + { + "display_name": f"{product.name} Wallet", + "qualified_name": f"{currency.value}:bp_wallet:{product.uuid}", + "normal_balance": Direction.CREDIT, + "account_type": AccountType.BP_WALLET, + "reference_type": "bp", + "reference_uuid": product.uuid, + "currency": currency, + } + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def setup_accounts( + product_factory: Callable[..., Product], + lm: LedgerManager, + user: User, + currency: LedgerCurrency, +) -> None: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + # BP's wallet and a revenue from their commissions account. + p1 = product_factory() + + account = LedgerAccount( + display_name=f"Revenue from {p1.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p1.uuid, + currency=currency, + ) + lm.get_account_or_create(account=account) + + account = LedgerAccount.model_validate( + { + "display_name": f"{p1.name} Wallet", + "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}", + "normal_balance": Direction.CREDIT, + "account_type": AccountType.BP_WALLET, + "reference_type": "bp", + "reference_uuid": p1.uuid, + "currency": currency, + } + ) + lm.get_account_or_create(account=account) + + # BP's wallet, user's wallet, and a revenue from their commissions account. + p2 = product_factory() + account = LedgerAccount( + display_name=f"Revenue from {p2.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) + + account = LedgerAccount( + display_name=f"{p2.name} Wallet", + qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.BP_WALLET, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) + + account = LedgerAccount( + display_name=f"{user.uuid} Wallet", + qualified_name=f"{currency.value}:user_wallet:{user.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.USER_WALLET, + reference_type="user", + reference_uuid=user.uuid, + currency="test", + ) + lm.get_account_or_create(account=account) + + +@pytest.fixture +def session_with_tx_factory( + session_factory: Callable[..., Session], + session_manager: SessionManager, + wall_manager: WallManager, + utc_hour_ago: datetime, + thl_lm: ThlLedgerManager, +) -> Callable[..., Session]: + + from generalresearch.models.thl.session import ( + Status, + StatusCode1, + ) + + def _inner( + user: User, + final_status: Status = Status.COMPLETE, + wall_req_cpi: Decimal = Decimal(".50"), + started: datetime = utc_hour_ago, + ) -> Session: + s: Session = session_factory( + user=user, + wall_count=2, + final_status=final_status, + wall_req_cpi=wall_req_cpi, + started=started, + ) + last_wall = s.wall_events[-1] + + wall_manager.finish( + wall=last_wall, + status=Status.COMPLETE, + status_code_1=StatusCode1.COMPLETE, + finished=last_wall.finished, + ) + + status, status_code_1 = s.determine_session_status() + _, _, bp_pay, user_pay = s.determine_payments() + session_manager.finish_with_status( + session=s, + finished=last_wall.finished, + payout=bp_pay, + user_payout=user_pay, + status=status, + status_code_1=status_code_1, + ) + + thl_lm.create_tx_task_complete( + wall=last_wall, + user=user, + created=last_wall.finished, + force=True, + ) + + thl_lm.create_tx_bp_payment(session=s, created=last_wall.finished, force=True) + + return s + + return _inner + + +@pytest.fixture +def adj_to_fail_with_tx_factory( + session_manager: SessionManager, + wall_manager: WallManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., None]: + from datetime import timedelta + + from generalresearch.models.thl.definitions import WallAdjustedStatus + + def _inner( + session: Session, + created: datetime, + ) -> None: + w1 = wall_manager.get_wall_events(session_id=session.id)[-1] + + # This is defined in `thl-grpc/thl/user_quality_history/recons.py:150` + # so we can't use it as part of this test anyway to add rows to the + # thl_taskadjustment table anyway.. until we created a + # TaskAdjustment Manager to put into generalresearch! + + # create_task_adjustment_event( + # wall, + # user, + # adjusted_status, + # amount_usd=amount_usd, + # alert_time=alert_time, + # ext_status_code=ext_status_code, + # ) + + wall_manager.adjust_status( + wall=w1, + adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, + adjusted_cpi=Decimal("0.00"), + adjusted_timestamp=created, + ) + + thl_lm.create_tx_task_adjustment( + wall=w1, + user=session.user, + created=created + timedelta(milliseconds=1), + ) + + session.wall_events = wall_manager.get_wall_events(session_id=session.id) + session_manager.adjust_status(session=session) + + thl_lm.create_tx_bp_adjustment( + session=session, created=created + timedelta(milliseconds=2) + ) + + return _inner + + +@pytest.fixture +def adj_to_complete_with_tx_factory( + session_manager: SessionManager, + wall_manager: WallManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., None]: + from datetime import timedelta + + from generalresearch.models.thl.definitions import WallAdjustedStatus + + def _inner( + session: Session, + created: datetime, + ) -> None: + w1 = wall_manager.get_wall_events(session_id=session.id)[-1] + + wall_manager.adjust_status( + wall=w1, + adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE, + adjusted_cpi=w1.req_cpi, + adjusted_timestamp=created, + ) + + thl_lm.create_tx_task_adjustment( + wall=w1, + user=session.user, + created=created + timedelta(milliseconds=1), + ) + + session.wall_events = wall_manager.get_wall_events(session_id=session.id) + session_manager.adjust_status(session=session) + + thl_lm.create_tx_bp_adjustment( + session=session, created=created + timedelta(milliseconds=2) + ) + + return _inner diff --git a/test_utils/models/network/__init__.py b/test_utils/models/network/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py new file mode 100644 index 0000000..abfbc18 --- /dev/null +++ b/test_utils/models/network/conftest.py @@ -0,0 +1,144 @@ +import os +from datetime import datetime, timedelta, timezone +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.managers.network.label import IPLabelManager +from generalresearch.managers.network.tool_run import ToolRunManager +from generalresearch.models.network.definitions import IPProtocol +from generalresearch.models.network.mtr.parser import parse_mtr_output +from generalresearch.models.network.mtr.result import MTRResult +from generalresearch.models.network.nmap.parser import parse_nmap_xml +from generalresearch.models.network.nmap.result import NmapResult +from generalresearch.models.network.rdns.parser import parse_rdns_output +from generalresearch.models.network.rdns.result import RDNSResult +from generalresearch.models.network.tool_run import MTRRun, NmapRun, RDNSRun, Status +from generalresearch.models.network.tool_run_command import ( + MTRRunCommand, + MTRRunCommandOptions, + NmapRunCommand, + NmapRunCommandOptions, + RDNSRunCommand, + RDNSRunCommandOptions, +) +from generalresearch.pg_helper import PostgresConfig + + +@pytest.fixture(scope="session") +def scan_group_id() -> str: + return uuid4().hex + + +@pytest.fixture(scope="session") +def iplabel_manager(thl_web_rw: PostgresConfig) -> IPLabelManager: + return IPLabelManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def toolrun_manager(thl_web_rw: PostgresConfig) -> ToolRunManager: + return ToolRunManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def nmap_raw_output(request: Request) -> str: + fp = os.path.join(request.config.rootpath, "data/nmaprun1.xml") + with open(fp) as f: + data = f.read() + return data + + +@pytest.fixture(scope="session") +def nmap_result(nmap_raw_output: str) -> NmapResult: + return parse_nmap_xml(nmap_raw_output) + + +@pytest.fixture(scope="session") +def nmap_run(nmap_result: NmapResult, scan_group_id: str): + r = nmap_result + config = NmapRunCommand( + command="nmap", + options=NmapRunCommandOptions( + ip=r.target_ip, ports="22-1000,11000,1100,3389,61232", top_ports=None + ), + ) + return NmapRun( + tool_version=r.version, + status=Status.SUCCESS, + ip=r.target_ip, + started_at=r.started_at, + finished_at=r.finished_at, + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + ) + + +@pytest.fixture(scope="session") +def dig_raw_output() -> str: + return "156.32.33.45.in-addr.arpa. 300 IN PTR scanme.nmap.org." + + +@pytest.fixture(scope="session") +def rdns_result(dig_raw_output: str) -> RDNSResult: + return parse_rdns_output(ip="45.33.32.156", raw=dig_raw_output) + + +@pytest.fixture(scope="session") +def rdns_run(rdns_result: RDNSResult, scan_group_id: str): + r = rdns_result + ip = "45.33.32.156" + utc_now = datetime.now(tz=timezone.utc) + config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip)) + return RDNSRun( + tool_version="1.2.3", + status=Status.SUCCESS, + ip=ip, + started_at=utc_now, + finished_at=utc_now + timedelta(seconds=1), + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + ) + + +@pytest.fixture(scope="session") +def mtr_raw_output(request: Request) -> str: + fp = os.path.join(request.config.rootpath, "data/mtr_fatbeam.json") + with open(fp) as f: + data = f.read() + return data + + +@pytest.fixture(scope="session") +def mtr_result(mtr_raw_output: str) -> MTRResult: + return parse_mtr_output(mtr_raw_output, port=443, protocol=IPProtocol.TCP) + + +@pytest.fixture(scope="session") +def mtr_run(mtr_result: MTRResult, scan_group_id: str): + r = mtr_result + utc_now = datetime.now(tz=timezone.utc) + config = MTRRunCommand( + command="mtr", + options=MTRRunCommandOptions( + ip=r.destination, protocol=IPProtocol.TCP, port=443 + ), + ) + + return MTRRun( + tool_version="1.2.3", + status=Status.SUCCESS, + ip=r.destination, + started_at=utc_now, + finished_at=utc_now + timedelta(seconds=1), + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + facility_id=1, + source_ip="1.2.3.4", + ) diff --git a/test_utils/models/thl/__init__.py b/test_utils/models/thl/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py new file mode 100644 index 0000000..cf8d2fa --- /dev/null +++ b/test_utils/models/thl/conftest.py @@ -0,0 +1,434 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from decimal import ROUND_DOWN, Decimal +from random import choice as rand_choice +from random import choice as rchoice +from random import randint, random +from typing import Any, Callable +from uuid import uuid4 + +import faker +import pytest +from pydantic import PositiveInt + +from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager +from generalresearch.managers.thl.payout import UserPayoutEventManager +from generalresearch.managers.thl.product import ProductManager +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.user_manager.user_manager import UserManager +from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager +from generalresearch.managers.thl.wall import WallManager +from generalresearch.models import DeviceType +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, +) +from generalresearch.models.legacy.bucket import Bucket +from generalresearch.models.thl.definitions import ( + PayoutStatus, +) +from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType +from generalresearch.models.thl.payout import UserPayoutEvent +from generalresearch.models.thl.product import ( + PayoutConfig, + Product, + ProfilingConfig, + SessionConfig, + SourcesConfig, + SupplyConfig, + UserCreateConfig, + UserHealthConfig, + UserWalletConfig, +) +from generalresearch.models.thl.session import ( + Session, + Source, + Status, + Wall, +) +from generalresearch.models.thl.user import User +from generalresearch.models.thl.user_iphistory import IPRecord +from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel +from generalresearch.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData + +fake = faker.Faker() + + +@pytest.fixture +def wall_status() -> Status: + return Status.COMPLETE + + +@pytest.fixture +def user_factory(user_manager: UserManager) -> Callable[..., User]: + + def _inner( + # --- Create dummy "optional" --- # + product_user_id: str | None = None, + # --- Optional --- # + product_id: UUIDStr | None = None, + product: Product | None = None, + created: datetime | None = None, + ) -> User: + + product_user_id = product_user_id or uuid4().hex + + return user_manager.create_user( + product_user_id=product_user_id, + product_id=product_id, + product=product, + created=created, + ) + + return _inner + + +@pytest.fixture +def wall_factory( + wall_manager: WallManager, session_factory: Session +) -> Callable[..., Wall]: + + def _inner( + session_id: int | None = None, + user_id: int | None = None, + started: datetime | None = None, + source: Source | None = None, + req_survey_id: str | None = None, + req_cpi: Decimal | None = None, + buyer_id: str | None = None, + uuid_id: str | None = None, + ): + """To be used in tests, where we don't care about certain fields""" + + user_id = user_id or fake.random_int(min=1, max=2_147_483_648) + started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + end_date=datetime.now(tz=timezone.utc), + tzinfo=timezone.utc, + ) + + if session_id is None: + # session = SessionManager(pg_config=self.pg_config).create_dummy( + # started=started + # ) + session = session_factory() + session_id = session.id + + source = source or rchoice(list(Source)) + req_survey_id = req_survey_id or uuid4().hex + req_cpi = req_cpi or Decimal(fake.random_int(min=1, max=150) / 100).quantize( + Decimal(".01"), rounding=ROUND_DOWN + ) + + return wall_manager.create( + session_id=session_id, + user_id=user_id, + started=started, + source=source, + req_survey_id=req_survey_id, + req_cpi=req_cpi, + buyer_id=buyer_id, + uuid_id=uuid_id, + ) + + return _inner + + +@pytest.fixture +def product_factory(product_manager: ProductManager) -> Callable[..., Product]: + + def _inner( + product_id: UUIDStr | None = None, + team_id: UUIDStr | None = None, + business_id: UUIDStr | None = None, + name: str | None = None, + redirect_url: str | None = None, + harmonizer_domain: str | None = None, + commission_pct: Decimal = Decimal("0.05000"), + sources_config: SourcesConfig | SupplyConfig | None = None, + payout_config: PayoutConfig | None = None, + session_config: SessionConfig | None = None, + profiling_config: ProfilingConfig | None = None, + user_wallet_config: UserWalletConfig | None = None, + user_create_config: UserCreateConfig | None = None, + user_health_config: UserHealthConfig | None = None, + ) -> Product: + """To be used in tests, where we don't care about certain fields""" + product_id = product_id if product_id else uuid4().hex + team_id = team_id if team_id else uuid4().hex + name = name if name else f"name-{product_id[:12]}" + redirect_url = redirect_url if redirect_url else "https://www.example.com/" + + return product_manager.create( + product_id=product_id, + team_id=team_id, + business_id=business_id, + name=name, + redirect_url=redirect_url, + harmonizer_domain=harmonizer_domain, + commission_pct=commission_pct, + sources_config=sources_config, + payout_config=payout_config, + session_config=session_config, + profiling_config=profiling_config, + user_wallet_config=user_wallet_config, + user_create_config=user_create_config, + user_health_config=user_health_config, + ) + + return _inner + + +@pytest.fixture +def session_factory(session_manager: SessionManager): + + def _inner( + # -- Create Dummy "optional" -- # + started: datetime | None = None, + user: User | None = None, + # -- Optional -- # + country_iso: str | None = None, + device_type: DeviceType | None = None, + ip: str | None = None, + bucket: Bucket | None = None, + url_metadata: dict[str, str] | None = None, + uuid_id: str | None = None, + ) -> Session: + """To be used in tests, where we don't care about certain fields""" + started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + end_date=datetime(year=2000, month=1, day=1, tzinfo=timezone.utc), + tzinfo=timezone.utc, + ) + user = user or User( + user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex + ) + + return session_manager.create( + started=started, + user=user, + country_iso=country_iso, + device_type=device_type, + ip=ip, + bucket=bucket, + url_metadata=url_metadata, + uuid_id=uuid_id, + ) + + return _inner + + +@pytest.fixture +def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: + + def _inner( + geoname_id: PositiveInt | None = None, + continent_code: str | None = None, + continent_name: str | None = None, + country_iso: str | None = None, + country_name: str | None = None, + subdivision_1_iso: str | None = None, + subdivision_1_name: str | None = None, + subdivision_2_iso: str | None = None, + subdivision_2_name: str | None = None, + city_name: str | None = None, + metro_code: int | None = None, + time_zone: str | None = None, + is_in_european_union: bool | None = None, + ) -> IPGeoname: + + return ipgeoname_manager.create( + geoname_id=geoname_id or randint(1, 999_999_999), + continent_code=continent_code or "na", + continent_name=continent_name or "North America", + country_iso=country_iso or "us", + country_name=country_name or "United States", + subdivision_1_iso=subdivision_1_iso or "fl", + subdivision_1_name=subdivision_1_name or "Florida", + subdivision_2_iso=subdivision_2_iso, + subdivision_2_name=subdivision_2_name, + city_name=city_name, + metro_code=metro_code, + time_zone=time_zone, + is_in_european_union=is_in_european_union, + ) + + return _inner + + +def ipinformation_factory( + ipinformation_manager: IPInformationManager, +) -> Callable[..., IPInformation]: + + def _inner( + ip: IPvAnyAddressStr | None = None, + geoname_id: PositiveInt | None = None, + country_iso: str | None = None, + registered_country_iso: str | None = None, + is_anonymous: bool | None = None, + is_anonymous_vpn: bool | None = None, + is_hosting_provider: bool | None = None, + is_public_proxy: bool | None = None, + is_tor_exit_node: bool | None = None, + is_residential_proxy: bool | None = None, + autonomous_system_number: PositiveInt | None = None, + autonomous_system_organization: str | None = None, + domain: str | None = None, + isp: str | None = None, + mobile_country_code: str | None = None, + mobile_network_code: str | None = None, + network: str | None = None, + organization: str | None = None, + static_ip_score: float | None = None, + user_type: UserType | None = None, + postal_code: str | None = None, + latitude: Decimal | None = None, + longitude: Decimal | None = None, + accuracy_radius: int | None = None, + ) -> IPInformation: + + return ipinformation_manager.create( + ip=ip or fake.ipv4_public(), + geoname_id=geoname_id, + country_iso=country_iso or fake.country_code(), + registered_country_iso=registered_country_iso, + is_anonymous=is_anonymous, + is_anonymous_vpn=is_anonymous_vpn, + is_hosting_provider=is_hosting_provider, + is_public_proxy=is_public_proxy, + is_tor_exit_node=is_tor_exit_node, + is_residential_proxy=is_residential_proxy, + autonomous_system_number=autonomous_system_number, + autonomous_system_organization=autonomous_system_organization, + domain=domain, + isp=isp, + mobile_country_code=mobile_country_code, + mobile_network_code=mobile_network_code, + network=network, + organization=organization, + static_ip_score=static_ip_score, + user_type=user_type, + postal_code=postal_code, + latitude=latitude, + longitude=longitude, + accuracy_radius=accuracy_radius, + ) + + return _inner + + +@pytest.fixture +def user_payout_event_factory( + user_payout_event_manager: UserPayoutEventManager, +) -> Callable[..., UserPayoutEvent]: + + def _inner( + uuid: UUIDStr | None = None, + debit_account_uuid: UUIDStr | None = None, + account_reference_type: str | None = None, + account_reference_uuid: UUIDStr | None = None, + cashout_method_uuid: UUIDStr | None = None, + description: str | None = None, + created: AwareDatetimeISO | None = None, + amount: PositiveInt | None = None, + status: PayoutStatus | None = None, + ext_ref_id: str | None = None, + payout_type: PayoutType | None = None, + request_data: dict[str, Any] | None = None, + order_data: dict[str, Any] | CashMailOrderData | None = None, + ) -> UserPayoutEvent: + + debit_account_uuid = debit_account_uuid or uuid4().hex + cashout_method_uuid = cashout_method_uuid or uuid4().hex + # account_reference_type = account_reference_type or f"acct-ref-{uuid4().hex}" + # account_reference_uuid = account_reference_uuid or uuid4().hex + # cashout_method_uuid = cashout_method_uuid or uuid4().hex + amount = amount or randint(a=99, b=9_999) + status = status or rand_choice(list(PayoutStatus)) + + description = description or f"desc-{uuid4().hex[:12]}" + # ext_ref_id = ext_ref_id or f"ext-ref-{uuid4().hex[:8]}" + payout_type = payout_type or rand_choice(list(PayoutType)) + request_data = request_data or {} + # order_data = order_data or None + + return user_payout_event_manager.create( + uuid=uuid, + debit_account_uuid=debit_account_uuid, + account_reference_type=account_reference_type, + account_reference_uuid=account_reference_uuid, + cashout_method_uuid=cashout_method_uuid, + description=description, + created=created, + amount=amount, + status=status, + ext_ref_id=ext_ref_id, + payout_type=payout_type, + request_data=request_data, + order_data=order_data, + ) + + return _inner + + +@pytest.fixture +def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecord]: + + def _inner( + user_id: PositiveInt, + ip: IPvAnyAddressStr | None = None, + forwarded_ip1: IPvAnyAddressStr | None = None, + forwarded_ip2: IPvAnyAddressStr | None = None, + forwarded_ip3: IPvAnyAddressStr | None = None, + forwarded_ip4: IPvAnyAddressStr | None = None, + forwarded_ip5: IPvAnyAddressStr | None = None, + forwarded_ip6: IPvAnyAddressStr | None = None, + ) -> IPRecord: + return iprecord_manager.create( + user_id=user_id, + ip=ip or fake.ipv4_public(), + forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), + forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None), + forwarded_ip3=( + forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None + ), + forwarded_ip4=forwarded_ip4, + forwarded_ip5=forwarded_ip5, + forwarded_ip6=forwarded_ip6, + ) + + return _inner + + +# class AuditLogManager(PostgresManager): + + +@pytest.fixture +def auditlog_factory(audit_log_manager: AuditLogManager): + + def _inner( + user_id: PositiveInt, + level: AuditLogLevel | None = None, + event_type: str | None = None, + event_msg: str | None = None, + event_value: float | None = None, + ) -> AuditLog: + + event_types = { + "offerwall-enter.blocked", + "offerwall-enter.rate-limited", + "offerwall-enter.url-modified", + } + + return audit_log_manager.create( + user_id=user_id, + level=level or rchoice(list(AuditLogLevel)), + event_type=event_type or rchoice(list(event_types)), + event_msg=event_msg, + event_value=event_value, + ) + + return _inner diff --git a/test_utils/models/upk/__init__.py b/test_utils/models/upk/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py new file mode 100644 index 0000000..c8855da --- /dev/null +++ b/test_utils/models/upk/conftest.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import os +import time +from typing import TYPE_CHECKING +from uuid import UUID + +import pandas as pd +import pytest + +from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.category import CategoryManager + + +def insert_data_from_csv( + thl_web_rw: PostgresConfig, + table_name: str, + fp: str | None = None, + disable_fk_checks: bool = False, + df: pd.DataFrame | None = None, +): + assert fp is not None or df is not None and not (fp is not None and df is not None) + if fp: + df = pd.read_csv(fp, dtype=str) + + assert isinstance(df, pd.DataFrame) + + df = df.where(pd.notnull(df), None) + cols = list(df.columns) + col_str = ", ".join(cols) + values_str = ", ".join(["%s"] * len(cols)) + if "id" in df.columns and len(df["id"].iloc[0]) == 36: + df["id"] = df["id"].map(lambda x: UUID(x).hex) + args = df.to_dict("tight")["data"] + + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + if disable_fk_checks: + c.execute("SET CONSTRAINTS ALL DEFERRED") + c.executemany( + f"INSERT INTO {table_name} ({col_str}) VALUES ({values_str})", + params_seq=args, + ) + conn.commit() + + +@pytest.fixture(scope="session") +def category_data( + thl_web_rw: PostgresConfig, category_manager: CategoryManager +) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_category.csv.gz") + insert_data_from_csv( + thl_web_rw, + fp=fp, + table_name="marketplace_category", + disable_fk_checks=True, + ) + # Don't strictly need to do this, but probably we should + category_manager.populate_caches() + cats = category_manager.categories.values() + path_id = {c.path: c.id for c in cats} + data = [ + {"id": c.id, "parent_id": path_id[c.parent_path]} for c in cats if c.parent_path + ] + query = """ + UPDATE marketplace_category + SET parent_id = %(parent_id)s + WHERE id = %(id)s; + """ + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.executemany(query=query, params_seq=data) + conn.commit() + + +@pytest.fixture(scope="session") +def property_data(thl_web_rw: PostgresConfig) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_property.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_property") + + +@pytest.fixture(scope="session") +def item_data(thl_web_rw: PostgresConfig) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_item.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_item") + + +@pytest.fixture(scope="session") +def propertycategoryassociation_data( + thl_web_rw: PostgresConfig, + category_data, + property_data, + category_manager: CategoryManager, +) -> None: + table_name = "marketplace_propertycategoryassociation" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + # Need to lookup category pk from uuid + category_manager.populate_caches() + df = pd.read_csv(fp, dtype=str) + df["category_id"] = df["category_id"].map( + lambda x: category_manager.categories[x].id + ) + insert_data_from_csv(thl_web_rw, df=df, table_name=table_name) + + +@pytest.fixture(scope="session") +def propertycountry_data(thl_web_rw: PostgresConfig, property_data) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_propertycountry.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_propertycountry") + + +@pytest.fixture(scope="session") +def propertymarketplaceassociation_data( + thl_web_rw: PostgresConfig, property_data +) -> None: + table_name = "marketplace_propertymarketplaceassociation" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) + + +@pytest.fixture(scope="session") +def propertyitemrange_data( + thl_web_rw: PostgresConfig, property_data, item_data +) -> None: + table_name = "marketplace_propertyitemrange" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) + + +@pytest.fixture(scope="session") +def question_data(thl_web_rw: PostgresConfig) -> None: + table_name = "marketplace_question" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv( + thl_web_rw, fp=fp, table_name=table_name, disable_fk_checks=True + ) + + +@pytest.fixture(scope="session") +def clear_upk_tables(thl_web_rw: PostgresConfig): + tables = [ + "marketplace_propertyitemrange", + "marketplace_propertymarketplaceassociation", + "marketplace_propertycategoryassociation", + "marketplace_category", + "marketplace_item", + "marketplace_property", + "marketplace_propertycountry", + "marketplace_question", + ] + table_str = ", ".join(tables) + + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.execute(f"TRUNCATE {table_str} RESTART IDENTITY CASCADE;") + conn.commit() + + +@pytest.fixture(scope="session") +def upk_data( + clear_upk_tables, + category_data, + property_data, + item_data, + propertycategoryassociation_data, + propertycountry_data, + propertymarketplaceassociation_data, + propertyitemrange_data, + question_data, +) -> None: + # Wait a second to make sure the HarmonizerCache refresh loop pulls these in + time.sleep(2) + + +def test_fixtures(upk_data): + pass diff --git a/test_utils/models/upk/marketplace_category.csv.gz b/test_utils/models/upk/marketplace_category.csv.gz new file mode 100644 index 0000000..0f8ec1c Binary files /dev/null and b/test_utils/models/upk/marketplace_category.csv.gz differ diff --git a/test_utils/models/upk/marketplace_item.csv.gz b/test_utils/models/upk/marketplace_item.csv.gz new file mode 100644 index 0000000..c12c5d8 Binary files /dev/null and b/test_utils/models/upk/marketplace_item.csv.gz differ diff --git a/test_utils/models/upk/marketplace_property.csv.gz b/test_utils/models/upk/marketplace_property.csv.gz new file mode 100644 index 0000000..a781d1d Binary files /dev/null and b/test_utils/models/upk/marketplace_property.csv.gz differ diff --git a/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz new file mode 100644 index 0000000..5b4ea19 Binary files /dev/null and b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz differ diff --git a/test_utils/models/upk/marketplace_propertycountry.csv.gz b/test_utils/models/upk/marketplace_propertycountry.csv.gz new file mode 100644 index 0000000..5d2a637 Binary files /dev/null and b/test_utils/models/upk/marketplace_propertycountry.csv.gz differ diff --git a/test_utils/models/upk/marketplace_propertyitemrange.csv.gz b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz new file mode 100644 index 0000000..84f4f0e Binary files /dev/null and b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz differ diff --git a/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz new file mode 100644 index 0000000..6b9fd1c Binary files /dev/null and b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz differ diff --git a/test_utils/models/upk/marketplace_question.csv.gz b/test_utils/models/upk/marketplace_question.csv.gz new file mode 100644 index 0000000..bcfc3ad Binary files /dev/null and b/test_utils/models/upk/marketplace_question.csv.gz differ diff --git a/tests/conftest.py b/tests/conftest.py index 2482269..6748592 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,8 +3,6 @@ pytest_plugins = [ "test_utils.conftest", # -- GRL IQ "test_utils.grliq.conftest", - "test_utils.grliq.managers.conftest", - "test_utils.grliq.models.conftest", # -- Incite "test_utils.incite.conftest", "test_utils.incite.collections.conftest", @@ -12,9 +10,17 @@ pytest_plugins = [ # -- Managers "test_utils.managers.conftest", "test_utils.managers.contest.conftest", + "test_utils.managers.gr.conftest", "test_utils.managers.ledger.conftest", "test_utils.managers.network.conftest", + "test_utils.managers.thl.conftest", "test_utils.managers.upk.conftest", # -- Models "test_utils.models.conftest", + "test_utils.models.contest.conftest", + "test_utils.models.gr.conftest", + "test_utils.models.ledger.conftest", + "test_utils.models.network.conftest", + "test_utils.models.thl.conftest", + "test_utils.models.upk.conftest", ] diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py index ac2792a..e4854e8 100644 --- a/tests/grliq/managers/test_forensic_data.py +++ b/tests/grliq/managers/test_forensic_data.py @@ -1,20 +1,23 @@ +from __future__ import annotations + from datetime import timedelta from typing import TYPE_CHECKING from uuid import uuid4 import pytest +from generalresearch.grliq.models.events import MouseEvent, TimingData +from generalresearch.grliq.models.forensic_data import GrlIqData +from generalresearch.grliq.models.forensic_result import ( + GrlIqCheckerResults, + GrlIqForensicCategoryResult, +) + if TYPE_CHECKING: from generalresearch.grliq.managers.forensic_data import ( GrlIqDataManager, GrlIqEventManager, ) - from generalresearch.grliq.models.events import MouseEvent, TimingData - from generalresearch.grliq.models.forensic_data import GrlIqData - from generalresearch.grliq.models.forensic_result import ( - GrlIqCheckerResults, - GrlIqForensicCategoryResult, - ) from generalresearch.models.thl.product import Product try: @@ -25,7 +28,7 @@ except ImportError: class TestGrlIqDataManager: - def test_create_dummy(self, grliq_dm: "GrlIqDataManager"): + def test_create_dummy(self, grliq_dm: GrlIqDataManager): from generalresearch.grliq.models.forensic_data import GrlIqData gd1: GrlIqData = grliq_dm.create_dummy(is_attempt_allowed=True) @@ -34,7 +37,7 @@ class TestGrlIqDataManager: assert isinstance(gd1.results, GrlIqCheckerResults) assert isinstance(gd1.category_result, GrlIqForensicCategoryResult) - def test_create(self, grliq_data: "GrlIqData", grliq_dm: "GrlIqDataManager"): + def test_create(self, grliq_data: GrlIqData, grliq_dm: GrlIqDataManager): grliq_dm.create(grliq_data) assert grliq_data.id is not None @@ -53,13 +56,13 @@ class TestGrlIqDataManager: def test_update_data(self): pass - def test_get_id(self, grliq_data: "GrlIqData", grliq_dm: "GrlIqDataManager"): + def test_get_id(self, grliq_data: GrlIqData, grliq_dm: GrlIqDataManager): grliq_dm.create(grliq_data) res = grliq_dm.get_data(forensic_id=grliq_data.id) assert res == grliq_data - def test_get_uuid(self, grliq_data: "GrlIqData", grliq_dm: "GrlIqDataManager"): + def test_get_uuid(self, grliq_data: GrlIqData, grliq_dm: GrlIqDataManager): grliq_dm.create(grliq_data) res = grliq_dm.get_data(forensic_uuid=grliq_data.uuid) @@ -73,7 +76,7 @@ class TestGrlIqDataManager: def test_get_unique_user_count_by_fingerprint(self): pass - def test_filter_data(self, grliq_data: "GrlIqData", grliq_dm: "GrlIqDataManager"): + def test_filter_data(self, grliq_data: GrlIqData, grliq_dm: GrlIqDataManager): grliq_dm.create(grliq_data) res = grliq_dm.filter_data(uuids=[grliq_data.uuid])[0] assert res == grliq_data @@ -100,7 +103,7 @@ class TestGrlIqDataManager: def test_make_filter_str(self): pass - def test_filter_count(self, grliq_dm: "GrlIqDataManager", product: "Product"): + def test_filter_count(self, grliq_dm: GrlIqDataManager, product: Product): res = grliq_dm.filter_count(product_id=product.uuid) assert isinstance(res, int) @@ -116,7 +119,7 @@ class TestGrlIqDataManager: class TestForensicDataGetAndFilter: - def test_events(self, grliq_dm: "GrlIqDataManager"): + def test_events(self, grliq_dm: GrlIqDataManager): """If load_events=True, the events and mouse_events attributes should be an array no matter what. An empty array means that the events were loaded, but there were no events available. @@ -141,7 +144,7 @@ class TestForensicDataGetAndFilter: assert len(instance.events) == 0 assert len(instance.mouse_events) == 0 - def test_timing(self, grliq_dm: "GrlIqDataManager", grliq_em: "GrlIqEventManager"): + def test_timing(self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager): forensic_uuid = uuid4().hex grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) @@ -161,7 +164,7 @@ class TestForensicDataGetAndFilter: assert isinstance(instance.timing_data, TimingData) def test_events_events( - self, grliq_dm: "GrlIqDataManager", grliq_em: "GrlIqEventManager" + self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager ): forensic_uuid = uuid4().hex grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) @@ -186,7 +189,7 @@ class TestForensicDataGetAndFilter: assert len(instance.keyboard_events) == 0 def test_events_click( - self, grliq_dm: "GrlIqDataManager", grliq_em: "GrlIqEventManager" + self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager ): forensic_uuid = uuid4().hex grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) diff --git a/tests/grliq/managers/test_forensic_results.py b/tests/grliq/managers/test_forensic_results.py index 68db732..a030451 100644 --- a/tests/grliq/managers/test_forensic_results.py +++ b/tests/grliq/managers/test_forensic_results.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -10,7 +12,7 @@ if TYPE_CHECKING: class TestGrlIqCategoryResultsReader: def test_filter_category_results( - self, grliq_dm: "GrlIqDataManager", grliq_crr: "GrlIqCategoryResultsReader" + self, grliq_dm: GrlIqDataManager, grliq_crr: GrlIqCategoryResultsReader ): from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py index cf4c405..a80afbe 100644 --- a/tests/models/admin/test_report_request.py +++ b/tests/models/admin/test_report_request.py @@ -1,4 +1,4 @@ -from datetime import timezone, datetime +from datetime import datetime, timezone import pandas as pd import pytest @@ -6,7 +6,7 @@ from pydantic import ValidationError class TestReportRequest: - def test_base(self, utc_60days_ago): + def test_base(self): from generalresearch.models.admin.request import ( ReportRequest, ReportType, @@ -24,7 +24,7 @@ class TestReportRequest: rr1 = ReportRequest.model_validate( { "start": datetime( - year=datetime.now().year, + year=datetime.now(tz=timezone.utc).year, month=1, day=1, hour=0, @@ -43,7 +43,7 @@ class TestReportRequest: rr2 = ReportRequest.model_validate( { "start": datetime( - year=datetime.now().year, + year=datetime.now(tz=timezone.utc).year, month=1, day=1, hour=6, @@ -81,29 +81,30 @@ class TestReportRequest: # interval='1d', include_open_bucket=True, # start_floor=datetime.datetime(2025, 7, 9, 0, 0, tzinfo=datetime.timezone.utc)).start_floor - def test_start_end_range(self, utc_90days_ago, utc_30days_ago): + def test_start_end_range(self, utc_90days_ago: datetime, utc_30days_ago: datetime): from generalresearch.models.admin.request import ReportRequest - with pytest.raises(expected_exception=ValidationError) as cm: + with pytest.raises(expected_exception=ValidationError): ReportRequest.model_validate( {"start": utc_30days_ago, "end": utc_90days_ago} ) - with pytest.raises(expected_exception=ValidationError) as cm: + with pytest.raises(expected_exception=ValidationError): ReportRequest.model_validate( { - "start": datetime(year=1990, month=1, day=1), - "end": datetime(year=1950, month=1, day=1), + "start": datetime(year=1990, month=1, day=1, tzinfo=timezone.utc), + "end": datetime(year=1950, month=1, day=1, tzinfo=timezone.utc), } ) def test_start_end_range_tz(self): - from generalresearch.models.admin.request import ReportRequest from zoneinfo import ZoneInfo + from generalresearch.models.admin.request import ReportRequest + pacific_tz = ZoneInfo("America/Los_Angeles") - with pytest.raises(expected_exception=ValidationError) as cm: + with pytest.raises(expected_exception=ValidationError): ReportRequest.model_validate( { "start": datetime(year=2000, month=1, day=1, tzinfo=pacific_tz), diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index b37f2c4..16e1f83 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -2,13 +2,11 @@ from typing import Optional from uuid import uuid4 import pytest -from pydantic import BaseModel, ValidationError, Field -from pydantic import MySQLDsn +from pydantic import BaseModel, Field, MySQLDsn, ValidationError from pydantic_core import Url from generalresearch.models.custom_types import DaskDsn, SentryDsn - # --- Test Pydantic Models --- @@ -27,7 +25,7 @@ class TestDaskDsn: from dask.distributed import Client m = SettingsModel(dask="tcp://dask-scheduler.internal") - + assert isinstance(m.dask, Url) assert m.dask.scheme == "tcp" assert m.dask.host == "dask-scheduler.internal" assert m.dask.port == 8786 @@ -72,6 +70,7 @@ class TestDaskDsn: def test_port(self): m = SettingsModel(dask="tcp://dask-scheduler.internal") + assert isinstance(m.dask, Url) assert m.dask.port == 8786 @@ -81,6 +80,7 @@ class TestSentryDsn: sentry=f"https://{uuid4().hex}@12345.ingest.us.sentry.io/9876543" ) + assert isinstance(m.sentry, Url) assert m.sentry.scheme == "https" assert m.sentry.host == "12345.ingest.us.sentry.io" assert m.sentry.port == 443 @@ -109,4 +109,5 @@ class TestSentryDsn: def test_port(self): test_url: str = f"https://{uuid4().hex}@12345.ingest.us.sentry.io/9876543" m = SettingsModel(sentry=test_url) + assert isinstance(m.sentry, Url) assert m.sentry.port == 443 diff --git a/tests/models/custom_types/test_uuid_str.py b/tests/models/custom_types/test_uuid_str.py index 91af9ae..02e6a8b 100644 --- a/tests/models/custom_types/test_uuid_str.py +++ b/tests/models/custom_types/test_uuid_str.py @@ -1,14 +1,15 @@ -from typing import Optional +from __future__ import annotations + from uuid import uuid4 import pytest -from pydantic import BaseModel, ValidationError, Field +from pydantic import BaseModel, Field, ValidationError from generalresearch.models.custom_types import UUIDStr class UUIDStrModel(BaseModel): - uuid_optional: Optional[UUIDStr] = Field(default_factory=lambda: uuid4().hex) + uuid_optional: UUIDStr | None = Field(default_factory=lambda: uuid4().hex) uuid: UUIDStr diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py index 23437f5..736c971 100644 --- a/tests/models/dynata/test_eligbility.py +++ b/tests/models/dynata/test_eligbility.py @@ -5,10 +5,10 @@ class TestEligibility: def test_evaluate_task_criteria(self): from generalresearch.models.dynata.survey import ( - DynataQuotaGroup, DynataFilterGroup, - DynataSurvey, + DynataQuotaGroup, DynataRequirements, + DynataSurvey, ) filters = [[["a", "b"], ["c", "d"]], [["e"], ["f"]]] @@ -137,10 +137,10 @@ class TestEligibility: def test_soft_pair(self): from generalresearch.models.dynata.survey import ( - DynataQuotaGroup, DynataFilterGroup, - DynataSurvey, + DynataQuotaGroup, DynataRequirements, + DynataSurvey, ) filters = [[["a", "b"], ["c", "d"]], [["e"], ["f"]]] @@ -186,7 +186,7 @@ class TestEligibility: } ) assert task.passes_filters(criteria_evaluation) - passes, condition_hashes = task.passes_filters_soft(criteria_evaluation) + passes, _ = task.passes_filters_soft(criteria_evaluation) assert passes # make 'e' & 'f' None, we don't pass the 2nd filtergroup diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index e906d8c..6c84a5d 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -3,17 +3,20 @@ import json import os from datetime import datetime, timezone from random import randint +from typing import Callable from uuid import uuid4 import pytest +from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.team import Membership, Team + SSO_ISSUER = "" class TestGRUser: - def test_init(self, gr_user): - from generalresearch.models.gr.authentication import GRUser + def test_init(self, gr_user: GRUser): assert isinstance(gr_user, GRUser) assert not gr_user.is_superuser @@ -26,8 +29,7 @@ class TestGRUser: def test_businesses(self): pass - def test_teams(self, gr_user, membership, gr_db, gr_redis_config): - from generalresearch.models.gr.team import Team + def test_teams(self, gr_user: GRUser, membership, gr_db, gr_redis_config): assert gr_user.teams is None @@ -40,11 +42,11 @@ class TestGRUser: def test_prefetch_team_duplicates( self, gr_user_token, - gr_user, - membership, + gr_user: GRUser, + membership: Membership, product_factory, membership_factory, - team, + team: Team, thl_web_rr, gr_redis_config, gr_db, @@ -61,10 +63,10 @@ class TestGRUser: def test_products( self, - gr_user, + gr_user: GRUser, product_factory, - team, - membership, + team: Team, + membership: Membership, gr_db, thl_web_rr, gr_redis_config, @@ -102,12 +104,12 @@ class TestGRUserMethods: def test_to_redis( self, - gr_user, + gr_user: GRUser, gr_redis, - team, + team: Team, business, product_factory, - membership_factory, + membership_factory: Callable[Membership], ): product_factory(team=team, business=business) membership_factory(team=team, gr_user=gr_user) @@ -122,7 +124,7 @@ class TestGRUserMethods: def test_set_cache( self, - gr_user, + gr_user: GRUser, gr_user_token, gr_redis, gr_db, @@ -145,7 +147,7 @@ class TestGRUserMethods: def test_set_cache_gr_user( self, - gr_user, + gr_user: GRUser, gr_user_token, gr_redis, gr_redis_config, @@ -203,9 +205,7 @@ class TestGRUserMethods: @pytest.mark.skip def test_set_cache_business_uuids( self, - gr_user, - membership, - gr_user_token, + gr_user: GRUser, gr_redis, gr_db, thl_web_rr, diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py new file mode 100644 index 0000000..323d7b6 --- /dev/null +++ b/tests/models/gr/test_base.py @@ -0,0 +1,25 @@ +from typing import Callable + +from pydantic import PostgresDsn + +from generalresearch.pg_helper import PostgresConfig + + +class TestGRPostgresDjangoCreation: + + def test_django_creation( + self, + django_db_factory: Callable[..., None], + ): + + dsn = django_db_factory("gr_carer") + assert isinstance(dsn, PostgresDsn) + + def test_django_tables(self, thl_web_rw: PostgresConfig): + res = thl_web_rw.execute_sql_query(query=""" + SELECT COUNT(*) + FROM information_schema.tables + WHERE table_schema = 'public'; + """) + assert len(res) == 1 + assert res[0]["count"] == 56 diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index e8bd06a..7a84f23 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -6,6 +6,7 @@ from uuid import uuid4 import pandas as pd import pytest +from dask.distributed import Client as DaskClient # noinspection PyUnresolvedReferences from distributed.utils_test import ( @@ -14,22 +15,27 @@ from distributed.utils_test import ( from pytest import approx from generalresearch.currency import USDCent +from generalresearch.managers.gr.business import BusinessBankAccountManager +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + BusinessContact, +) from generalresearch.models.thl.finance import ( BusinessBalances, ProductBalances, ) - -# from test_utils.incite.conftest import mnt_filepath -from test_utils.managers.conftest import ( - business_bank_account_manager, - lm, - thl_lm, -) +from generalresearch.pg_helper import PostgresConfig class TestBusinessBankAccount: - def test_init(self, business, business_bank_account_manager): + def test_init( + self, + business: Business, + business_bank_account_manager: BusinessBankAccountManager, + ): from generalresearch.models.gr.business import ( BusinessBankAccount, TransferMethod, @@ -42,7 +48,13 @@ class TestBusinessBankAccount: ) assert isinstance(instance, BusinessBankAccount) - def test_business(self, business_bank_account, business, gr_db, gr_redis_config): + def test_business( + self, + business_bank_account: BusinessBankAccount, + business: Business, + gr_db, + gr_redis_config, + ): from generalresearch.models.gr.business import Business assert business_bank_account.business is None @@ -56,16 +68,13 @@ class TestBusinessBankAccount: class TestBusinessAddress: - def test_init(self, business_address): - from generalresearch.models.gr.business import BusinessAddress - + def test_init(self, business_address: BusinessAddress): assert isinstance(business_address, BusinessAddress) class TestBusinessContact: def test_init(self): - from generalresearch.models.gr.business import BusinessContact bc = BusinessContact(name="abc", email="test@abc.com") assert isinstance(bc, BusinessContact) @@ -104,7 +113,7 @@ class TestBusiness: user_factory, session_with_tx_factory, pop_ledger_merge, - client_no_amm, + client_no_amm: DaskClient, ledger_collection, mnt_filepath, create_main_accounts, @@ -220,11 +229,11 @@ class TestBusiness: def test_balance( self, - business, + business: Business, mnt_filepath, - client_no_amm, - thl_web_rr, - lm, + client_no_amm: DaskClient, + thl_web_rr: PostgresConfig, + ledger_manager, pop_ledger_merge, ): assert business.balance is None @@ -232,7 +241,7 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_balance( thl_pg_config=thl_web_rr, - lm=lm, + lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -248,7 +257,7 @@ class TestBusiness: business, product_factory, thl_web_rr, - thl_lm, + thl_ledger_manager, business_payout_event_manager, ): assert business.payouts is None @@ -256,17 +265,17 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert "Must provide product_uuids" in str(cm.value) p = product_factory(business=business) - thl_lm.get_account_or_create_bp_wallet(product=p) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p) business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert isinstance(business.payouts, list) @@ -274,17 +283,17 @@ class TestBusiness: def test_payouts( self, - business, - product_factory, + business: Business, + product_factory: Callable[Product], bp_payout_factory, - thl_lm, + thl_ledger_manager, thl_web_rr, business_payout_event_manager, create_main_accounts, ): create_main_accounts() p = product_factory(business=business) - thl_lm.get_account_or_create_bp_wallet(product=p) + thl_ledger_manager.get_account_or_create_bp_wallet(product=p) business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) bp_payout_factory( @@ -293,7 +302,7 @@ class TestBusiness: business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert len(business.payouts) == 1 @@ -306,10 +315,12 @@ class TestBusiness: skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + business_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) business.prebuild_payouts( thl_pg_config=thl_web_rr, - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert len(business.payouts) == 1 @@ -370,7 +381,7 @@ class TestBusiness: self, business, thl_web_rr, - thl_lm, + thl_ledger_manager, mnt_filepath, client_no_amm, pop_ledger_merge, @@ -496,7 +507,7 @@ class TestBusinessBalance: mnt_filepath, bp_payout_factory, thl_lm, - lm, + ledger_manager, duration, offset, start, -- cgit v1.2.3 From c4a44873540ca4c0a3ab19b9beef4cfc6e0252a7 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 21 Aug 2026 11:19:36 -0700 Subject: last files for Greg to merge --- test_utils/conftest.py | 14 ++++++++++++-- tests/models/gr/test_base.py | 39 ++++++++++++++++++++++++++++++--------- 2 files changed, 42 insertions(+), 11 deletions(-) (limited to 'tests') diff --git a/test_utils/conftest.py b/test_utils/conftest.py index cd7f282..378b9cc 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -4,6 +4,7 @@ import os import shutil import stat import subprocess +import sys import tempfile from datetime import datetime, timedelta, timezone from os.path import join as pjoin @@ -174,7 +175,7 @@ def git_key_path( @pytest.fixture(scope="session") -def gr_models(git_key_path: Path) -> Callable[..., Path]: +def gr_repo(git_key_path: Path) -> Callable[..., Path]: repo_url = "ssh://code.g-r-l.com/general-research/gr-carer.git" repo_path = Path("/tmp/gr-carer") @@ -202,7 +203,9 @@ def gr_models(git_key_path: Path) -> Callable[..., Path]: @pytest.fixture(scope="session") def django_db_factory( - postgres_instance: PostgresDsn, postgres_instance_dict: PostgresDict + postgres_instance: PostgresDsn, + postgres_instance_dict: PostgresDict, + gr_repo: Callable[..., Path], ) -> Callable[..., PostgresDsn]: import django @@ -211,6 +214,13 @@ def django_db_factory( def _inner(django_project: str = "generalresearch.thl_django"): + if "gr" in django_project: + # We need model files that are NOT in this repo. + gr_path = gr_repo() + sys.path.insert(0, str(gr_path)) + + print(sys.path) + # 1. Bootstrapping Django settings if not django_settings.configured: django_settings.configure( diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index 323d7b6..a9f01a8 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -1,5 +1,8 @@ +import subprocess +from pathlib import Path from typing import Callable +import pytest from pydantic import PostgresDsn from generalresearch.pg_helper import PostgresConfig @@ -7,19 +10,37 @@ from generalresearch.pg_helper import PostgresConfig class TestGRPostgresDjangoCreation: + def test_git(self, git_key_path: Path, gr_repo: Callable[..., Path]): + repo_path = gr_repo() + + try: + # Run the git command inside the target directory + result = subprocess.run( + ["git", "rev-parse", "--is-inside-work-tree"], + cwd=repo_path, + capture_output=True, + text=True, + check=True, + ) + # Check if the output string is exactly "true" + assert result.stdout.strip() == "true" + + except (subprocess.CalledProcessError, FileNotFoundError) as e: + pytest.fail(f"Directory is not a Git repo or Git is not installed: {e}") + def test_django_creation( self, django_db_factory: Callable[..., None], ): - dsn = django_db_factory("gr_carer") + dsn = django_db_factory("gr") assert isinstance(dsn, PostgresDsn) - def test_django_tables(self, thl_web_rw: PostgresConfig): - res = thl_web_rw.execute_sql_query(query=""" - SELECT COUNT(*) - FROM information_schema.tables - WHERE table_schema = 'public'; - """) - assert len(res) == 1 - assert res[0]["count"] == 56 + # def test_django_tables(self, thl_web_rw: PostgresConfig): + # res = thl_web_rw.execute_sql_query(query=""" + # SELECT COUNT(*) + # FROM information_schema.tables + # WHERE table_schema = 'public'; + # """) + # assert len(res) == 1 + # assert res[0]["count"] == 56 -- cgit v1.2.3