diff options
| -rw-r--r-- | generalresearch/managers/thl/userhealth.py | 65 | ||||
| -rw-r--r-- | generalresearch/models/thl/ipinfo.py | 10 | ||||
| -rw-r--r-- | test_utils/managers/conftest.py | 107 | ||||
| -rw-r--r-- | test_utils/models/thl/conftest.py | 174 | ||||
| -rw-r--r-- | tests/managers/thl/test_ipinfo.py | 201 | ||||
| -rw-r--r-- | tests/managers/thl/test_userhealth.py | 189 |
6 files changed, 246 insertions, 500 deletions
diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 0df67aa..1aeab44 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -33,6 +33,10 @@ if TYPE_CHECKING: class UserIpHistoryManager(PostgresManagerWithRedis): + def __init__(self, geoip_info_manager: GeoIpInfoManager | None, **kwargs) -> None: + super().__init__(**kwargs) + self.geoip_info_manager = geoip_info_manager + def get_redis_key(self, user_id: int) -> str: return f"generalreserach:user-ip-history-v2:{user_id}" @@ -74,9 +78,8 @@ class UserIpHistoryManager(PostgresManagerWithRedis): iph = UserIPHistory(user=user, ips=records) self.set_user_ip_history_cache(user_id=user.user_id, iph=iph) - def get_user_ip_history( - self, user: UserRef | User, geoip_info_manager: GeoIpInfoManager | None = None - ) -> UserIPHistory: + def get_user_ip_history(self, user: UserRef | User) -> UserIPHistory: + assert self.geoip_info_manager is not None, "GeoIpInfoManager is required" user = user if isinstance(user, UserRef) else user.to_user_ref() iph = self.get_user_ip_history_cache(user_id=user.user_id) @@ -88,19 +91,15 @@ class UserIpHistoryManager(PostgresManagerWithRedis): # todo: we may get dns records from somewhere else here ... iph = UserIPHistory(user=user, ips=records) self.set_user_ip_history_cache(user_id=user.user_id, iph=iph) - if geoip_info_manager: - iph.enrich_ips(geoip_info_manager=geoip_info_manager) + iph.enrich_ips(geoip_info_manager=self.geoip_info_manager) return iph def get_user_latest_ip_record( self, user: UserRef | User, exclude_anon: bool = False, - geoip_info_manager: GeoIpInfoManager | None = None, ) -> UserIPRecord | None: - iphistory = self.get_user_ip_history( - user=user, geoip_info_manager=geoip_info_manager - ) + iphistory = self.get_user_ip_history(user=user) if not iphistory or not iphistory.ips: return None @@ -108,7 +107,6 @@ class UserIpHistoryManager(PostgresManagerWithRedis): # This logic was changed at some point? Or the py-utils # get_user_latest_ip_record was changed. Please be careful here... if exclude_anon: - assert geoip_info_manager is not None, "Must pass geoip_info_manager" for ipr in iphistory.ips: if not ipr.information.is_anonymous: return ipr @@ -120,11 +118,8 @@ class UserIpHistoryManager(PostgresManagerWithRedis): self, user: UserRef | User, exclude_anon: bool = False, - geoip_info_manager: GeoIpInfoManager | None = None, ) -> str | None: - record = self.get_user_latest_ip_record( - user=user, exclude_anon=exclude_anon, geoip_info_manager=geoip_info_manager - ) + record = self.get_user_latest_ip_record(user=user, exclude_anon=exclude_anon) if record: return record.ip return None @@ -132,25 +127,18 @@ class UserIpHistoryManager(PostgresManagerWithRedis): def get_user_latest_country( self, user: UserRef | User, - geoip_info_manager: GeoIpInfoManager, exclude_anon: bool = False, ) -> str | None: """Get the country the user is in, based off their latest ip.""" - ipr = self.get_user_latest_ip_record( - user, geoip_info_manager=geoip_info_manager, exclude_anon=exclude_anon - ) + ipr = self.get_user_latest_ip_record(user, exclude_anon=exclude_anon) # The ipr.information should exist, but it is possible the user has # no IP history at all, so the record is None return ipr.country_iso if ipr is not None else None - def is_user_anonymous( - self, user: UserRef | User, geoip_info_manager: GeoIpInfoManager - ) -> bool | None: + def is_user_anonymous(self, user: UserRef | User) -> bool | None: # Get the user's latest ip. is it marked as anonymous? # Return None if the user has no IP history at all - ipr = self.get_user_latest_ip_record( - user, geoip_info_manager=geoip_info_manager - ) + ipr = self.get_user_latest_ip_record(user) if ipr: return ipr.is_anonymous if ipr.is_anonymous is not None else False return None @@ -175,11 +163,12 @@ class IPRecordManager(PostgresManagerWithRedis): redis_config=self.redis_config, cache_prefix=self.cache_prefix, permissions=self.permissions, + geoip_info_manager=None, ) def create_unpack( self, - user: UserRef, + user_id: int, ip: IPvAnyAddressStr, forwarded_ips: list[str], ) -> IPRecord: @@ -188,22 +177,22 @@ class IPRecordManager(PostgresManagerWithRedis): padded = list(forwarded_ips) + [None] * (6 - len(forwarded_ips)) - return self.create(user, ip, *padded) + return self.create(user_id, ip, *padded) def create( self, - user: UserRef, + user_id: int, ip: IPvAnyAddressStr, - forwarded_ip1: IPvAnyAddressStr, - forwarded_ip2: IPvAnyAddressStr, - forwarded_ip3: IPvAnyAddressStr, - forwarded_ip4: IPvAnyAddressStr, - forwarded_ip5: IPvAnyAddressStr, - forwarded_ip6: IPvAnyAddressStr, + forwarded_ip1: IPvAnyAddressStr | None, + forwarded_ip2: IPvAnyAddressStr | None, + forwarded_ip3: IPvAnyAddressStr | None, + forwarded_ip4: IPvAnyAddressStr | None, + forwarded_ip5: IPvAnyAddressStr | None, + forwarded_ip6: IPvAnyAddressStr | None, ) -> IPRecord: data = { - "user_id": user.user_id, + "user_id": user_id, "ip": ipaddress.ip_address(ip).exploded, "created": datetime.now(tz=UTC), } @@ -245,7 +234,7 @@ class IPRecordManager(PostgresManagerWithRedis): """, params=data, ) - self.recreate_user_ip_history_cache(user=user) + self.delete_user_ip_history_cache(user_id=user_id) return IPRecord.from_mysql(data) @@ -310,8 +299,10 @@ class IPRecordManager(PostgresManagerWithRedis): return [IPRecord.from_mysql(i) for i in res] - def recreate_user_ip_history_cache(self, user: UserRef): - return self.user_ip_history_manager.recreate_user_ip_history_cache(user=user) + def delete_user_ip_history_cache(self, user_id: int): + return self.user_ip_history_manager.delete_user_ip_history_cache( + user_id=user_id + ) class AuditLogManager(PostgresManager): diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 9fe86e4..0dff6b4 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -10,9 +10,9 @@ from pydantic import ( BaseModel, ConfigDict, Field, + IPvAnyAddress, PositiveInt, field_validator, - IPvAnyAddress, ) from generalresearch.models.custom_types import ( @@ -122,11 +122,13 @@ class IPGeoname(BaseModel): return cls.model_validate(d) -class IPInformation(BaseModel): +class GeoIPInformation(BaseModel): """ Fields we'll always pull from GRIP's mmdb files at minimum """ + model_config = ConfigDict(extra="ignore") + ip: IPvAnyAddressStr = Field() country_iso: CountryISOLike | None = Field( @@ -162,7 +164,3 @@ class IPInformation(BaseModel): "(e.g., 'residential', 'business').", examples=[AccessType.RESIDENTIAL], ) - - -class GeoIPInformation(IPInformation): - model_config = ConfigDict(extra="ignore") diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 391e6bf..72be79e 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -1,19 +1,33 @@ from __future__ import annotations import random +from collections import defaultdict from collections.abc import Callable from datetime import datetime from typing import TYPE_CHECKING +from unittest.mock import Mock from uuid import uuid4 import pytest +from grip_client.mmdb.models import ( + GRIPAnonymousRecord, + GRIPAsnRecord, + GRIPCountryRecord, + GRIPMMDBLookupResult, +) from generalresearch.managers.thl.cashout_method import ( CashoutMethodManager, ) +from generalresearch.managers.thl.ipinfo import GeoIpInfoManager from generalresearch.managers.thl.user_streak import ( UserStreakManager, ) +from generalresearch.managers.thl.userhealth import ( + AuditLogManager, + IPRecordManager, + UserIpHistoryManager, +) from generalresearch.models.definitions import Source from generalresearch.models.thl.wallet.cashout_method import ( CashoutMethod, @@ -24,15 +38,6 @@ from generalresearch.models.thl.wallet.definitions import Currency, PayoutType if TYPE_CHECKING: from generalresearch.managers.spectrum.survey import SpectrumSurveyManager from generalresearch.managers.thl.buyer import BuyerManager - from generalresearch.managers.thl.ipinfo import ( - GeoIpInfoManager, - IPGeonameManager, - ) - from generalresearch.managers.thl.userhealth import ( - AuditLogManager, - IPRecordManager, - UserIpHistoryManager, - ) from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig @@ -51,41 +56,39 @@ def audit_log_manager(thl_web_rw: PostgresConfig) -> AuditLogManager: return AuditLogManager(pg_config=thl_web_rw) -@pytest.fixture(scope="session") -def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ipinfo import IPGeonameManager - - return IPGeonameManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") +@pytest.fixture def ip_record_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig + thl_web_rw: PostgresConfig, + thl_redis_config: RedisConfig, + geoip_info_manager: GeoIpInfoManager, ) -> IPRecordManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path from generalresearch.managers.thl.userhealth import IPRecordManager - return IPRecordManager(pg_config=thl_web_rw, redis_config=thl_redis_config) + return IPRecordManager( + pg_config=thl_web_rw, + redis_config=thl_redis_config, + geoip_info_manager=geoip_info_manager, + ) -@pytest.fixture(scope="session") +@pytest.fixture def user_iphistory_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig + thl_web_rw: PostgresConfig, + thl_redis_config: RedisConfig, + geoip_info_manager: GeoIpInfoManager, ) -> UserIpHistoryManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.userhealth import ( - UserIpHistoryManager, + return UserIpHistoryManager( + pg_config=thl_web_rw, + redis_config=thl_redis_config, + geoip_info_manager=geoip_info_manager, ) - return UserIpHistoryManager(pg_config=thl_web_rw, redis_config=thl_redis_config) - @pytest.fixture(scope="function") def user_iphistory_manager_clear_cache(user_iphistory_manager, user: User): @@ -96,16 +99,42 @@ def user_iphistory_manager_clear_cache(user_iphistory_manager, user: User): user_iphistory_manager.delete_user_ip_history_cache(user_id=user.user_id) -@pytest.fixture(scope="session") -def geoipinfo_manager( - thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig -) -> GeoIpInfoManager: - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path +@pytest.fixture +def grip_lookup_results() -> dict[str, GRIPMMDBLookupResult]: + return defaultdict( + lambda: GRIPMMDBLookupResult( + country=GRIPCountryRecord(), + anonymous=GRIPAnonymousRecord(), + asn=GRIPAsnRecord(), + ), + { + "8.8.8.8": GRIPMMDBLookupResult( + country=GRIPCountryRecord(country_iso="US"), + anonymous=GRIPAnonymousRecord(is_anonymous=False), + asn=GRIPAsnRecord( + asn=15169, + network_operator="Google", + ), + ), + "1.1.1.1": GRIPMMDBLookupResult( + country=GRIPCountryRecord(country_iso="AU"), + anonymous=GRIPAnonymousRecord(is_anonymous=True), + asn=GRIPAsnRecord( + asn=13335, + network_operator="Cloudflare", + ), + ), + }, + ) - from generalresearch.managers.thl.ipinfo import GeoIpInfoManager - return GeoIpInfoManager(pg_config=thl_web_rw, redis_config=thl_redis_config) +@pytest.fixture(scope="function") +def geoip_info_manager( + grip_lookup_results: dict[str, GRIPMMDBLookupResult], +) -> GeoIpInfoManager: + manager = GeoIpInfoManager(grip_token="test-token") + manager.grip_mmdb.lookup = Mock(side_effect=grip_lookup_results.__getitem__) + return manager @pytest.fixture(scope="session") @@ -193,10 +222,10 @@ def example_tango_cashout_methods( last_updated=datetime.fromisoformat("2021-06-23T20:45:38.239182Z"), is_live=True, type=PayoutType.TANGO, - ext_id='U025035', + ext_id="U025035", name="Safeway eGift Card $25", data=TangoCashoutMethodData( - value_type="fixed", countries=["US"], utid='U025035' + value_type="fixed", countries=["US"], utid="U025035" ), user=None, image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b694446-1200w-326ppi.png", @@ -209,7 +238,7 @@ def example_tango_cashout_methods( last_updated=datetime.fromisoformat("2021-06-23T20:45:38.239182Z"), is_live=True, type=PayoutType.TANGO, - ext_id='U006961', + ext_id="U006961", name="Amazon.it Gift Certificate", data=TangoCashoutMethodData( value_type="variable", countries=["IT"], utid="U006961" diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index d20fdc4..433003c 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -14,10 +14,9 @@ from grip_client.enums import AccessType from pydantic import PositiveInt from generalresearch.currency import USDCent -from generalresearch.managers.thl.payout import ( - BusinessPayoutEventManager, - UserPayoutEventManager, -) +from generalresearch.managers.thl.ipinfo import GeoIpInfoManager +from generalresearch.managers.thl.payout import BusinessPayoutEventManager +from generalresearch.managers.thl.wallet.user_payout import UserPayoutEventManager from generalresearch.models.custom_types import ( AwareDatetimeISO, IPvAnyAddressStr, @@ -29,6 +28,7 @@ from generalresearch.models.thl.definitions import ( PayoutStatus, Status, ) +from generalresearch.models.thl.ipinfo import GeoIPInformation from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.user import User from generalresearch.models.thl.userhealth import AuditLogLevel @@ -36,10 +36,6 @@ from generalresearch.models.thl.wallet.definitions import PayoutType from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: - from generalresearch.managers.thl.ipinfo import ( - IPGeonameManager, - IPInformationManager, - ) from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.session import SessionManager @@ -50,7 +46,6 @@ if TYPE_CHECKING: from generalresearch.models.gr.business import Business from generalresearch.models.gr.team import Team from generalresearch.models.legacy.bucket import Bucket - from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation from generalresearch.models.thl.payout import BrokerageProductPayoutEvent from generalresearch.models.thl.product import ( PayoutConfig, @@ -438,143 +433,44 @@ def unsaved_product(product_factory: Callable[..., Product]) -> Product: return product_factory(save=False) -# --- IP Geoname --- - - -@pytest.fixture -def ip_geoname_factory( - ip_geoname_manager: IPGeonameManager, -) -> Callable[..., IPGeoname]: - - def _inner( - save: bool = True, - 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: - if save: - return ip_geoname_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, - ) - else: - raise ValueError("Unsaved IPGeoname not yet supported") - - return _inner - - -@pytest.fixture() -def ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: - return ip_geoname_factory(save=True) - - -@pytest.fixture() -def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: - return ip_geoname_factory(save=True) - - -# --- IP Information --- +# --- GeoIP Information --- @pytest.fixture -def ip_information_factory( - ip_information_manager: IPInformationManager, -) -> Callable[..., IPInformation]: +def geoip_information_factory( + geoip_info_manager: GeoIpInfoManager, +) -> Callable[..., GeoIPInformation]: def _inner( - save: bool = True, ip: IPvAnyAddressStr | None = None, - geoname_id: PositiveInt | None = None, country_iso: str | None = None, - registered_country_iso: str | None = None, + country_name: 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: AccessType | None = None, - postal_code: str | None = None, - latitude: Decimal | None = None, - longitude: Decimal | None = None, - accuracy_radius: int | None = None, - ) -> IPInformation: - - if save: - return ip_information_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, - ) - else: - raise ValueError("Unsaved IP Information not supported yet") + access_type: AccessType | None = None, + ) -> GeoIPInformation: + + return GeoIPInformation( + country_iso=country_iso or fake.country_code("alpha-2"), + access_type=access_type, + is_anonymous=is_anonymous, + autonomous_system_number=autonomous_system_number, + autonomous_system_organization=autonomous_system_organization, + country_name=country_name, + ip=ip or fake.ipv4_public(), + subdivision_1_iso=None, + subdivision_1_name=None, + ) return _inner @pytest.fixture -def ip_information( - ip_information_factory: Callable[..., IPInformation], -) -> IPInformation: - return ip_information_factory(save=True) - - -@pytest.fixture() -def unsaved_ip_information( - ip_information_factory: Callable[..., IPInformation], -) -> IPInformation: - return ip_information_factory(save=False) +def geoip_information( + geoip_information_factory: Callable[..., GeoIPInformation], +) -> GeoIPInformation: + return geoip_information_factory() # --- IP Record --- @@ -584,7 +480,7 @@ def unsaved_ip_information( def ip_record_factory(ip_record_manager: IPRecordManager) -> Callable[..., IPRecord]: def _inner( - user_id: PositiveInt, + user: User, save: bool = True, ip: IPvAnyAddressStr | None = None, forwarded_ip1: IPvAnyAddressStr | None = None, @@ -597,15 +493,13 @@ def ip_record_factory(ip_record_manager: IPRecordManager) -> Callable[..., IPRec if save: return ip_record_manager.create( - user_id=user_id, + user_id=user.to_user_ref().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_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, diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index c021eb9..f39d50e 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -1,160 +1,47 @@ -from collections.abc import Callable -from typing import TYPE_CHECKING - -import faker - -from generalresearch.managers.thl.ipinfo import ( - GeoIpInfoManager, - IPGeonameManager, - IPInformationManager, -) -from generalresearch.models.thl.ipinfo import ( - GeoIPInformation, - IPGeoname, - IPInformation, -) - -if TYPE_CHECKING: - from generalresearch.pg_helper import PostgresConfig - from generalresearch.redis_helper import RedisConfig - -fake = faker.Faker() - - -class TestIPGeonameManager: - - def test_init( - self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager - ): - - instance = IPGeonameManager(pg_config=thl_web_rr) - assert isinstance(instance, IPGeonameManager) - assert isinstance(ip_geoname_manager, IPGeonameManager) - - def test_create( - self, - ip_geoname_factory: Callable[..., IPGeoname], - ip_geoname_manager: IPGeonameManager, - ): - - instance = ip_geoname_factory() - assert isinstance(instance, IPGeoname) - - res = ip_geoname_manager.fetch_geoname_ids(filter_ids=[instance.geoname_id]) - assert res[0].model_dump_json() == instance.model_dump_json() - - -class TestIPInformationManager: - - def test_init( - self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager - ): - instance = IPInformationManager(pg_config=thl_web_rr) - assert isinstance(instance, IPInformationManager) - assert isinstance(ip_information_manager, IPInformationManager) - - def test_create( - self, - ip_information_factory: Callable[..., IPInformation], - ip_information_manager: IPInformationManager, - ): - instance = ip_information_factory() - assert isinstance(instance, IPInformation) - - res = ip_information_manager.fetch_ip_information(filter_ips=[instance.ip]) - assert res[0].model_dump_json() == instance.model_dump_json() - - def test_prefetch_geoname( - self, - ip_information: IPInformation, - ip_geoname: IPGeoname, - thl_web_rr: PostgresConfig, - ): - assert isinstance(ip_information, IPInformation) - - assert ip_information.geoname_id == ip_geoname.geoname_id - assert ip_information.geoname is None - - ip_information.prefetch_geoname(pg_config=thl_web_rr) - assert isinstance(ip_information.geoname, IPGeoname) +from generalresearch.managers.thl.ipinfo import GeoIpInfoManager +from generalresearch.models.thl.ipinfo import GeoIPInformation class TestGeoIpInfoManager: - def test_init( - self, - thl_web_rr: PostgresConfig, - thl_redis_config: RedisConfig, - geoipinfo_manager: GeoIpInfoManager, - ): - instance = GeoIpInfoManager(pg_config=thl_web_rr, redis_config=thl_redis_config) - assert isinstance(instance, GeoIpInfoManager) - assert isinstance(geoipinfo_manager, GeoIpInfoManager) - - def test_multi( - self, - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, - geoipinfo_manager: GeoIpInfoManager, - ): - ip = fake.ipv4_public() - ip_information_factory(ip=ip, geoname=ip_geoname) - ips = [ip] - - # This only looks up in redis. They don't exist yet - res = geoipinfo_manager.get_cache_multi(ip_addresses=ips) - assert res == {ip: None} - - # Looks up in redis, if not exists, looks in mysql, then sets - # the caches that didn't exist. - res = geoipinfo_manager.get_multi(ip_addresses=ips) - assert res[ip] is not None - - ip2 = fake.ipv4_public() - ip_information_factory(ip=ip2, geoname=ip_geoname) - ips = [ip, ip2] - res = geoipinfo_manager.get_cache_multi(ip_addresses=ips) - assert res[ip] is not None - assert res[ip2] is None - res = geoipinfo_manager.get_multi(ip_addresses=ips) - assert res[ip] is not None - assert res[ip2] is not None - res = geoipinfo_manager.get_cache_multi(ip_addresses=ips) - assert res[ip] is not None - assert res[ip2] is not None - - def test_multi_ipv6( - self, - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, - geoipinfo_manager: GeoIpInfoManager, - ): - ip = fake.ipv6() - # Make another IP that will be in the same /64 block. - ip2 = ip[:-1] + "a" if ip[-1] != "a" else ip[:-1] + "b" - ip_information_factory(ip=ip, geoname=ip_geoname) - ips = [ip, ip2] - print(f"{ips=}") - - # This only looks up in redis. They don't exist yet - res = geoipinfo_manager.get_cache_multi(ip_addresses=ips) - assert res == {ip: None, ip2: None} - - # Looks up in redis, if not exists, looks in mysql, then sets - # the caches that didn't exist. - res = geoipinfo_manager.get_multi(ip_addresses=ips) - - res1 = res[ip] - assert isinstance(res1, GeoIPInformation) - assert res1.ip == ip - assert res1.lookup_prefix == "/64" - - res2 = res[ip2] - assert isinstance(res2, GeoIPInformation) - assert res2.ip == ip2 - assert res2.lookup_prefix == "/64" - # they should be the same basically, except for the ip - - def test_doesnt_exist(self, geoipinfo_manager: GeoIpInfoManager): - ip = fake.ipv4_public() - res = geoipinfo_manager.get_multi(ip_addresses=[ip]) - assert res == {ip: None} + def test_get(self, geoipinfo_manager: GeoIpInfoManager): + result = geoip_info_manager.get("8.8.8.8") + + assert result == GeoIPInformation( + ip="8.8.8.8", + country_iso="us", + is_anonymous=False, + autonomous_system_number=15169, + autonomous_system_organization="Google", + access_type=None, + ) + geoip_info_manager.grip_mmdb.lookup.assert_called_once_with("8.8.8.8") + + def test_get_multi(self, geoipinfo_manager: GeoIpInfoManager): + result = geoip_info_manager.get_multi(["8.8.8.8", "1.1.1.1", "8.8.8.8"]) + + assert result == { + "8.8.8.8": GeoIPInformation( + ip="8.8.8.8", + country_iso="us", + is_anonymous=False, + autonomous_system_number=15169, + autonomous_system_organization="Google", + access_type=None, + ), + "1.1.1.1": GeoIPInformation( + ip="1.1.1.1", + country_iso="au", + is_anonymous=True, + autonomous_system_number=13335, + autonomous_system_organization="Cloudflare", + access_type=None, + ), + } + assert geoip_info_manager.grip_mmdb.lookup.call_count == 2 + assert { + call.args[0] for call in geoip_info_manager.grip_mmdb.lookup.call_args_list + } == {"8.8.8.8", "1.1.1.1"} + + def test_get_multi_empty(self, geoipinfo_manager: GeoIpInfoManager): + assert geoip_info_manager.get_multi([]) == {} + geoip_info_manager.grip_mmdb.lookup.assert_not_called() diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index 268b110..8256185 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -3,11 +3,13 @@ from __future__ import annotations from collections.abc import Callable from datetime import UTC, datetime from typing import TYPE_CHECKING +from unittest.mock import Mock from uuid import uuid4 import faker import pytest +from generalresearch.managers.thl.ipinfo import GeoIpInfoManager from generalresearch.managers.thl.userhealth import ( AuditLogManager, IPRecordManager, @@ -23,10 +25,6 @@ from generalresearch.models.thl.user_iphistory import ( from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel if TYPE_CHECKING: - from generalresearch.models.thl.ipinfo import ( - IPGeoname, - IPInformation, - ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig @@ -36,7 +34,6 @@ fake = faker.Faker() class TestAuditLog: - def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager: AuditLogManager): alm = AuditLogManager(pg_config=thl_web_rr) @@ -56,7 +53,7 @@ class TestAuditLog: user_id=user.user_id, level=level, event_type=uuid4().hex ) assert isinstance(instance, AuditLog) - assert instance.id != 1 + assert instance.id != 0 def test_get_by_id(self, audit_log: AuditLog, audit_log_manager: AuditLogManager): @@ -225,14 +222,18 @@ class TestAuditLog: class TestIPRecordManager: - def test_init( self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, + geoip_info_manager: GeoIpInfoManager, ip_record_manager: IPRecordManager, ): - instance = IPRecordManager(pg_config=thl_web_rr, redis_config=thl_redis_config) + instance = IPRecordManager( + pg_config=thl_web_rr, + redis_config=thl_redis_config, + geoip_info_manager=geoip_info_manager, + ) assert isinstance(instance, IPRecordManager) assert isinstance(ip_record_manager, IPRecordManager) @@ -240,10 +241,11 @@ class TestIPRecordManager: self, ip_record_manager: IPRecordManager, user: User, - ip_information: IPInformation, ip_record_factory: Callable[..., IPRecord], ): - instance = ip_record_factory(user_id=user.user_id, ip=ip_information.ip) + ip = fake.ipv4_public() + + instance = ip_record_factory(user=user, ip=ip) assert isinstance(instance, IPRecord) assert isinstance(instance.forwarded_ips, list) @@ -260,40 +262,12 @@ class TestIPRecordManager: def test_prefetch_info( self, ip_record_factory: Callable[..., IPRecord], - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, user: User, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, ): - - ip = fake.ipv4_public() - ip_information_factory(ip=ip, geoname=ip_geoname) - ipr: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) - assert isinstance(ipr, IPRecord) - - assert ipr.information is None - assert len(ipr.forwarded_ip_records) >= 1 - assert isinstance(ipr.forwarded_ip_records, list) - fipr = ipr.forwarded_ip_records[0] - assert fipr.information is None - - ipr.prefetch_ipinfo( - pg_config=thl_web_rr, - redis_config=thl_redis_config, - include_forwarded=True, - ) - assert isinstance(ipr.information, GeoIPInformation) - assert ipr.information.ip == ipr.ip == ip - assert fipr.information is None, "the ipinfo doesn't exist in the db yet" - - ip_information_factory(ip=fipr.ip, geoname=ip_geoname) - ipr.prefetch_ipinfo( - pg_config=thl_web_rr, - redis_config=thl_redis_config, - include_forwarded=True, - ) - assert fipr.information is not None + # No more prefetch info here. Moved into UserIPHistory.enrich_ips + pass @pytest.mark.usefixtures("user_iphistory_manager_clear_cache") @@ -302,57 +276,86 @@ class TestUserIpHistoryManager: self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, + geoip_info_manager: GeoIpInfoManager, user_iphistory_manager: UserIpHistoryManager, ): instance = UserIpHistoryManager( - pg_config=thl_web_rr, redis_config=thl_redis_config + pg_config=thl_web_rr, + redis_config=thl_redis_config, + geoip_info_manager=geoip_info_manager, ) assert isinstance(instance, UserIpHistoryManager) assert isinstance(user_iphistory_manager, UserIpHistoryManager) - def test_latest_record( + def test_latest_record_and_enrich( self, user_iphistory_manager: UserIpHistoryManager, user: User, ip_record_factory: Callable[..., IPRecord], - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, + geoip_information_factory: Callable[..., GeoIPInformation], + geoip_info_manager: GeoIpInfoManager, ): ip = fake.ipv4_public() - ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) - ipr1: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) + information = geoip_information_factory( + ip=ip, is_anonymous=True, country_iso="de" + ) + lookup_results = {ip: information} + geoip_info_manager.get_multi = Mock( + side_effect=lambda ip_addresses: { + address: lookup_results[address] for address in ip_addresses + } + ) - ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) - assert isinstance(ipr, IPRecord) + ipr1 = ip_record_factory(user=user, ip=ip) + ipr = user_iphistory_manager.get_user_latest_ip_record( + user=user, + ) assert ipr.ip == ipr1.ip assert ipr.is_anonymous assert isinstance(ipr.information, GeoIPInformation) - assert ipr.information.lookup_prefix == "/32" - ip = fake.ipv6() - ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id) - ipr2: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) + assert ( + user_iphistory_manager.get_user_latest_country( + user=user, + ) + == "de" + ) + + ip2 = fake.ipv6() + ipr2: IPRecord = ip_record_factory(user=user, ip=ip2) + lookup_results[ipr2.ip] = geoip_information_factory( + ip=ipr2.ip, + country_iso="us", + is_anonymous=False, + ) - ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) - assert isinstance(ipr, IPRecord) + ipr = user_iphistory_manager.get_user_latest_ip_record( + user=user, + ) assert ipr.ip == ipr2.ip assert isinstance(ipr.information, GeoIPInformation) - assert ipr.information.lookup_prefix == "/64" assert ipr.information is not None assert not ipr.is_anonymous - country_iso = user_iphistory_manager.get_user_latest_country(user=user) - assert country_iso == ip_geoname.country_iso + assert ( + user_iphistory_manager.get_user_latest_country( + user=user, + ) + == "us" + ) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + iph = user_iphistory_manager.get_user_ip_history(user=user) assert isinstance(iph, UserIPHistory) assert isinstance(iph.ips, list) assert iph.ips[0].information is not None assert iph.ips[1].information is not None - assert iph.ips[0].country_iso == country_iso - assert iph.ips[0].is_anonymous - assert iph.ips[0].ip == ipr1.ip - assert iph.ips[1].ip == ipr2.ip + assert iph.ips[0].country_iso == "us" + assert iph.ips[1].country_iso == "de" + assert not iph.ips[0].is_anonymous + assert iph.ips[1].is_anonymous + # ordered by created DESCENDING!!!!!!!!!!!!!!1 + assert iph.ips[0].ip == ipr2.ip + assert iph.ips[1].ip == ipr1.ip def test_virgin( self, @@ -360,65 +363,9 @@ class TestUserIpHistoryManager: user_iphistory_manager: UserIpHistoryManager, ip_record_factory: Callable[..., IPRecord], ): - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + iph = user_iphistory_manager.get_user_ip_history(user=user) assert len(iph.ips) == 0 - ip_record_factory(user_id=user.user_id, ip=fake.ipv4_public()) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + ip_record_factory(user=user, ip=fake.ipv4_public()) + iph = user_iphistory_manager.get_user_ip_history(user=user) assert len(iph.ips) == 1 - - def test_out_of_order( - self, - ip_record_factory: Callable[..., IPRecord], - user: User, - user_iphistory_manager: UserIpHistoryManager, - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, - ): - # Create the user-ip association BEFORE the ip even exists in the ipinfo table - ip = fake.ipv4_public() - ip_record_factory(user_id=user.user_id, ip=ip) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) - assert isinstance(iph, UserIPHistory) - assert isinstance(iph.ips, list) - assert len(iph.ips) == 1 - ipr = iph.ips[0] - assert ipr.information is None - assert not ipr.is_anonymous - - ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) - assert isinstance(iph, UserIPHistory) - assert isinstance(iph.ips, list) - assert len(iph.ips) == 1 - ipr = iph.ips[0] - assert ipr.information is not None - assert ipr.is_anonymous - - def test_out_of_order_ipv6( - self, - ip_record_factory: Callable[..., IPRecord], - user: User, - user_iphistory_manager: UserIpHistoryManager, - ip_information_factory: Callable[..., IPInformation], - ip_geoname: IPGeoname, - ): - # Create the user-ip association BEFORE the ip even exists in the ipinfo table - ip = fake.ipv6() - ip_record_factory(user_id=user.user_id, ip=ip) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) - assert isinstance(iph, UserIPHistory) - assert isinstance(iph.ips, list) - assert len(iph.ips) == 1 - ipr = iph.ips[0] - assert ipr.information is None - assert not ipr.is_anonymous - - ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) - iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) - assert isinstance(iph, UserIPHistory) - assert isinstance(iph.ips, list) - assert len(iph.ips) == 1 - ipr = iph.ips[0] - assert ipr.information is not None - assert ipr.is_anonymous |
