aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--.gitignore3
-rw-r--r--Jenkinsfile41
-rw-r--r--generalresearch/config.py20
-rw-r--r--generalresearch/grliq/managers/forensic_data.py330
-rw-r--r--generalresearch/incite/base.py170
-rw-r--r--generalresearch/managers/gr/authentication.py79
-rw-r--r--generalresearch/managers/gr/business.py175
-rw-r--r--generalresearch/managers/gr/team.py93
-rw-r--r--generalresearch/managers/leaderboard/__init__.py8
-rw-r--r--generalresearch/managers/thl/delete_request.py178
-rw-r--r--generalresearch/managers/thl/ipinfo.py106
-rw-r--r--generalresearch/managers/thl/payout.py64
-rw-r--r--generalresearch/managers/thl/product.py42
-rw-r--r--generalresearch/managers/thl/session.py40
-rw-r--r--generalresearch/managers/thl/user_manager/user_manager.py20
-rw-r--r--generalresearch/managers/thl/userhealth.py66
-rw-r--r--generalresearch/managers/thl/wall.py51
-rw-r--r--generalresearch/models/custom_types.py41
-rw-r--r--generalresearch/models/gr/business.py2
-rw-r--r--generalresearch/models/network/nmap/parser.py2
-rw-r--r--generalresearch/models/network/rdns/parser.py2
-rw-r--r--pyproject.toml2
-rw-r--r--test_utils/conftest.py401
-rw-r--r--test_utils/grliq/conftest.py124
-rw-r--r--test_utils/incite/collections/conftest.py38
-rw-r--r--test_utils/incite/conftest.py37
-rw-r--r--test_utils/incite/mergers/conftest.py105
-rw-r--r--test_utils/managers/conftest.py578
-rw-r--r--test_utils/managers/contest/conftest.py294
-rw-r--r--test_utils/managers/gr/__init__.py (renamed from test_utils/grliq/managers/__init__.py)0
-rw-r--r--test_utils/managers/gr/conftest.py110
-rw-r--r--test_utils/managers/ledger/conftest.py777
-rw-r--r--test_utils/managers/network/conftest.py143
-rw-r--r--test_utils/managers/thl/__init__.py (renamed from test_utils/grliq/models/__init__.py)0
-rw-r--r--test_utils/managers/thl/conftest.py258
-rw-r--r--test_utils/managers/upk/conftest.py188
-rw-r--r--test_utils/models/conftest.py249
-rw-r--r--test_utils/models/contest/__init__.py (renamed from test_utils/grliq/managers/conftest.py)0
-rw-r--r--test_utils/models/contest/conftest.py292
-rw-r--r--test_utils/models/gr/__init__.py (renamed from test_utils/grliq/models/conftest.py)0
-rw-r--r--test_utils/models/gr/conftest.py213
-rw-r--r--test_utils/models/ledger/__init__.py0
-rw-r--r--test_utils/models/ledger/conftest.py724
-rw-r--r--test_utils/models/network/__init__.py0
-rw-r--r--test_utils/models/network/conftest.py144
-rw-r--r--test_utils/models/thl/__init__.py0
-rw-r--r--test_utils/models/thl/conftest.py434
-rw-r--r--test_utils/models/upk/__init__.py0
-rw-r--r--test_utils/models/upk/conftest.py178
-rw-r--r--test_utils/models/upk/marketplace_category.csv.gz (renamed from test_utils/managers/upk/marketplace_category.csv.gz)bin100990 -> 100990 bytes
-rw-r--r--test_utils/models/upk/marketplace_item.csv.gz (renamed from test_utils/managers/upk/marketplace_item.csv.gz)bin3225 -> 3225 bytes
-rw-r--r--test_utils/models/upk/marketplace_property.csv.gz (renamed from test_utils/managers/upk/marketplace_property.csv.gz)bin3315 -> 3315 bytes
-rw-r--r--test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz (renamed from test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz)bin2079 -> 2079 bytes
-rw-r--r--test_utils/models/upk/marketplace_propertycountry.csv.gz (renamed from test_utils/managers/upk/marketplace_propertycountry.csv.gz)bin71359 -> 71359 bytes
-rw-r--r--test_utils/models/upk/marketplace_propertyitemrange.csv.gz (renamed from test_utils/managers/upk/marketplace_propertyitemrange.csv.gz)bin65389 -> 65389 bytes
-rw-r--r--test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz (renamed from test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz)bin4272 -> 4272 bytes
-rw-r--r--test_utils/models/upk/marketplace_question.csv.gz (renamed from test_utils/managers/upk/marketplace_question.csv.gz)bin283465 -> 283465 bytes
-rw-r--r--tests/conftest.py10
-rw-r--r--tests/grliq/managers/test_forensic_data.py35
-rw-r--r--tests/grliq/managers/test_forensic_results.py4
-rw-r--r--tests/incite/collections/test_df_collection_item_thl_web.py26
-rw-r--r--tests/incite/collections/test_df_collection_thl_marketplaces.py2
-rw-r--r--tests/incite/collections/test_df_collection_thl_web.py27
-rw-r--r--tests/incite/test_collection_base.py6
-rw-r--r--tests/models/admin/test_report_request.py23
-rw-r--r--tests/models/custom_types/test_aware_datetime.py7
-rw-r--r--tests/models/custom_types/test_dsn.py9
-rw-r--r--tests/models/custom_types/test_uuid_str.py7
-rw-r--r--tests/models/dynata/test_eligbility.py10
-rw-r--r--tests/models/gr/test_authentication.py36
-rw-r--r--tests/models/gr/test_base.py46
-rw-r--r--tests/models/gr/test_business.py75
-rw-r--r--tests/models/thl/test_product.py53
-rw-r--r--tests/pytest.ini3
-rw-r--r--tests/test_postgres.py68
75 files changed, 3705 insertions, 3564 deletions
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 <<EOF
- CREATE DATABASE "${env.DB_NAME}" WITH TEMPLATE = template0 ENCODING = 'UTF8';
- EOF
- """
- sh """
- PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres <<EOF
- CREATE DATABASE "${env.GRLIQ_DB_NAME}" WITH TEMPLATE = template0 ENCODING = 'UTF8';
- EOF
- """
- sh """
- PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres <<EOF
- CREATE DATABASE "${env.GR_DB_NAME}" WITH TEMPLATE = template0 ENCODING = 'UTF8';
- EOF
- """
- sh """
- mysql -h ${env.DB_MARIA_HOST} -u ${env.DB_USER} -p${env.DB_PASSWORD} --ssl=0 -e 'CREATE DATABASE `${env.SPECTRUM_DB_NAME}`;'
- """
-
script {
env.REDIS_DB = new Random().nextInt(1024).toString()
env.REDIS = "${env.REDIS}:6379/${env.REDIS_DB}"
diff --git a/generalresearch/config.py b/generalresearch/config.py
index 551f75a..44f3db7 100644
--- a/generalresearch/config.py
+++ b/generalresearch/config.py
@@ -5,9 +5,9 @@ from datetime import datetime, timezone
from pathlib import Path
from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn
-from pydantic_settings import BaseSettings
+from pydantic_settings import BaseSettings, SettingsConfigDict
-from generalresearch.models.custom_types import DaskDsn, SentryDsn
+from generalresearch.models.custom_types import DaskDsn, InternalHostname, SentryDsn
os.environ["DISABLE_PANDERA_IMPORT_WARNING"] = "True"
@@ -39,8 +39,24 @@ def is_debug() -> bool:
class GRLBaseSettings(BaseSettings):
+ model_config = SettingsConfigDict(
+ env_file=(".env.test", ".env.testing", ".env.staging", ".env.prod"),
+ env_file_encoding="utf-8",
+ extra="allow",
+ )
+
debug: bool = Field(default=True)
+ # --- Pytest ---
+
+ testing_postgres: InternalHostname | None = Field(default=None)
+ 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)
redis_timeout: float = Field(default=0.10)
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 8468cdc..d5138cd 100644
--- a/generalresearch/managers/leaderboard/__init__.py
+++ b/generalresearch/managers/leaderboard/__init__.py
@@ -1,12 +1,12 @@
-from typing import Dict
-from zoneinfo import ZoneInfo
+from __future__ import annotations
import pytz
-from cachetools import cached, LRUCache
+from cachetools import LRUCache, cached
+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/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/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/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py
index fb4b20a..0efbf99 100644
--- a/generalresearch/managers/thl/payout.py
+++ b/generalresearch/managers/thl/payout.py
@@ -3,7 +3,8 @@ from __future__ import annotations
from collections import defaultdict
from collections.abc import Collection
from datetime import datetime, timezone
-from random import choice as rand_choice, randint
+from random import choice as rand_choice
+from random import randint
from typing import Any
from uuid import uuid4
@@ -16,8 +17,8 @@ from pydantic import AwareDatetime, NonNegativeInt, PositiveInt
from generalresearch.currency import USDCent
from generalresearch.decorators import LOG
from generalresearch.managers.base import (
- PostgresManagerWithRedis,
Permission,
+ PostgresManagerWithRedis,
)
from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerTransactionConditionFailedError,
@@ -89,15 +90,13 @@ class PayoutEventManager(PostgresManagerWithRedis):
payout_event.update(status=status, ext_ref_id=ext_ref_id, order_data=order_data)
d = payout_event.model_dump_postgres()
- query = sql.SQL(
- """
+ query = sql.SQL("""
UPDATE event_payout SET
status = %(status)s,
ext_ref_id = %(ext_ref_id)s,
order_data = %(order_data)s
WHERE uuid = %(uuid)s;
- """
- )
+ """)
with self.pg_config.make_connection() as conn:
with conn.cursor() as c:
c.execute(query=query, params=d)
@@ -327,53 +326,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
@@ -610,15 +562,13 @@ class BrokerageProductPayoutEventManager(PayoutEventManager):
)
except LedgerTransactionConditionFailedError as e:
if e.args[0] == "duplicate tag":
- raise ValueError(
- f"""Payout event already exists! {e}
+ raise ValueError(f"""Payout event already exists! {e}
You are trying to create a tx that already exists. We can't know
if this is a new payout event with the same ref id, or you're
trying to run the same one twice ... So not setting the existing
payout event to FAILED, b/c the existing one is not failed!
Doing nothing ...
- """
- ) from e
+ """) from e
self.update(payout_event=bp_pe, status=PayoutStatus.FAILED)
raise
except Exception as e:
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/custom_types.py b/generalresearch/models/custom_types.py
index ea96741..84bf8e3 100644
--- a/generalresearch/models/custom_types.py
+++ b/generalresearch/models/custom_types.py
@@ -1,6 +1,10 @@
+from __future__ import annotations
+
import json
+import re
+import sys as _sys
from datetime import datetime, timedelta, timezone
-from typing import Any, Literal, Optional, Set
+from typing import Any, Literal
from uuid import UUID
from pydantic import (
@@ -14,14 +18,31 @@ from pydantic import (
)
from pydantic.functional_serializers import PlainSerializer
from pydantic.functional_validators import AfterValidator, BeforeValidator
-from pydantic.networks import UrlConstraints, IPvAnyNetwork
-from pydantic_core import Url
+from pydantic.networks import IPvAnyNetwork, UrlConstraints
+from pydantic_core import MultiHostHost, Url
from typing_extensions import Annotated
from generalresearch.models import DeviceType, Source
-# if TYPE_CHECKING:
-# from generalresearch.models import DeviceType
+HOSTNAME_REGEX = re.compile(
+ r"^[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?(\.[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?)*$"
+)
+
+
+def validate_hostname(v: str) -> str:
+ if not HOSTNAME_REGEX.match(v):
+ raise ValueError("Invalid internal hostname format")
+ return v
+
+
+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:
@@ -31,7 +52,7 @@ def convert_datetime_to_iso_8601_with_z_suffix(dt: datetime) -> str:
return dt.strftime("%Y-%m-%dT%H:%M:%S.%fZ")
-def convert_str_dt(v: Any) -> Optional[AwareDatetime]:
+def convert_str_dt(v: Any) -> AwareDatetime | None:
# By default, pydantic is unable to handle tz-aware isoformat str. Attempt
# to parse a str that was dumped using the iso8601 format with Z suffix.
if v is not None and type(v) is str:
@@ -158,7 +179,7 @@ from_comma_sep_str = BeforeValidator(
# This is a set of DeviceType, that serializes and de-serializes into a
# (sorted) comma-separated str
-DeviceTypes = Annotated[Set[DeviceType], enum_to_comma_sep_str, from_comma_sep_str]
+DeviceTypes = Annotated[set[DeviceType], enum_to_comma_sep_str, from_comma_sep_str]
# This is a set of alphanumeric strings, that serializes and de-serializes
# into a (sorted) comma-separated str
@@ -223,9 +244,9 @@ InfluxDsn = Annotated[
),
]
-AlphaNumStrSet = Annotated[Set[AlphaNumStr], to_comma_sep_str, from_comma_sep_str]
-IPLikeStrSet = Annotated[Set[IPLikeStr], to_comma_sep_str, from_comma_sep_str]
-UUIDStrSet = Annotated[Set[UUIDStr], to_comma_sep_str, from_comma_sep_str]
+AlphaNumStrSet = Annotated[set[AlphaNumStr], to_comma_sep_str, from_comma_sep_str]
+IPLikeStrSet = Annotated[set[IPLikeStr], to_comma_sep_str, from_comma_sep_str]
+UUIDStrSet = Annotated[set[UUIDStr], to_comma_sep_str, from_comma_sep_str]
list_models_to_json_str = PlainSerializer(
lambda x: json.dumps([y.model_dump(mode="json") for y in x]),
diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py
index a842757..51317da 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/pyproject.toml b/pyproject.toml
index 79b2382..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",
@@ -27,6 +28,7 @@ dependencies = [
"pytest",
"pylibmc",
"pymemcache",
+ "pytz",
"redis",
"requests",
"scipy",
diff --git a/test_utils/conftest.py b/test_utils/conftest.py
index 232c1fc..378b9cc 100644
--- a/test_utils/conftest.py
+++ b/test_utils/conftest.py
@@ -1,34 +1,31 @@
+from __future__ import annotations
+
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
from pathlib import Path
-from typing import TYPE_CHECKING, Callable, Generator
+from typing import Callable, Generator
from uuid import uuid4
-import django
import pytest
-import redis
from _pytest.config import Config
-from django.conf import settings as django_settings
-from django.core.management import call_command
from dotenv import load_dotenv
-from pydantic import MariaDBDsn, PostgresDsn
-from redis import Redis
+from pydantic import MariaDBDsn, PostgresDsn, TypeAdapter
+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) -> str:
+def env_file_path(pytestconfig: Config) -> Path:
root_path = pytestconfig.rootpath
env_file = ".env.test"
@@ -40,18 +37,16 @@ def env_file_path(pytestconfig: Config) -> str:
for env_path in candidates:
if os.path.exists(env_path):
load_dotenv(dotenv_path=env_path, override=True)
- return os.path.normpath(env_path)
+ return Path(os.path.normpath(env_path))
raise AssertionError(f"No .env.test file found in: {', '.join(candidates)}")
@pytest.fixture(scope="session")
-def settings(env_file_path: str) -> "GRLBaseSettings":
+def settings(env_file_path: Path) -> GRLBaseSettings:
from generalresearch.config import GRLBaseSettings
- print(f"{env_file_path=}")
-
- s = GRLBaseSettings(_env_file=env_file_path)
+ s = GRLBaseSettings()
if s.thl_mkpl_rr_db is not None:
if s.spectrum_rw_db is None:
@@ -68,17 +63,31 @@ def settings(env_file_path: str) -> "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 is simplified, and only based off a single host. We don't want to
- create multiple migrated tmp databases for each rw/rr/ro connection
+ This does not create any tables, or schema definitions within the instance.
+ What this does is simply:
+
+ 1. Create a database on a known, consistent, staging or unittest
+ defined Postgres server.
+
+ 2. Return the PostgresDsn of that table
+
+ 3. On shutdown, go ahead and delete that database after the
+ tests have finished.
"""
- assert settings.thl_web_rw_db
- # assert settings.thl_web_rw_db.host
+ msg = "Must define Postgres test settings"
+ assert settings.testing_postgres, msg
+ assert settings.testing_postgres_user, msg
+ assert settings.testing_postgres_pass, msg
- dsn: PostgresDsn = settings.thl_web_rw_db
+ db_uri, db_user, db_pass = (
+ settings.testing_postgres,
+ settings.testing_postgres_user,
+ settings.testing_postgres_pass,
+ )
# Connect to default DB to create the new one
from psycopg import connect
@@ -86,138 +95,167 @@ def postgres_instance(settings: "GRLBaseSettings") -> Generator[PostgresDsn]:
now = datetime.now(timezone.utc)
ts: str = now.strftime("%Y-%m-%d")
-
db_name = f"unittest-{ts}-{uuid4().hex[:6]}"
- print("XXX", str(dsn))
- conn = connect(str(dsn))
+
+ db_path_connect = f"postgres://{db_user}:{db_pass}@{db_uri}"
+ db_path = f"{db_path_connect}/{db_name}"
+
+ # The DATABASE does NOT yet exist on the Postgres SERVER, thus
+ # we first must connect only to the SERVER (eg: default postgres path used)
+ conn = connect(f"{db_path_connect}/postgres")
conn.autocommit = True
cur = conn.cursor()
cur.execute(SQL("CREATE DATABASE {}").format(Identifier(db_name)))
cur.close()
conn.close()
- host = dsn.hosts()[0]
- db_url = (
- f"postgres://{host['username']}:{host['password']}@{host['host']}/{db_name}"
- )
-
- yield PostgresDsn(db_url)
+ yield PostgresDsn(db_path)
# Teardown: drop the DB after the session
- conn = connect(str(dsn))
+ 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(scope="session")
-def django_db_setup(settings: "GRLBaseSettings") -> Callable[..., None]:
+def postgres_instance_dict(
+ postgres_instance: PostgresDsn,
+) -> 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,
+ )
- def _inner():
- assert settings.thl_web_rw_db
- dsn: PostgresDsn = settings.thl_web_rw_db
- host = dsn.hosts()[0]
+@pytest.fixture(scope="session")
+def postgres_instance_host(
+ postgres_instance_dict: PostgresDict,
+) -> Generator[InternalHostname]:
+ adapter = TypeAdapter(InternalHostname)
+ value = adapter.validate_python(postgres_instance_dict["host"])
+ yield value
- # 1. Bootstrapping Django settings
- if not django_settings.configured:
- django_settings.configure(
- DATABASES={
- "default": {
- "ENGINE": "django.db.backends.postgresql",
- # PostgresDsn stores path as "/dbname"
- "NAME": str(dsn.path).lstrip("/"),
- "USER": host["username"],
- "PASSWORD": host["password"],
- "HOST": host["host"],
- "PORT": "5432",
- }
- },
- INSTALLED_APPS=[
- "django.contrib.postgres",
- "django.contrib.contenttypes",
- "generalresearch.thl_django",
- ],
- )
- django.setup()
- from django.apps import apps
+# @pytest.fixture(scope="session")
+# def git_key_path(settings: GRLBaseSettings) -> Path:
+# return Path('/tmp/')
- for model in apps.get_models():
- print(f"Discovered model: {model._meta.label}")
- # 2. Run migrations directly during fixture activation
- call_command("migrate")
+@pytest.fixture(scope="session")
+def git_key_path(
+ settings: GRLBaseSettings,
+) -> Generator[Path]:
- return _inner
+ 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)
-@pytest.fixture(scope="session")
-def thl_web_rr(
- settings: "GRLBaseSettings", postgres_instance: PostgresDsn, django_db_setup
-) -> PostgresConfig:
- dsn = settings.thl_web_rr_db
- assert dsn
- assert dsn.path
+ yield Path(key_path)
+
+ os.unlink(key_path)
- if dsn.path not in ["/", "/postgres"]:
- assert "/unittest-" in dsn.path
- db_path = postgres_instance.path
- host = dsn.hosts()[0]
- db_url = f"postgres://{host['username']}:{host['password']}@{host['host']}{db_path}"
+@pytest.fixture(scope="session")
+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")
+
+ 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,
+ )
- # Run Migrations now.
- django_db_setup()
+ return repo_path
- return PostgresConfig(
- dsn=PostgresDsn(db_url),
- connect_timeout=1,
- statement_timeout=5,
- )
+ return _inner
@pytest.fixture(scope="session")
-def thl_web_rw(
- settings: "GRLBaseSettings", postgres_instance: PostgresDsn, django_db_setup
-) -> PostgresConfig:
- dsn = settings.thl_web_rw_db
- assert dsn
- assert dsn.path
+def django_db_factory(
+ postgres_instance: PostgresDsn,
+ postgres_instance_dict: PostgresDict,
+ gr_repo: Callable[..., Path],
+) -> Callable[..., PostgresDsn]:
- if dsn.path not in ["/", "/postgres"]:
- assert "/unittest-" in dsn.path
+ import django
+ from django.conf import settings as django_settings
+ from django.core.management import call_command
- db_path = postgres_instance.path
- host = dsn.hosts()[0]
- db_url = f"postgres://{host['username']}:{host['password']}@{host['host']}{db_path}"
+ def _inner(django_project: str = "generalresearch.thl_django"):
- # Run Migrations now.
- django_db_setup()
+ 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))
- return PostgresConfig(
- dsn=PostgresDsn(db_url),
- connect_timeout=1,
- statement_timeout=5,
- )
+ print(sys.path)
+ # 1. Bootstrapping Django settings
+ if not django_settings.configured:
+ django_settings.configure(
+ DATABASES={
+ "default": {
+ "ENGINE": "django.db.backends.postgresql",
+ "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",
+ django_project,
+ ],
+ )
+ django.setup()
-@pytest.fixture(scope="session")
-def gr_db(settings: "GRLBaseSettings") -> PostgresConfig:
- dsn = settings.gr_db
- assert dsn
- assert dsn.path
+ # for model in apps.get_models():
+ # print(f"Discovered model: {model._meta.label}")
- if dsn.path not in ["/", "/postgres"]:
- assert "/unittest-" in dsn.path
+ # 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 PostgresConfig(dsn=settings.gr_db, connect_timeout=5, statement_timeout=2)
+ return _inner
@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
@@ -233,114 +271,15 @@ def spectrum_rw(settings: "GRLBaseSettings") -> SqlHelper:
)
-@pytest.fixture(scope="session")
-def grliq_db(settings: "GRLBaseSettings") -> PostgresConfig:
- dsn = settings.grliq_db
- assert dsn
- assert dsn.path
-
- if dsn.path not in ["/", "/postgres"]:
- assert "/unittest-" in dsn.path
-
- # 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=settings.grliq_db,
- connect_timeout=2,
- statement_timeout=2,
- )
-
-
-@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)
@@ -351,30 +290,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)
@@ -424,6 +355,8 @@ def delete_df_collection(
)
case _:
+ assert coll.data_type
+
thl_web_rw.execute_write(
query=f"DELETE FROM {coll.data_type.value};",
)
@@ -435,23 +368,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/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 058093e..12e57c5 100644
--- a/test_utils/incite/conftest.py
+++ b/test_utils/incite/conftest.py
@@ -1,18 +1,17 @@
+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
from _pytest.fixtures import SubRequest
from faker import Faker
-# from test_utils.managers.ledger.conftest import session_with_tx_factory
-# from test_utils.models.conftest import session_factory
-
if TYPE_CHECKING:
from generalresearch.config import GRLBaseSettings
from generalresearch.incite.base import GRLDatasets
@@ -33,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)
@@ -56,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,
@@ -72,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,
@@ -88,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.
@@ -114,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
@@ -125,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
@@ -145,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,
@@ -204,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
@@ -518,38 +116,6 @@ def geoipinfo_manager(
@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
assert "/unittest-" in 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/grliq/managers/__init__.py b/test_utils/managers/gr/__init__.py
index e69de29..e69de29 100644
--- a/test_utils/grliq/managers/__init__.py
+++ b/test_utils/managers/gr/__init__.py
diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py
new file mode 100644
index 0000000..37da164
--- /dev/null
+++ b/test_utils/managers/gr/conftest.py
@@ -0,0 +1,110 @@
+from __future__ import annotations
+
+from typing import Callable
+
+import pytest
+import redis.asyncio as redis_async
+from pydantic import PostgresDsn
+from redis import Redis
+
+from generalresearch.config import GRLBaseSettings
+from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager
+from generalresearch.managers.gr.business import (
+ BusinessAddressManager,
+ BusinessBankAccountManager,
+ BusinessManager,
+)
+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_redis_async(settings: GRLBaseSettings) -> redis_async.Redis:
+ assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis)
+
+ return redis_async.Redis.from_url(
+ str(settings.gr_redis),
+ decode_responses=True,
+ socket_timeout=0.20,
+ socket_connect_timeout=0.20,
+ )
+
+
+@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_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig:
+
+ return PostgresConfig(
+ dsn=django_db_factory("gr_carer"),
+ connect_timeout=1,
+ statement_timeout=5,
+ )
+
+
+# === Managers ===
+
+
+@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
+
+ from generalresearch.managers.gr.authentication import GRUserManager
+
+ return GRUserManager(pg_config=gr_db, redis_config=gr_redis_config)
+
+
+@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)
+
+
+@pytest.fixture(scope="session")
+def gr_business_bank_account_manager(
+ gr_db: PostgresConfig,
+) -> BusinessBankAccountManager:
+ return BusinessBankAccountManager(pg_config=gr_db)
+
+
+@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/ledger/conftest.py b/test_utils/managers/ledger/conftest.py
index 105085d..ce8348e 100644
--- a/test_utils/managers/ledger/conftest.py
+++ b/test_utils/managers/ledger/conftest.py
@@ -1,739 +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,
- ext_ref_id=uuid4().hex
+ 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/grliq/models/__init__.py b/test_utils/managers/thl/__init__.py
index e69de29..e69de29 100644
--- a/test_utils/grliq/models/__init__.py
+++ b/test_utils/managers/thl/__init__.py
diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py
new file mode 100644
index 0000000..5b70961
--- /dev/null
+++ b/test_utils/managers/thl/conftest.py
@@ -0,0 +1,258 @@
+from __future__ import annotations
+
+from typing import Callable
+
+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/models/conftest.py b/test_utils/models/conftest.py
index 6a8e4cf..89f6f32 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
@@ -13,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
@@ -56,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
@@ -67,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:
@@ -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,25 +172,24 @@ def wall(
@pytest.fixture
def session_factory(
- wall_factory: Callable[..., "Wall"],
- session_manager: "SessionManager",
- wall_manager: "WallManager",
+ 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,
started: datetime = utc_hour_ago,
- ) -> "Session":
+ ) -> Session:
if wall_req_cpis:
assert len(wall_req_cpis) == wall_count
if wall_statuses:
@@ -236,24 +237,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 +282,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 +297,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 +309,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 +326,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 +348,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 +358,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 +368,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,17 +380,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,
- ) -> "BrokerageProductPayoutEvent":
+ 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:
from generalresearch.currency import USDCent
product = product or product_manager.create_dummy()
@@ -411,120 +413,49 @@ 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":
- 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":
+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,
@@ -537,14 +468,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
)
@@ -552,10 +483,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,
@@ -568,25 +499,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(
@@ -597,7 +528,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(
@@ -608,7 +539,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
@@ -616,10 +547,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/test_utils/grliq/managers/conftest.py b/test_utils/models/contest/__init__.py
index e69de29..e69de29 100644
--- a/test_utils/grliq/managers/conftest.py
+++ b/test_utils/models/contest/__init__.py
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/grliq/models/conftest.py b/test_utils/models/gr/__init__.py
index e69de29..e69de29 100644
--- a/test_utils/grliq/models/conftest.py
+++ b/test_utils/models/gr/__init__.py
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
--- /dev/null
+++ b/test_utils/models/ledger/__init__.py
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
--- /dev/null
+++ b/test_utils/models/network/__init__.py
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
--- /dev/null
+++ b/test_utils/models/thl/__init__.py
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
--- /dev/null
+++ b/test_utils/models/upk/__init__.py
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/managers/upk/marketplace_category.csv.gz b/test_utils/models/upk/marketplace_category.csv.gz
index 0f8ec1c..0f8ec1c 100644
--- a/test_utils/managers/upk/marketplace_category.csv.gz
+++ b/test_utils/models/upk/marketplace_category.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_item.csv.gz b/test_utils/models/upk/marketplace_item.csv.gz
index c12c5d8..c12c5d8 100644
--- a/test_utils/managers/upk/marketplace_item.csv.gz
+++ b/test_utils/models/upk/marketplace_item.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_property.csv.gz b/test_utils/models/upk/marketplace_property.csv.gz
index a781d1d..a781d1d 100644
--- a/test_utils/managers/upk/marketplace_property.csv.gz
+++ b/test_utils/models/upk/marketplace_property.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz
index 5b4ea19..5b4ea19 100644
--- a/test_utils/managers/upk/marketplace_propertycategoryassociation.csv.gz
+++ b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_propertycountry.csv.gz b/test_utils/models/upk/marketplace_propertycountry.csv.gz
index 5d2a637..5d2a637 100644
--- a/test_utils/managers/upk/marketplace_propertycountry.csv.gz
+++ b/test_utils/models/upk/marketplace_propertycountry.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_propertyitemrange.csv.gz b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz
index 84f4f0e..84f4f0e 100644
--- a/test_utils/managers/upk/marketplace_propertyitemrange.csv.gz
+++ b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz
index 6b9fd1c..6b9fd1c 100644
--- a/test_utils/managers/upk/marketplace_propertymarketplaceassociation.csv.gz
+++ b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz
Binary files differ
diff --git a/test_utils/managers/upk/marketplace_question.csv.gz b/test_utils/models/upk/marketplace_question.csv.gz
index bcfc3ad..bcfc3ad 100644
--- a/test_utils/managers/upk/marketplace_question.csv.gz
+++ b/test_utils/models/upk/marketplace_question.csv.gz
Binary files 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/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/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/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_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py
index a23413c..530142e 100644
--- a/tests/models/custom_types/test_aware_datetime.py
+++ b/tests/models/custom_types/test_aware_datetime.py
@@ -1,10 +1,11 @@
+from __future__ import annotations
+
import logging
from datetime import datetime, timezone
-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
@@ -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/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..a9f01a8
--- /dev/null
+++ b/tests/models/gr/test_base.py
@@ -0,0 +1,46 @@
+import subprocess
+from pathlib import Path
+from typing import Callable
+
+import pytest
+from pydantic import PostgresDsn
+
+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")
+ 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,
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
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..3b3ddd0
--- /dev/null
+++ b/tests/test_postgres.py
@@ -0,0 +1,68 @@
+import socket
+import subprocess
+from typing import Callable
+
+from pydantic import PostgresDsn
+
+from generalresearch.models.custom_types import InternalHostname, PostgresDict
+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
+
+
+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