diff options
| author | stuppie | 2026-09-04 12:52:46 -0600 |
|---|---|---|
| committer | stuppie | 2026-09-04 12:52:46 -0600 |
| commit | 4b705346968e38671ac9601cfa8444c944ecc8bd (patch) | |
| tree | 8aa812e35846a72b3965d136642ea84c689ec22a | |
| parent | 493f131d6ffe1629ffb3610165271c3c87b32d3b (diff) | |
| download | generalresearch-4b705346968e38671ac9601cfa8444c944ecc8bd.tar.gz generalresearch-4b705346968e38671ac9601cfa8444c944ecc8bd.zip | |
remove a lot of if type_checking; fix model_rebuild errors. Remove all network models/managers/test/fixtures
77 files changed, 169 insertions, 3078 deletions
diff --git a/generalresearch/managers/network/__init__.py b/generalresearch/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/generalresearch/managers/network/__init__.py +++ /dev/null diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py deleted file mode 100644 index cdea016..0000000 --- a/generalresearch/managers/network/label.py +++ /dev/null @@ -1,152 +0,0 @@ -from __future__ import annotations - -from collections.abc import Collection -from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING - -from psycopg import sql -from pydantic import TypeAdapter - -from generalresearch.managers.base import PostgresManager -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - IPvAnyNetwork, - IPvAnyNetworkStr, -) -from generalresearch.models.network.label import IPLabel - -if TYPE_CHECKING: - - from generalresearch.models.network.label import IPLabelKind, IPLabelSource - - -class IPLabelManager(PostgresManager): - def create(self, ip_label: IPLabel) -> IPLabel: - query = sql.SQL(""" - INSERT INTO network_iplabel ( - ip, labeled_at, created_at, - label_kind, source, confidence, - provider, metadata - ) VALUES ( - %(ip)s, %(labeled_at)s, %(created_at)s, - %(label_kind)s, %(source)s, %(confidence)s, - %(provider)s, %(metadata)s - ) RETURNING id;""") - params = ip_label.model_dump_postgres() - with self.pg_config.make_connection() as conn, conn.cursor() as c: - c.execute(query, params) - _pk = c.fetchone()["id"] - return ip_label - - def make_filter_str( - self, - ips: Collection[IPvAnyNetworkStr] | None = None, - ip_in_network: IPvAnyAddressStr | None = None, - label_kind: IPLabelKind | None = None, - source: IPLabelSource | None = None, - labeled_at: AwareDatetimeISO | None = None, - labeled_after: AwareDatetimeISO | None = None, - labeled_before: AwareDatetimeISO | None = None, - provider: str | None = None, - ): - filters = [] - params = {} - if labeled_after or labeled_before: - time_end = labeled_before or datetime.now(tz=UTC) - time_start = labeled_after or datetime(2017, 1, 1, tzinfo=UTC) - assert time_start.tzinfo.utcoffset(time_start) == timedelta(), "must be UTC" - assert time_end.tzinfo.utcoffset(time_end) == timedelta(), "must be UTC" - filters.append("labeled_at BETWEEN %(time_start)s AND %(time_end)s") - params["time_start"] = time_start - params["time_end"] = time_end - if labeled_at: - assert labeled_at.tzinfo.utcoffset(labeled_at) == timedelta(), "must be UTC" - filters.append("labeled_at == %(labeled_at)s") - params["labeled_at"] = labeled_at - if label_kind: - filters.append("label_kind = %(label_kind)s") - params["label_kind"] = label_kind.value - if source: - filters.append("source = %(source)s") - params["source"] = source.value - if provider: - filters.append("provider = %(provider)s") - params["provider"] = provider - if ips is not None: - filters.append("ip = ANY(%(ips)s)") - params["ips"] = list(ips) - if ip_in_network: - """ - Return matching networks. - e.g. ip = '13f9:c462:e039:a38c::1', might return rows - where ip = '13f9:c462:e039::/48' or '13f9:c462:e039:a38c::/64' - """ - filters.append("ip >>= %(ip_in_network)s") - params["ip_in_network"] = ip_in_network - - filter_str = "WHERE " + " AND ".join(filters) if filters else "" - return filter_str, params - - def filter( - self, - ips: Collection[IPvAnyNetworkStr] | None = None, - ip_in_network: IPvAnyAddressStr | None = None, - label_kind: IPLabelKind | None = None, - source: IPLabelSource | None = None, - labeled_at: AwareDatetimeISO | None = None, - labeled_after: AwareDatetimeISO | None = None, - labeled_before: AwareDatetimeISO | None = None, - provider: str | None = None, - ) -> list[IPLabel]: - filter_str, params = self.make_filter_str( - ips=ips, - ip_in_network=ip_in_network, - label_kind=label_kind, - source=source, - labeled_at=labeled_at, - labeled_after=labeled_after, - labeled_before=labeled_before, - provider=provider, - ) - query = f""" - SELECT - ip, labeled_at, created_at, - label_kind, source, confidence, - provider, metadata - FROM network_iplabel - {filter_str} - """ - res = self.pg_config.execute_sql_query(query, params) - return [IPLabel.model_validate(rec) for rec in res] - - def get_most_specific_matching_network(self, ip: IPvAnyAddressStr) -> IPvAnyNetwork: - """ - e.g. ip = 'b5f4:dc2:f136:70d5:5b6e:9a85:c7d4:3517', might return - 'b5f4:dc2:f136:70d5::/64' - """ - ip = TypeAdapter(IPvAnyAddressStr).validate_python(ip) - - query = """ - SELECT ip - FROM network_iplabel - WHERE ip >>= %(ip)s - ORDER BY masklen(ip) DESC - LIMIT 1;""" - res = self.pg_config.execute_sql_query(query, {"ip": ip}) - if res: - return IPvAnyNetwork(res[0]["ip"]) - - def test_join(self, ip): - query = """ - SELECT - to_jsonb(i) AS ipinfo, - to_jsonb(l) AS iplabel - FROM thl_ipinformation i - LEFT JOIN network_iplabel l - ON l.ip >>= i.ip - WHERE i.ip = %(ip)s - ORDER BY masklen(l.ip) DESC;""" - params = {"ip": ip} - res = self.pg_config.execute_sql_query(query, params) - return res diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py deleted file mode 100644 index 7b79d96..0000000 --- a/generalresearch/managers/network/mtr.py +++ /dev/null @@ -1,53 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import MTRRun - - -class MTRRunManager(PostgresManager): - - def _create(self, run: MTRRun, c: Cursor | None = None) -> None: - """ - Do not use this directly. Must only be used in the context of a toolrun - """ - query = sql.SQL(""" - INSERT INTO network_mtr ( - run_id, source_ip, facility_id, - protocol, port, parsed, - started_at, ip, scan_group_id - ) - VALUES ( - %(run_id)s, %(source_ip)s, %(facility_id)s, - %(protocol)s, %(port)s, %(parsed)s, - %(started_at)s, %(ip)s, %(scan_group_id)s - ); - """) - params = run.model_dump_postgres() - - query_hops = sql.SQL(""" - INSERT INTO network_mtrhop ( - hop, ip, domain, asn, mtr_run_id - ) VALUES ( - %(hop)s, %(ip)s, %(domain)s, - %(asn)s, %(mtr_run_id)s - ) - """) - mtr_run = run.parsed - params_hops = [h.model_dump_postgres(run_id=run.id) for h in mtr_run.hops] - - if c: - c.execute(query, params) - if params_hops: - c.executemany(query_hops, params_hops) - - else: - with self.pg_config.make_connection() as conn, conn.cursor() as _c: - _c.execute(query, params) - if params_hops: - _c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py deleted file mode 100644 index 96c6009..0000000 --- a/generalresearch/managers/network/nmap.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import NmapRun - - -class NmapRunManager(PostgresManager): - - def _create(self, run: NmapRun, c: Cursor | None = None) -> None: - """ - Insert a PortScan + PortScanPorts from a Pydantic NmapResult. - Do not use this directly. Must only be used in the context of a toolrun - """ - query = sql.SQL(""" - INSERT INTO network_portscan ( - run_id, xml_version, host_state, - host_state_reason, latency_ms, distance, - uptime_seconds, last_boot, - parsed, scan_group_id, open_tcp_ports, - started_at, ip, open_udp_ports - ) - VALUES ( - %(run_id)s, %(xml_version)s, %(host_state)s, - %(host_state_reason)s, %(latency_ms)s, %(distance)s, - %(uptime_seconds)s, %(last_boot)s, - %(parsed)s, %(scan_group_id)s, %(open_tcp_ports)s, - %(started_at)s, %(ip)s, %(open_udp_ports)s - ); - """) - params = run.model_dump_postgres() - - query_ports = sql.SQL(""" - INSERT INTO network_portscanport ( - port_scan_id, protocol, port, - state, reason, reason_ttl, - service_name - ) VALUES ( - %(port_scan_id)s, %(protocol)s, %(port)s, - %(state)s, %(reason)s, %(reason_ttl)s, - %(service_name)s - ) - """) - nmap_run = run.parsed - params_ports = [p.model_dump_postgres(run_id=run.id) for p in nmap_run.ports] - - if c: - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) - else: - with self.pg_config.make_connection() as conn, conn.cursor(): - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py deleted file mode 100644 index 1800364..0000000 --- a/generalresearch/managers/network/rdns.py +++ /dev/null @@ -1,37 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import RDNSRun - - -class RDNSRunManager(PostgresManager): - - def _create(self, run: RDNSRun, c: Cursor | None = None) -> None: - """ - Do not use this directly. Must only be used in the context of a toolrun - """ - query = """ - INSERT INTO network_rdnsresult ( - run_id, primary_hostname, primary_domain, - hostname_count, hostnames, - ip, started_at, scan_group_id - ) - VALUES ( - %(run_id)s, %(primary_hostname)s, %(primary_domain)s, - %(hostname_count)s, %(hostnames)s, - %(ip)s, %(started_at)s, %(scan_group_id)s - ); - """ - params = run.model_dump_postgres() - if c: - c.execute(query, params) - - else: - with self.pg_config.make_connection() as conn, conn.cursor() as _c: - _c.execute(query, params) diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py deleted file mode 100644 index ec06305..0000000 --- a/generalresearch/managers/network/tool_run.py +++ /dev/null @@ -1,138 +0,0 @@ -from __future__ import annotations - -from collections.abc import Collection -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager -from generalresearch.managers.network.mtr import MTRRunManager -from generalresearch.managers.network.nmap import NmapRunManager -from generalresearch.managers.network.rdns import RDNSRunManager -from generalresearch.models.network.rdns.result import RDNSResult -from generalresearch.models.network.tool_run import ( - MTRRun, - NmapRun, - RDNSRun, - ToolRun, -) - -if TYPE_CHECKING: - from generalresearch.managers.base import Permission - from generalresearch.models.network.tool_run import ToolName - from generalresearch.pg_helper import PostgresConfig - - -class ToolRunManager(PostgresManager): - def __init__( - self, - pg_config: PostgresConfig, - permissions: Collection[Permission] | None = None, - ): - super().__init__(pg_config=pg_config, permissions=permissions) - self.nmap_manager = NmapRunManager(self.pg_config) - self.rdns_manager = RDNSRunManager(self.pg_config) - self.mtr_manager = MTRRunManager(self.pg_config) - - def _create_tool_run(self, run: NmapRun | RDNSRun | MTRRun, c: Cursor): - query = sql.SQL(""" - INSERT INTO network_toolrun ( - ip, scan_group_id, tool_class, - tool_name, tool_version, started_at, - finished_at, status, raw_command, - config - ) - VALUES ( - %(ip)s, %(scan_group_id)s, %(tool_class)s, - %(tool_name)s, %(tool_version)s, %(started_at)s, - %(finished_at)s, %(status)s, %(raw_command)s, - %(config)s - ) RETURNING id; - """) - params = run.model_dump_postgres() - c.execute(query, params) - run_id = c.fetchone()["id"] - run.id = run_id - - def create_tool_run(self, run: NmapRun | RDNSRun | MTRRun): - if type(run) is NmapRun: - return self.create_nmap_run(run) - elif type(run) is RDNSRun: - return self.create_rdns_run(run) - elif type(run) is MTRRun: - return self.create_mtr_run(run) - else: - raise ValueError("unrecognized run type") - - def get_latest_runs_by_tool(self, ip: str) -> dict[ToolName, ToolRun]: - query = """ - SELECT DISTINCT ON (tool_name) * - FROM network_toolrun - WHERE ip = %(ip)s - ORDER BY tool_name, started_at DESC; - """ - params = {"ip": ip} - res = self.pg_config.execute_sql_query(query, params=params) - runs = [ToolRun.model_validate(x) for x in res] - return {r.tool_name: r for r in runs} - - def create_nmap_run(self, run: NmapRun) -> NmapRun: - """ - Insert a PortScan + PortScanPorts from a Pydantic NmapResult. - """ - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.nmap_manager._create(run, c=c) - return run - - def get_nmap_run(self, id: int) -> NmapRun: - query = """ - SELECT tr.*, np.parsed - FROM network_toolrun tr - JOIN network_portscan np ON tr.id = np.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - return NmapRun.model_validate(res) - - def create_rdns_run(self, run: RDNSRun) -> RDNSRun: - """ - Insert a RDnsRun + RDNSResult - """ - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.rdns_manager._create(run, c=c) - return run - - def get_rdns_run(self, id: int) -> RDNSRun: - query = """ - SELECT tr.*, hostnames - FROM network_toolrun tr - JOIN network_rdnsresult np ON tr.id = np.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - parsed = RDNSResult.model_validate( - {"ip": res["ip"], "hostnames": res["hostnames"]} - ) - res["parsed"] = parsed - return RDNSRun.model_validate(res) - - def create_mtr_run(self, run: MTRRun) -> MTRRun: - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.mtr_manager._create(run, c=c) - return run - - def get_mtr_run(self, id: int) -> MTRRun: - query = """ - SELECT tr.*, mtr.parsed, mtr.source_ip, mtr.facility_id - FROM network_toolrun tr - JOIN network_mtr mtr ON tr.id = mtr.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - return MTRRun.model_validate(res) diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 0f5453b..89c7871 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -3,11 +3,12 @@ from __future__ import annotations import json from datetime import UTC, datetime from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator +from generalresearch.models.cint import CintQuestionIdType from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import Source from generalresearch.models.string_utils import remove_nbsp @@ -15,12 +16,9 @@ from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.cint import CintQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) class CintQuestionType(StrEnum): diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index cd429dd..b2a8935 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -4,7 +4,7 @@ import json import logging from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -18,6 +18,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.cint import CintQuestionIdType from generalresearch.models.custom_types import ( AlphaNumStr, AwareDatetimeISO, @@ -31,10 +32,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.cint import CintQuestionIdType - - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) @@ -76,9 +73,9 @@ class CintQuota(BaseModel): @model_validator(mode="after") def validate_condition_len(self) -> Self: if self.quota_type == "total": - assert ( - self.condition_hashes is None - ), "total quota should not have conditions" + assert self.condition_hashes is None, ( + "total quota should not have conditions" + ) elif self.quota_type == "client": assert len(self.condition_hashes) > 0, "quota must have conditions" return self diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py index 31a0173..4ae8de4 100644 --- a/generalresearch/models/cint/task_collection.py +++ b/generalresearch/models/cint/task_collection.py @@ -1,19 +1,15 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator +from generalresearch.models.cint.survey import CintSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.cint.survey import CintSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 491157b..6ff397f 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -26,7 +26,7 @@ from generalresearch.models.custom_types import ( CoercedStr, DeviceTypes, ) -from generalresearch.models.definitions import Source +from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.demographics import ( Gender, @@ -37,10 +37,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.definitions import TaskCalculationType - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) @@ -136,7 +132,7 @@ class DynataCondition(MarketplaceCondition): if cell["kind"] == "RANGE": d["values"] = [ - f"{cell["range"]["from"] or "inf"}-{cell["range"]["to"] or "inf"}" + f"{cell['range']['from'] or 'inf'}-{cell['range']['to'] or 'inf'}" ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index c6cdc19..e6f0548 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,14 +8,12 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.dynata import DynataStatus +from generalresearch.models.dynata.survey import DynataSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.dynata.survey import DynataSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index 6423399..4af0639 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -4,21 +4,19 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.innovate import InnovateQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.innovate import InnovateQuestionID - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 7232228..3ca990c 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -6,7 +6,6 @@ from datetime import UTC, date from decimal import Decimal from functools import cached_property from typing import ( - TYPE_CHECKING, Annotated, Any, Literal, @@ -33,12 +32,14 @@ from generalresearch.models.custom_types import ( from generalresearch.models.definitions import ( LogicalOperator, Source, + TaskCalculationType, ) from generalresearch.models.innovate import ( InnovateDuplicateCheckLevel, InnovateQuotaStatus, InnovateStatus, ) +from generalresearch.models.innovate.question import InnovateQuestionID from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -46,13 +47,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.definitions import ( - TaskCalculationType, - ) - from generalresearch.models.innovate.question import InnovateQuestionID - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/innovate/task_collection.py b/generalresearch/models/innovate/task_collection.py index a647d30..7bf9d0f 100644 --- a/generalresearch/models/innovate/task_collection.py +++ b/generalresearch/models/innovate/task_collection.py @@ -1,20 +1,16 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.innovate import InnovateStatus +from generalresearch.models.innovate.survey import InnovateSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.innovate.survey import InnovateSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/legacy/offerwall.py b/generalresearch/models/legacy/offerwall.py index 013b506..da28663 100644 --- a/generalresearch/models/legacy/offerwall.py +++ b/generalresearch/models/legacy/offerwall.py @@ -1,34 +1,27 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.legacy.bucket import ( + BucketBase, + MarketplaceBucket, + OneShotOfferwallBucket, + OneShotSoftPairOfferwallBucket, + SingleEntryBucket, + SoftPairBucket, + TimeBucksBucket, + TopNBucket, + TopNPlusBucket, + TopNPlusRecontactBucket, + WXETOfferwallBucket, +) from generalresearch.models.legacy.definitions import OfferwallReason from generalresearch.models.thl.payout_format import ( PayoutFormatField, + PayoutFormatType, ) - -if TYPE_CHECKING: - - from generalresearch.models.legacy.bucket import ( - BucketBase, - MarketplaceBucket, - OneShotOfferwallBucket, - OneShotSoftPairOfferwallBucket, - SingleEntryBucket, - SoftPairBucket, - TimeBucksBucket, - TopNBucket, - TopNPlusBucket, - TopNPlusRecontactBucket, - WXETOfferwallBucket, - ) - from generalresearch.models.thl.payout_format import ( - PayoutFormatType, - ) - from generalresearch.models.thl.profiling.upk_question import UpkQuestion +from generalresearch.models.thl.profiling.upk_question import UpkQuestion """ Not Done: diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index e6803f0..caa6aae 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -17,17 +17,15 @@ from sentry_sdk import capture_exception from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.api_status import StatusResponse +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestionOut, +) +from generalresearch.models.thl.session import Wall +from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) + from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestionOut, - ) - from generalresearch.models.thl.session import Wall - from generalresearch.models.thl.user import User class UpkQuestionResponse(StatusResponse): diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py index c1b9e52..288f0d2 100644 --- a/generalresearch/models/lucid/question.py +++ b/generalresearch/models/lucid/question.py @@ -2,20 +2,18 @@ from __future__ import annotations import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.lucid import LucidQuestionIdType from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) - -if TYPE_CHECKING: - from generalresearch.models.lucid import LucidQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py index 02b31ab..bca471b 100644 --- a/generalresearch/models/lucid/survey.py +++ b/generalresearch/models/lucid/survey.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt @@ -11,14 +11,12 @@ from generalresearch.models.custom_types import ( UUIDStr, ) from generalresearch.models.definitions import Source +from generalresearch.models.thl.locales import CountryISO, LanguageISO from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.thl.locales import CountryISO, LanguageISO - class LucidCondition(MarketplaceCondition): model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore") diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py index 909992f..7c676fb 100644 --- a/generalresearch/models/morning/question.py +++ b/generalresearch/models/morning/question.py @@ -1,20 +1,18 @@ import json from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.locales import Localelator from generalresearch.models.definitions import Source +from generalresearch.models.morning import MorningQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) -if TYPE_CHECKING: - from generalresearch.models.morning import MorningQuestionID - # todo: we could validate that the country_iso / language_iso exists ... locale_helper = Localelator() diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 01cc4ff..8738631 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -6,7 +6,6 @@ from datetime import UTC from decimal import Decimal from functools import cached_property from typing import ( - TYPE_CHECKING, Annotated, Any, Literal, @@ -30,24 +29,20 @@ from generalresearch.models.custom_types import ( UUIDStrCoerce, ) from generalresearch.models.definitions import Source -from generalresearch.models.morning import MorningStatus +from generalresearch.models.morning import MorningQuestionID, MorningStatus +from generalresearch.models.morning.question import MorningQuestion from generalresearch.models.thl.demographics import Gender +from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISOs, +) from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.morning import MorningQuestionID - from generalresearch.models.morning.question import MorningQuestion - from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISOs, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index 1117937..9303a2f 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -1,20 +1,16 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.morning import MorningStatus +from generalresearch.models.morning.survey import MorningBid from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.morning.survey import MorningBid - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/network/__init__.py b/generalresearch/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/generalresearch/models/network/__init__.py +++ /dev/null diff --git a/generalresearch/models/network/definitions.py b/generalresearch/models/network/definitions.py deleted file mode 100644 index 2e1ab91..0000000 --- a/generalresearch/models/network/definitions.py +++ /dev/null @@ -1,70 +0,0 @@ -from __future__ import annotations - -from enum import StrEnum -from ipaddress import ip_address, ip_network - -CGNAT_NET = ip_network("100.64.0.0/10") - - -class IPProtocol(StrEnum): - TCP = "tcp" - UDP = "udp" - SCTP = "sctp" - IP = "ip" - ICMP = "icmp" - ICMPv6 = "icmpv6" - - def to_number(self) -> int: - # https://www.iana.org/assignments/protocol-numbers/protocol-numbers.xhtml - return { - self.TCP: 6, - self.UDP: 17, - self.SCTP: 132, - self.IP: 4, - self.ICMP: 1, - self.ICMPv6: 58, - }[self] - - -class IPKind(StrEnum): - PUBLIC = "public" - PRIVATE = "private" - CGNAT = "carrier_nat" - LOOPBACK = "loopback" - LINK_LOCAL = "link_local" - MULTICAST = "multicast" - RESERVED = "reserved" - UNSPECIFIED = "unspecified" - - -def get_ip_kind(ip: str | None) -> IPKind | None: - if not ip: - return None - - ip_obj = ip_address(ip) - - if ip_obj in CGNAT_NET: - return IPKind.CGNAT - - if ip_obj.is_loopback: - return IPKind.LOOPBACK - - if ip_obj.is_link_local: - return IPKind.LINK_LOCAL - - if ip_obj.is_multicast: - return IPKind.MULTICAST - - if ip_obj.is_unspecified: - return IPKind.UNSPECIFIED - - if ip_obj.is_private: - return IPKind.PRIVATE - - if ip_obj.is_reserved: - return IPKind.RESERVED - - if ip_obj.is_global: - return IPKind.PUBLIC - - return None diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py deleted file mode 100644 index c27f36f..0000000 --- a/generalresearch/models/network/label.py +++ /dev/null @@ -1,129 +0,0 @@ -from __future__ import annotations - -import ipaddress -from enum import StrEnum -from ipaddress import IPv4Network, IPv6Network -from typing import TYPE_CHECKING - -from pydantic import ( - BaseModel, - ConfigDict, - Field, - IPvAnyNetwork, - computed_field, - field_validator, -) - -from generalresearch.models.custom_types import now_utc_factory - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - - -class IPTrustClass(StrEnum): - TRUSTED = "trusted" - UNTRUSTED = "untrusted" - # Note: use case of unknown is for e.g. Spur says this IP is a residential proxy - # on 2026-1-1, and then has no annotation a month later. It doesn't mean - # the IP is TRUSTED, but we want to record that Spur now doesn't claim UNTRUSTED. - UNKNOWN = "unknown" - - -class IPLabelKind(StrEnum): - # --- UNTRUSTED --- - RESIDENTIAL_PROXY = "residential_proxy" - DATACENTER_PROXY = "datacenter_proxy" - ISP_PROXY = "isp_proxy" - MOBILE_PROXY = "mobile_proxy" - PROXY = "proxy" - HOSTING = "hosting" - VPN = "vpn" - RELAY = "relay" - TOR_EXIT = "tor_exit" - BAD_ACTOR = "bad_actor" - # --- TRUSTED --- - TRUSTED_USER = "trusted_user" - # --- UNKNOWN --- - UNKNOWN = "unknown" - - -class IPLabelSource(StrEnum): - # We got this IP from our own use of a proxy service - INTERNAL_USE = "internal_use" - - # An external "security" service flagged this IP - GRIP = "grip" - SPUR = "spur" - IPINFO = "ipinfo" - MAXMIND = "maxmind" - - MANUAL = "manual" - - -class IPLabel(BaseModel): - """ - Stores *ground truth* about an IP at a specific time. - To be used for model training and evaluation. - """ - - model_config = ConfigDict(validate_assignment=True) - - ip: IPvAnyNetwork = Field() - - labeled_at: AwareDatetimeISO = Field(default_factory=now_utc_factory) - created_at: AwareDatetimeISO | None = Field(default=None) - - label_kind: IPLabelKind = Field() - source: IPLabelSource = Field() - - confidence: float = Field(default=1.0, ge=0.0, le=1.0) - - # Optionally, if this is untrusted, which service is providing the proxy/vpn service - provider: str | None = Field( - default=None, examples=["geonode", "gecko"], max_length=128 - ) - - metadata: IPLabelMetadata | None = Field(default=None) - - @field_validator("ip", mode="before") - @classmethod - def normalize_and_validate_network( - cls, v: IPvAnyNetwork - ) -> IPv4Network | IPv6Network | None: - net = ipaddress.ip_network(address=v, strict=False) - - if isinstance(net, ipaddress.IPv6Network) and net.prefixlen > 64: - raise ValueError("IPv6 network must be /64 or larger") - - return net - - @field_validator("provider", mode="before") - @classmethod - def provider_format(cls, v: str | None) -> str | None: - if v is None: - return v - return v.lower().strip() - - @computed_field() - @property - def trust_class(self) -> IPTrustClass: - if self.label_kind == IPLabelKind.UNKNOWN: - return IPTrustClass.UNKNOWN - if self.label_kind == IPLabelKind.TRUSTED_USER: - return IPTrustClass.TRUSTED - return IPTrustClass.UNTRUSTED - - def model_dump_postgres(self): - d = self.model_dump(mode="json") - d["metadata"] = self.metadata.model_dump_json() if self.metadata else None - return d - - -class IPLabelMetadata(BaseModel): - """ - To be expanded. Just for storing some things from Spur for now - """ - - model_config = ConfigDict(validate_assignment=True, extra="allow") - - services: list[str] | None = Field(min_length=1, examples=[["RDP"]]) diff --git a/generalresearch/models/network/mtr/__init__.py b/generalresearch/models/network/mtr/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/generalresearch/models/network/mtr/__init__.py +++ /dev/null diff --git a/generalresearch/models/network/mtr/command.py b/generalresearch/models/network/mtr/command.py deleted file mode 100644 index 7e74f20..0000000 --- a/generalresearch/models/network/mtr/command.py +++ /dev/null @@ -1,76 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.parser import parse_mtr_output - -if TYPE_CHECKING: - from generalresearch.models.network.mtr.result import MTRResult - from generalresearch.models.network.tool_run_command import MTRRunCommand - -SUPPORTED_PROTOCOLS = { - IPProtocol.TCP, - IPProtocol.UDP, - IPProtocol.SCTP, - IPProtocol.ICMP, -} -PROTOCOLS_W_PORT = {IPProtocol.TCP, IPProtocol.UDP, IPProtocol.SCTP} - - -def build_mtr_command( - ip: str, - protocol: IPProtocol | None = None, - port: int | None = None, - report_cycles: int | None = 10, -) -> str: - # https://manpages.ubuntu.com/manpages/focal/man8/mtr.8.html - # e.g. "mtr -r -c 2 -b -z -j -T -P 443 74.139.70.149" - args = ["mtr", "--report", "--show-ips", "--aslookup", "--json"] - if report_cycles is not None: - args.extend(["-c", str(int(report_cycles))]) - if port is not None: - if protocol is None: - protocol = IPProtocol.TCP - assert protocol in PROTOCOLS_W_PORT, "port only allowed for TCP/SCTP/UDP traces" - args.extend(["--port", str(int(port))]) - if protocol: - assert protocol in SUPPORTED_PROTOCOLS, f"unsupported protocol: {protocol}" - # default is ICMP (no args) - arg_map = { - IPProtocol.TCP: "--tcp", - IPProtocol.UDP: "--udp", - IPProtocol.SCTP: "--sctp", - } - if protocol in arg_map: - args.append(arg_map[protocol]) - args.append(ip) - return " ".join(args) - - -def get_mtr_version() -> str: - proc = subprocess.run( - ["mtr", "-v"], - capture_output=True, - text=True, - check=False, - ) - # e.g. mtr 0.95 - ver_str = proc.stdout.strip() - return ver_str.split(" ", 1)[1] - - -def run_mtr(config: MTRRunCommand) -> MTRResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_mtr_output( - raw, protocol=config.options.protocol, port=config.options.port - ) diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py deleted file mode 100644 index 5ab7632..0000000 --- a/generalresearch/models/network/mtr/execute.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from datetime import UTC, datetime -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.command import ( - get_mtr_version, - run_mtr, -) -from generalresearch.models.network.tool_run import ( - MTRRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - MTRRunCommandOptions, -) -from generalresearch.models.network.utils import get_source_ip - - -def execute_mtr( - ip: str, - scan_group_id: UUIDStr | None = None, - protocol: IPProtocol | None = IPProtocol.ICMP, - port: int | None = None, - report_cycles: int = 10, -) -> MTRRun: - config = MTRRunCommand( - options=MTRRunCommandOptions( - ip=ip, - report_cycles=report_cycles, - protocol=protocol, - port=port, - ), - ) - - started_at = datetime.now(tz=UTC) - tool_version = get_mtr_version() - result = run_mtr(config) - finished_at = datetime.now(tz=UTC) - - return MTRRun( - tool_name=ToolName.MTR, - tool_class=ToolClass.TRACEROUTE, - tool_version=tool_version, - status=Status.SUCCESS, - ip=ip, - started_at=started_at, - finished_at=finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - source_ip=get_source_ip(), - facility_id=1, - ) diff --git a/generalresearch/models/network/mtr/parser.py b/generalresearch/models/network/mtr/parser.py deleted file mode 100644 index 30c22bf..0000000 --- a/generalresearch/models/network/mtr/parser.py +++ /dev/null @@ -1,19 +0,0 @@ -import json -from typing import Any - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.result import MTRResult - - -def parse_mtr_output(raw: str, port: int, protocol: IPProtocol) -> MTRResult: - data = parse_mtr_raw_output(raw) - data["port"] = port - data["protocol"] = protocol - return MTRResult.model_validate(data) - - -def parse_mtr_raw_output(raw: str) -> dict[str, Any]: - data = json.loads(raw)["report"] - data.update(data.pop("mtr")) - data["hops"] = data.pop("hubs") - return data diff --git a/generalresearch/models/network/mtr/result.py b/generalresearch/models/network/mtr/result.py deleted file mode 100644 index d17136c..0000000 --- a/generalresearch/models/network/mtr/result.py +++ /dev/null @@ -1,175 +0,0 @@ -from __future__ import annotations - -import re -from functools import cached_property -from ipaddress import ip_address -from typing import TYPE_CHECKING - -import tldextract -from pydantic import ( - BaseModel, - ConfigDict, - Field, - computed_field, - field_validator, - model_validator, -) - -from generalresearch.models.network.definitions import ( - IPProtocol, - get_ip_kind, -) - -if TYPE_CHECKING: - from generalresearch.models.network.definitions import IPKind - -HOST_RE = re.compile(r"^(?P<hostname>.+?) \((?P<ip>[^)]+)\)$") - - -class MTRHop(BaseModel): - model_config = ConfigDict(populate_by_name=True) - - hop: int = Field(alias="count") - host: str - asn: int | None = Field(default=None, alias="ASN") - - loss_pct: float = Field(alias="Loss%") - sent: int = Field(alias="Snt") - - last_ms: float = Field(alias="Last") - avg_ms: float = Field(alias="Avg") - best_ms: float = Field(alias="Best") - worst_ms: float = Field(alias="Wrst") - stdev_ms: float = Field(alias="StDev") - - hostname: str | None = Field( - default=None, examples=["fixed-187-191-8-145.totalplay.net"] - ) - ip: str | None = None - - @field_validator("asn", mode="before") - @classmethod - def normalize_asn(cls, v: str): - if v is None or v == "AS???": - return None - if type(v) is int: - return v - return int(v.replace("AS", "")) - - @model_validator(mode="after") - def parse_host(self): - host = self.host.strip() - - # hostname (ip) - m = HOST_RE.match(host) - if m: - self.hostname = m.group("hostname") - self.ip = m.group("ip") - return self - - # ip only - try: - ip_address(host) - self.ip = host - self.hostname = None - return self - except ValueError: - pass - - # hostname only - self.hostname = host - self.ip = None - return self - - @cached_property - def ip_kind(self) -> IPKind | None: - return get_ip_kind(self.ip) - - @cached_property - def icmp_rate_limited(self): - if self.avg_ms == 0: - return False - return self.stdev_ms > self.avg_ms or self.worst_ms > self.best_ms * 10 - - @computed_field(examples=["totalplay.net"]) - @cached_property - def domain(self) -> str | None: - if self.hostname: - return tldextract.extract(self.hostname).top_domain_under_public_suffix - - def model_dump_postgres(self, run_id: int): - # Writes for the network_mtrhop table - d = {"mtr_run_id": run_id} - data = self.model_dump( - mode="json", - include={ - "hop", - "ip", - "domain", - "asn", - }, - ) - d.update(data) - return d - - -class MTRResult(BaseModel): - model_config = ConfigDict(populate_by_name=True) - - source: str = Field(description="Hostname of the system running mtr.", alias="src") - destination: str = Field( - description="Destination hostname or IP being traced.", alias="dst" - ) - tos: int = Field(description="IP Type-of-Service (TOS) value used for probes.") - tests: int = Field(description="Number of probes sent per hop.") - psize: int = Field(description="Probe packet size in bytes.") - bitpattern: str = Field(description="Payload byte pattern used in probes (hex).") - - # Protocol used for the traceroute - protocol: IPProtocol = Field(default=IPProtocol.ICMP) - # The target port number for TCP/SCTP/UDP traces - port: int | None = Field(default=None) - - hops: list[MTRHop] = Field() - - def model_dump_postgres(self): - # Writes for the network_mtr table - d = self.model_dump( - mode="json", - include={"port"}, - ) - d["protocol"] = self.protocol.to_number() - d["parsed"] = self.model_dump_json(indent=0) - return d - - def print_report(self) -> None: - print( - f"MTR Report → {self.destination} {self.protocol.name} {self.port or ''}\n" - ) - host_max_len = max(len(h.host) for h in self.hops) - - header = ( - f"{'Hop':>3} " - f"{'Host':<{host_max_len}} " - f"{'Kind':<10} " - f"{'ASN':<8} " - f"{'Loss%':>6} {'Sent':>5} " - f"{'Last':>7} {'Avg':>7} {'Best':>7} {'Worst':>7} {'StDev':>7}" - ) - print(header) - print("-" * len(header)) - - for hop in self.hops: - print( - f"{hop.hop:>3} " - f"{hop.host:<{host_max_len}} " - f"{hop.ip_kind or '???':<10} " - f"{hop.asn or '???':<8} " - f"{hop.loss_pct:6.1f} " - f"{hop.sent:5d} " - f"{hop.last_ms:7.1f} " - f"{hop.avg_ms:7.1f} " - f"{hop.best_ms:7.1f} " - f"{hop.worst_ms:7.1f} " - f"{hop.stdev_ms:7.1f}" - ) diff --git a/generalresearch/models/network/nmap/__init__.py b/generalresearch/models/network/nmap/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/generalresearch/models/network/nmap/__init__.py +++ /dev/null diff --git a/generalresearch/models/network/nmap/command.py b/generalresearch/models/network/nmap/command.py deleted file mode 100644 index 3509b8d..0000000 --- a/generalresearch/models/network/nmap/command.py +++ /dev/null @@ -1,51 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.nmap.parser import parse_nmap_xml - -if TYPE_CHECKING: - from generalresearch.models.network.nmap.result import NmapResult - from generalresearch.models.network.tool_run_command import NmapRunCommand - - -def build_nmap_command( - ip: str, - no_ping: bool = True, - enable_advanced: bool = True, - timing: int = 4, - ports: str | None = None, - top_ports: int | None = None, -) -> str: - # e.g. "nmap -Pn -T4 -A --top-ports 1000 -oX - scanme.nmap.org" - # https://linux.die.net/man/1/nmap - args = ["nmap"] - assert 0 <= timing <= 5 - args.append(f"-T{timing}") - if no_ping: - args.append("-Pn") - if enable_advanced: - args.append("-A") - if ports is not None: - assert top_ports is None - args.extend(["-p", ports]) - if top_ports is not None: - assert ports is None - args.extend(["--top-ports", str(top_ports)]) - - args.extend(["-oX", "-", ip]) - return " ".join(args) - - -def run_nmap(config: NmapRunCommand) -> NmapResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_nmap_xml(raw) diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py deleted file mode 100644 index 09ec28b..0000000 --- a/generalresearch/models/network/nmap/execute.py +++ /dev/null @@ -1,56 +0,0 @@ -from __future__ import annotations - -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.nmap.command import run_nmap -from generalresearch.models.network.tool_run import ( - NmapRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - NmapRunCommand, - NmapRunCommandOptions, -) - - -def execute_nmap( - ip: str, - top_ports: int | None = 1000, - ports: str | None = None, - no_ping: bool = True, - enable_advanced: bool = True, - timing: int = 4, - scan_group_id: UUIDStr | None = None, -) -> NmapRun: - config = NmapRunCommand( - options=NmapRunCommandOptions( - top_ports=top_ports, - ports=ports, - no_ping=no_ping, - enable_advanced=enable_advanced, - timing=timing, - ip=ip, - ) - ) - result = run_nmap(config) - assert result.exit_status == "success" - assert result.target_ip == ip, f"{result.target_ip=}, {ip=}" - assert result.command_line == config.to_command_str() - - run = NmapRun( - tool_name=ToolName.NMAP, - tool_class=ToolClass.PORT_SCAN, - tool_version=result.version, - status=Status.SUCCESS, - ip=ip, - started_at=result.started_at, - finished_at=result.finished_at, - raw_command=result.command_line, - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - ) - return run diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py deleted file mode 100644 index 866b4bd..0000000 --- a/generalresearch/models/network/nmap/parser.py +++ /dev/null @@ -1,414 +0,0 @@ -from __future__ import annotations - -import xml.etree.ElementTree as ET -from datetime import UTC, datetime -from typing import Any - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.nmap.result import ( - NmapHostname, - NmapHostScript, - NmapHostState, - NmapHostStatusReason, - NmapOSClass, - NmapOSMatch, - NmapPort, - NmapPortStats, - NmapResult, - NmapScanInfo, - NmapScanType, - NmapScript, - NmapService, - NmapTrace, - NmapTraceHop, - PortState, - PortStateReason, -) - - -class NmapParserException(Exception): - def __init__(self, msg): - self.msg = msg - - def __str__(self): - return self.msg - - -class NmapXmlParser: - """ - Example: https://nmap.org/book/output-formats-xml-output.html - Full DTD: https://nmap.org/book/nmap-dtd.html - """ - - @classmethod - def parse_xml(cls, nmap_data: str) -> NmapResult: - """ - Expects a full nmap scan report. - """ - - try: - root = ET.fromstring(nmap_data) - except ET.ParseError as e: - emsg = f"Wrong XML structure: cannot parse data: {e}" - raise NmapParserException(emsg) - - if root.tag != "nmaprun": - raise NmapParserException("Unpexpected data structure for XML " "root node") - return cls._parse_xml_nmaprun(root) - - @classmethod - def _parse_xml_nmaprun(cls, root: ET.Element) -> NmapResult: - """ - This method parses out a full nmap scan report from its XML root - node: <nmaprun>. We expect there is only 1 host in this report! - - :param root: Element from xml.ElementTree (top of XML the document) - """ - cls._validate_nmap_root(root) - host_count = len(root.findall(".//host")) - assert host_count == 1, f"Expected 1 host, got {host_count}" - - xml_str = ET.tostring(root, encoding="unicode").replace("\n", "") - nmap_data = {"raw_xml": xml_str} - nmap_data.update(cls._parse_nmaprun(root)) - - nmap_data["scan_infos"] = [ - cls._parse_scaninfo(scaninfo_el) - for scaninfo_el in root.findall(".//scaninfo") - ] - - nmap_data.update(cls._parse_runstats(root)) - - nmap_data.update(cls._parse_xml_host(root.find(".//host"))) - - return NmapResult.model_validate(nmap_data) - - @classmethod - def _validate_nmap_root(cls, root: ET.Element) -> None: - allowed = { - "scaninfo", - "host", - "runstats", - "verbose", - "debugging", - "taskprogress", - } - - found = {child.tag for child in root} - unexpected = found - allowed - if unexpected: - raise ValueError( - f"Unexpected top-level tags in nmap XML: {sorted(unexpected)}" - ) - - @classmethod - def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo: - data = {} - data["type"] = NmapScanType(scaninfo_el.attrib["type"]) - data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"]) - data["num_services"] = scaninfo_el.attrib["numservices"] - data["services"] = scaninfo_el.attrib["services"] - return NmapScanInfo.model_validate(data) - - @classmethod - def _parse_runstats(cls, root: ET.Element) -> dict: - runstats = root.find("runstats") - if runstats is None: - return {} - - finished = runstats.find("finished") - if finished is None: - return {} - - finished_at = None - ts = finished.attrib.get("time") - if ts: - finished_at = datetime.fromtimestamp(int(ts), tz=UTC) - - return { - "finished_at": finished_at, - "exit_status": finished.attrib.get("exit"), - } - - @classmethod - def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict: - nmap_data = {} - nmaprun = dict(nmaprun_el.attrib) - nmap_data["command_line"] = nmaprun["args"] - nmap_data["started_at"] = datetime.fromtimestamp( - float(nmaprun["start"]), tz=UTC - ) - nmap_data["version"] = nmaprun["version"] - nmap_data["xmloutputversion"] = nmaprun["xmloutputversion"] - return nmap_data - - @classmethod - def _parse_xml_host(cls, host_el: ET.Element) -> dict: - """ - Receives a <host> XML tag representing a scanned host with - its services. - """ - data = {} - - # <status state="up" reason="user-set" reason_ttl="0"/> - status_el = host_el.find("status") - data["host_state"] = NmapHostState(status_el.attrib["state"]) - data["host_state_reason"] = NmapHostStatusReason(status_el.attrib["reason"]) - host_state_reason_ttl = status_el.attrib.get("reason_ttl") - if host_state_reason_ttl: - data["host_state_reason_ttl"] = int(host_state_reason_ttl) - - # <address addr="108.171.53.1" addrtype="ipv4"/> - address_el = host_el.find("address") - data["target_ip"] = address_el.attrib["addr"] - - data["hostnames"] = cls._parse_hostnames(host_el.find("hostnames")) - - data["ports"], data["port_stats"] = cls._parse_xml_ports(host_el.find("ports")) - - uptime = host_el.find("uptime") - if uptime is not None: - data["uptime_seconds"] = int(uptime.attrib["seconds"]) - - distance = host_el.find("distance") - if distance is not None: - data["distance"] = int(distance.attrib["value"]) - - tcpsequence = host_el.find("tcpsequence") - if tcpsequence is not None: - data["tcp_sequence_index"] = int(tcpsequence.attrib["index"]) - data["tcp_sequence_difficulty"] = tcpsequence.attrib["difficulty"] - ipidsequence = host_el.find("ipidsequence") - if ipidsequence is not None: - data["ipid_sequence_class"] = ipidsequence.attrib["class"] - tcptssequence = host_el.find("tcptssequence") - if tcptssequence is not None: - data["tcp_timestamp_class"] = tcptssequence.attrib["class"] - - times_elem = host_el.find("times") - if times_elem is not None: - data.update( - { - "srtt_us": int(times_elem.attrib.get("srtt", 0)) or None, - "rttvar_us": int(times_elem.attrib.get("rttvar", 0)) or None, - "timeout_us": int(times_elem.attrib.get("to", 0)) or None, - } - ) - - hostscripts_el = host_el.find("hostscript") - if hostscripts_el is not None: - data["host_scripts"] = [ - NmapHostScript(id=el.attrib["id"], output=el.attrib.get("output")) - for el in hostscripts_el.findall("script") - ] - - data["os_matches"] = cls._parse_os_matches(host_el) - - data["trace"] = cls._parse_trace(host_el) - - return data - - @classmethod - def _parse_os_matches(cls, host_el: ET.Element) -> list[NmapOSMatch] | None: - os_elem = host_el.find("os") - if os_elem is None: - return None - - matches: list[NmapOSMatch] = [] - - for m in os_elem.findall("osmatch"): - classes: list[NmapOSClass] = [] - - for c in m.findall("osclass"): - cpes = [e.text.strip() for e in c.findall("cpe") if e.text] - - classes.append( - NmapOSClass( - vendor=c.attrib.get("vendor"), - osfamily=c.attrib.get("osfamily"), - osgen=c.attrib.get("osgen"), - accuracy=( - int(c.attrib["accuracy"]) - if "accuracy" in c.attrib - else None - ), - cpe=cpes or None, - ) - ) - - matches.append( - NmapOSMatch( - name=m.attrib["name"], - accuracy=int(m.attrib["accuracy"]), - classes=classes, - ) - ) - - return matches or None - - @classmethod - def _parse_hostnames(cls, hostnames_el: ET.Element) -> list[NmapHostname]: - """ - Parses the hostnames element. - e.g. <hostnames> - <hostname name="108-171-53-1.aceips.com" type="PTR"/> - </hostnames> - """ - return [ - cls._parse_hostname(hname) for hname in hostnames_el.findall("hostname") - ] - - @classmethod - def _parse_hostname(cls, hostname_el: ET.Element) -> NmapHostname: - """ - Parses the hostname element. - e.g. <hostname name="108-171-53-1.aceips.com" type="PTR"/> - - :param hostname_el: <hostname> XML tag from a nmap scan - """ - return NmapHostname.model_validate(dict(hostname_el.attrib)) - - @classmethod - def _parse_xml_ports( - cls, ports_elem: ET.Element - ) -> tuple[list[NmapPort], NmapPortStats]: - """ - Parses the list of scanned services from a targeted host. - """ - ports: list[NmapPort] = [] - stats = NmapPortStats() - - # handle extraports first - for e in ports_elem.findall("extraports"): - state = PortState(e.attrib["state"]) - count = int(e.attrib["count"]) - - key = state.value.replace("|", "_") - setattr(stats, key, getattr(stats, key) + count) - - for port_elem in ports_elem.findall("port"): - port = cls._parse_xml_port(port_elem) - ports.append(port) - key = port.state.value.replace("|", "_") - setattr(stats, key, getattr(stats, key) + 1) - return ports, stats - - @classmethod - def _parse_xml_service(cls, service_elem: ET.Element) -> NmapService: - svc = { - "name": service_elem.attrib.get("name"), - "product": service_elem.attrib.get("product"), - "version": service_elem.attrib.get("version"), - "extrainfo": service_elem.attrib.get("extrainfo"), - "method": service_elem.attrib.get("method"), - "conf": ( - int(service_elem.attrib["conf"]) - if "conf" in service_elem.attrib - else None - ), - "cpe": [e.text.strip() for e in service_elem.findall("cpe")], - } - - return NmapService.model_validate(svc) - - @classmethod - def _parse_xml_script(cls, script_elem: ET.Element) -> NmapScript: - output = script_elem.attrib.get("output") - if output: - output = output.strip() - script = { - "id": script_elem.attrib["id"], - "output": output, - } - - elements: dict[str, Any] = {} - - # handle <elem key="...">value</elem> - for elem in script_elem.findall(".//elem"): - key = elem.attrib.get("key") - if key: - elements[key.strip()] = elem.text.strip() - - script["elements"] = elements - return NmapScript.model_validate(script) - - @classmethod - def _parse_xml_port(cls, port_elem: ET.Element) -> NmapPort: - """ - <port protocol="tcp" portid="61232"> - <state state="open" reason="syn-ack" reason_ttl="47"/> - <service name="socks5" extrainfo="Username/password authentication required" method="probed" conf="10"/> - <script id="socks-auth-info" output="
 Username and password"> - <table> - <elem key="name">Username and password</elem> - <elem key="method">2</elem> - </table> - </script> - </port> - """ - state_elem = port_elem.find("state") - - port = { - "port": int(port_elem.attrib["portid"]), - "protocol": port_elem.attrib["protocol"], - "state": PortState(state_elem.attrib["state"]), - "reason": ( - PortStateReason(state_elem.attrib["reason"]) - if "reason" in state_elem.attrib - else None - ), - "reason_ttl": ( - int(state_elem.attrib["reason_ttl"]) - if "reason_ttl" in state_elem.attrib - else None - ), - } - - service_elem = port_elem.find("service") - if service_elem is not None: - port["service"] = cls._parse_xml_service(service_elem) - - port["scripts"] = [] - for script_elem in port_elem.findall("script"): - port["scripts"].append(cls._parse_xml_script(script_elem)) - - return NmapPort.model_validate(port) - - @classmethod - def _parse_trace(cls, host_elem: ET.Element) -> NmapTrace | None: - trace_elem = host_elem.find("trace") - if trace_elem is None: - return None - - port_attr = trace_elem.attrib.get("port") - proto_attr = trace_elem.attrib.get("proto") - - hops: list[NmapTraceHop] = [] - - for hop_elem in trace_elem.findall("hop"): - ttl = hop_elem.attrib.get("ttl") - if ttl is None: - continue # ttl is required by the DTD but guard anyway - - rtt = hop_elem.attrib.get("rtt") - ipaddr = hop_elem.attrib.get("ipaddr") - host = hop_elem.attrib.get("host") - - hops.append( - NmapTraceHop( - ttl=int(ttl), - ipaddr=ipaddr, - rtt_ms=float(rtt) if rtt is not None else None, - host=host, - ) - ) - - return NmapTrace( - port=int(port_attr) if port_attr is not None else None, - protocol=IPProtocol(proto_attr) if proto_attr is not None else None, - hops=hops, - ) - - -def parse_nmap_xml(raw) -> NmapResult: - return NmapXmlParser.parse_xml(raw) diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py deleted file mode 100644 index 57c2e8b..0000000 --- a/generalresearch/models/network/nmap/result.py +++ /dev/null @@ -1,436 +0,0 @@ -from __future__ import annotations - -import json -from datetime import timedelta -from enum import StrEnum -from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal - -from pydantic import BaseModel, Field, computed_field - -from generalresearch.models.network.definitions import IPProtocol - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr - - -class PortState(StrEnum): - OPEN = "open" - CLOSED = "closed" - FILTERED = "filtered" - UNFILTERED = "unfiltered" - OPEN_FILTERED = "open|filtered" - CLOSED_FILTERED = "closed|filtered" - # Added by me, does not get returned. Used for book-keeping - NOT_SCANNED = "not_scanned" - - -class PortStateReason(StrEnum): - SYN_ACK = "syn-ack" - RESET = "reset" - CONN_REFUSED = "conn-refused" - NO_RESPONSE = "no-response" - SYN = "syn" - FIN = "fin" - - ICMP_NET_UNREACH = "net-unreach" - ICMP_HOST_UNREACH = "host-unreach" - ICMP_PROTO_UNREACH = "proto-unreach" - ICMP_PORT_UNREACH = "port-unreach" - - ADMIN_PROHIBITED = "admin-prohibited" - HOST_PROHIBITED = "host-prohibited" - NET_PROHIBITED = "net-prohibited" - - ECHO_REPLY = "echo-reply" - TIME_EXCEEDED = "time-exceeded" - - -class NmapScanType(StrEnum): - SYN = "syn" - CONNECT = "connect" - ACK = "ack" - WINDOW = "window" - MAIMON = "maimon" - FIN = "fin" - NULL = "null" - XMAS = "xmas" - UDP = "udp" - SCTP_INIT = "sctpinit" - SCTP_COOKIE_ECHO = "sctpcookieecho" - - -class NmapHostState(StrEnum): - UP = "up" - DOWN = "down" - UNKNOWN = "unknown" - - -class NmapHostStatusReason(StrEnum): - USER_SET = "user-set" - SYN_ACK = "syn-ack" - RESET = "reset" - ECHO_REPLY = "echo-reply" - ARP_RESPONSE = "arp-response" - NO_RESPONSE = "no-response" - NET_UNREACH = "net-unreach" - HOST_UNREACH = "host-unreach" - PROTO_UNREACH = "proto-unreach" - PORT_UNREACH = "port-unreach" - ADMIN_PROHIBITED = "admin-prohibited" - LOCALHOST_RESPONSE = "localhost-response" - - -class NmapOSClass(BaseModel): - vendor: str = None - osfamily: str = None - osgen: str | None = None - accuracy: int = None - cpe: list[str] | None = None - - -class NmapOSMatch(BaseModel): - name: str - accuracy: int - classes: list[NmapOSClass] = Field(default_factory=list) - - @property - def best_class(self) -> NmapOSClass | None: - if not self.classes: - return None - return max(self.classes, key=lambda m: m.accuracy) - - -class NmapScript(BaseModel): - """ - <script id="socks-auth-info" output="
 Username and password"> - <table> - <elem key="name">Username and password</elem> - <elem key="method">2</elem> - </table> - </script> - """ - - id: str - output: str | None = None - elements: dict[str, Any] = Field(default_factory=dict) - - -class NmapService(BaseModel): - # <service name="socks5" extrainfo="Username/password authentication required" method="probed" conf="10"/> - name: str | None = None - product: str | None = None - version: str | None = None - extrainfo: str | None = None - method: str | None = None - conf: int | None = None - cpe: list[str] = Field(default_factory=list) - - def model_dump_postgres(self): - d = self.model_dump(mode="json") - d["service_name"] = self.name - return d - - -class NmapPort(BaseModel): - port: int = Field() - protocol: IPProtocol = Field() - # Closed ports will not have a NmapPort record - state: PortState = Field() - reason: PortStateReason | None = Field(default=None) - reason_ttl: int | None = Field(default=None) - - service: NmapService | None = None - scripts: list[NmapScript] = Field(default_factory=list) - - def model_dump_postgres(self, run_id: int): - # Writes for the network_portscanport table - d = {"port_scan_id": run_id} - data = self.model_dump( - mode="json", - include={ - "port", - "state", - "reason", - "reason_ttl", - }, - ) - d.update(data) - d["protocol"] = self.protocol.to_number() - if self.service: - d.update(self.service.model_dump_postgres()) - return d - - -class NmapHostScript(BaseModel): - id: str = Field() - output: str | None = Field(default=None) - - -class NmapTraceHop(BaseModel): - """ - One hop observed during Nmap's traceroute. - - Example XML: - <hop ttl="7" ipaddr="62.115.192.20" rtt="17.17" host="gdl-b2-link.ip.twelve99.net"/> - """ - - ttl: int = Field() - - ipaddr: str | None = Field( - default=None, - description="IP address of the responding router or host", - ) - - rtt_ms: float | None = Field( - default=None, - description="Round-trip time in milliseconds for the probe reaching this hop.", - ) - - host: str | None = Field( - default=None, - description="Reverse DNS hostname for the hop if Nmap resolved one.", - ) - - -class NmapTrace(BaseModel): - """ - Traceroute information collected by Nmap. - - Nmap performs a single traceroute per host using probes matching the scan - type (typically TCP) directed at a chosen destination port. - - Example XML: - <trace port="61232" proto="tcp"> - <hop ttl="1" ipaddr="192.168.86.1" rtt="3.83"/> - ... - </trace> - """ - - port: int | None = Field( - default=None, - description="Destination port used for traceroute probes (may be absent depending on scan type).", - ) - protocol: IPProtocol | None = Field( - default=None, - description="Transport protocol used for the traceroute probes (tcp, udp, etc.).", - ) - - hops: list[NmapTraceHop] = Field( - default_factory=list, - description="Ordered list of hops observed during the traceroute.", - ) - - @property - def destination(self) -> NmapTraceHop | None: - return self.hops[-1] if self.hops else None - - -class NmapHostname(BaseModel): - # <hostname name="108-171-53-1.aceips.com" type="PTR"/> - name: str - type: Literal["PTR", "user"] | None = None - - -class NmapPortStats(BaseModel): - """ - This is counts across all protocols scanned (tcp/udp) - """ - - open: int = 0 - closed: int = 0 - filtered: int = 0 - unfiltered: int = 0 - open_filtered: int = 0 - closed_filtered: int = 0 - - -class NmapScanInfo(BaseModel): - """ - We could have multiple protocols in one run. - <scaninfo type="syn" protocol="tcp" numservices="983" services="22-1000,1100,3389,11000,61232"/> - <scaninfo type="syn" protocol="udp" numservices="983" services="1100"/> - """ - - type: NmapScanType = Field() - protocol: IPProtocol = Field() - num_services: int = Field() - services: str = Field() - - @cached_property - def port_set(self) -> set[int]: - """ - Expand the Nmap services string into a set of port numbers. - Example: - "22-25,80,443" -> {22,23,24,25,80,443} - """ - ports: set[int] = set() - for part in self.services.split(","): - if "-" in part: - start, end = part.split("-", 1) - ports.update(range(int(start), int(end) + 1)) - else: - ports.add(int(part)) - return ports - - -class NmapResult(BaseModel): - """ - A Nmap Run. Expects that we've only scanned ONE host. - """ - - command_line: str = Field() - started_at: AwareDatetimeISO = Field() - version: str = Field() - xmloutputversion: str = Field() - - scan_infos: list[NmapScanInfo] = Field(min_length=1) - - # comes from <runstats> - finished_at: AwareDatetimeISO | None = Field(default=None) - exit_status: Literal["success", "error"] | None = Field(default=None) - - ##### - # Everything below here is from within the *single* host we've scanned - ##### - - # <status state="up" reason="user-set" reason_ttl="0"/> - host_state: NmapHostState = Field() - host_state_reason: NmapHostStatusReason = Field() - host_state_reason_ttl: int | None = None - - # <address addr="108.171.53.1" addrtype="ipv4"/> - target_ip: IPvAnyAddressStr = Field() - - hostnames: list[NmapHostname] = Field() - - ports: list[NmapPort] = [] - port_stats: NmapPortStats = Field() - - # <uptime seconds="4063775" lastboot="Fri Jan 16 12:12:06 2026"/> - uptime_seconds: int | None = Field(default=None) - # <distance value="11"/> - distance: int | None = Field(description="approx number of hops", default=None) - - # <tcpsequence index="263" difficulty="Good luck!"> - tcp_sequence_index: int | None = None - tcp_sequence_difficulty: str | None = None - - # <ipidsequence class="All zeros"> - ipid_sequence_class: str | None = None - - # <tcptssequence class="1000HZ" > - tcp_timestamp_class: str | None = None - - # <times srtt="54719" rttvar="23423" to="148411"/> - srtt_us: int | None = Field( - default=None, description="smoothed RTT estimate (microseconds µs)" - ) - rttvar_us: int | None = Field( - default=None, description="RTT variance (microseconds µs)" - ) - timeout_us: int | None = Field( - default=None, description="probe timeout (microseconds µs)" - ) - - os_matches: list[NmapOSMatch] | None = Field(default=None) - - host_scripts: list[NmapHostScript] = Field(default_factory=list) - - trace: NmapTrace | None = Field(default=None) - - raw_xml: str | None = None - - @computed_field - @property - def last_boot(self) -> AwareDatetimeISO | None: - if self.uptime_seconds: - return self.started_at - timedelta(seconds=self.uptime_seconds) - - @property - def scan_info_tcp(self): - return next( - filter(lambda x: x.protocol == IPProtocol.TCP, self.scan_infos), None - ) - - @property - def scan_info_udp(self): - return next( - filter(lambda x: x.protocol == IPProtocol.UDP, self.scan_infos), None - ) - - @property - def latency_ms(self) -> float | None: - return self.srtt_us / 1000 if self.srtt_us is not None else None - - @property - def best_os_match(self) -> NmapOSMatch | None: - if not self.os_matches: - return None - return max(self.os_matches, key=lambda m: m.accuracy) - - def filter_ports(self, protocol: IPProtocol, state: PortState) -> list[NmapPort]: - return [p for p in self.ports if p.protocol == protocol and p.state == state] - - @property - def tcp_open_ports(self) -> list[int]: - """ - Returns a list of open TCP port numbers. - """ - return [ - p.port - for p in self.filter_ports(protocol=IPProtocol.TCP, state=PortState.OPEN) - ] - - @property - def udp_open_ports(self) -> list[int]: - """ - Returns a list of open UDP port numbers. - """ - return [ - p.port - for p in self.filter_ports(protocol=IPProtocol.UDP, state=PortState.OPEN) - ] - - @cached_property - def _port_index(self) -> dict[tuple[IPProtocol, int], NmapPort]: - return {(p.protocol, p.port): p for p in self.ports} - - def get_port_state( - self, port: int, protocol: IPProtocol = IPProtocol.TCP - ) -> PortState: - # Explicit (only if scanned and not closed) - if (protocol, port) in self._port_index: - return self._port_index[(protocol, port)].state - - # Check if we even scanned it - scaninfo = next((s for s in self.scan_infos if s.protocol == protocol), None) - if scaninfo and port in scaninfo.port_set: - return PortState.CLOSED - - # We didn't scan it - return PortState.NOT_SCANNED - - def model_dump_postgres(self): - # Writes for the network_portscan table - d = {} - data = self.model_dump( - mode="json", - include={ - "started_at", - "host_state", - "host_state_reason", - "distance", - "uptime_seconds", - "raw_xml", - }, - ) - d.update(data) - d["ip"] = self.target_ip - d["xml_version"] = self.xmloutputversion - d["latency_ms"] = self.latency_ms - d["last_boot"] = self.last_boot - d["parsed"] = self.model_dump_json(indent=0) - d["open_tcp_ports"] = json.dumps(self.tcp_open_ports) - d["open_udp_ports"] = json.dumps(self.udp_open_ports) - return d diff --git a/generalresearch/models/network/rdns/__init__.py b/generalresearch/models/network/rdns/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/generalresearch/models/network/rdns/__init__.py +++ /dev/null diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py deleted file mode 100644 index 2250449..0000000 --- a/generalresearch/models/network/rdns/command.py +++ /dev/null @@ -1,38 +0,0 @@ -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.rdns.parser import parse_rdns_output - -if TYPE_CHECKING: - from generalresearch.models.network.rdns.result import RDNSResult - from generalresearch.models.network.tool_run_command import RDNSRunCommand - - -def run_rdns(config: RDNSRunCommand) -> RDNSResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_rdns_output(ip=config.options.ip, raw=raw) - - -def build_rdns_command(ip: str) -> str: - # e.g. dig +noall +answer -x 1.2.3.4 - return f"dig +noall +answer -x {ip}" - - -def get_dig_version() -> str: - proc = subprocess.run( - ["dig", "-v"], - capture_output=True, - text=True, - check=False, - ) - # e.g. DiG 9.18.39-0ubuntu0.22.04.2-Ubuntu - ver_str = proc.stderr.strip() + proc.stdout.strip() - return ver_str.split("-", 1)[0].split(" ", 1)[1] diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py deleted file mode 100644 index d6de84b..0000000 --- a/generalresearch/models/network/rdns/execute.py +++ /dev/null @@ -1,44 +0,0 @@ -from __future__ import annotations - -from datetime import UTC, datetime -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.rdns.command import ( - get_dig_version, - run_rdns, -) -from generalresearch.models.network.tool_run import ( - RDNSRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - RDNSRunCommand, - RDNSRunCommandOptions, -) - - -def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None): - started_at = datetime.now(tz=UTC) - tool_version = get_dig_version() - config = RDNSRunCommand(options=RDNSRunCommandOptions(ip=ip)) - result = run_rdns(config) - finished_at = datetime.now(tz=UTC) - - run = RDNSRun( - tool_name=ToolName.DIG, - tool_class=ToolClass.RDNS, - tool_version=tool_version, - status=Status.SUCCESS, - ip=ip, - started_at=started_at, - finished_at=finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - ) - - return run diff --git a/generalresearch/models/network/rdns/parser.py b/generalresearch/models/network/rdns/parser.py deleted file mode 100644 index 31a5ed6..0000000 --- a/generalresearch/models/network/rdns/parser.py +++ /dev/null @@ -1,21 +0,0 @@ -import ipaddress -import re - -from generalresearch.models.custom_types import IPvAnyAddressStr -from generalresearch.models.network.rdns.result import RDNSResult - -PTR_RE = re.compile(r"\sPTR\s+([^\s]+)\.") - - -def parse_rdns_output(ip: IPvAnyAddressStr, raw: str) -> RDNSResult: - hostnames: list[str] = [] - - for line in raw.splitlines(): - m = PTR_RE.search(line) - if m: - hostnames.append(m.group(1)) - - return RDNSResult( - ip=ipaddress.ip_address(ip), - hostnames=hostnames, - ) diff --git a/generalresearch/models/network/rdns/result.py b/generalresearch/models/network/rdns/result.py deleted file mode 100644 index 46af643..0000000 --- a/generalresearch/models/network/rdns/result.py +++ /dev/null @@ -1,52 +0,0 @@ -from __future__ import annotations - -import json -from functools import cached_property - -import tldextract -from pydantic import BaseModel, Field, computed_field, model_validator - -from generalresearch.models.custom_types import IPvAnyAddressStr - - -class RDNSResult(BaseModel): - - ip: IPvAnyAddressStr = Field() - - hostnames: list[str] = Field(default_factory=list) - - @model_validator(mode="after") - def validate_hostname_prop(self): - assert len(self.hostnames) == self.hostname_count - if self.hostnames: - assert self.hostnames[0] == self.primary_hostname - assert self.primary_domain in self.primary_hostname - return self - - @computed_field(examples=["fixed-187-191-8-145.totalplay.net"]) - @cached_property - def primary_hostname(self) -> str | None: - if self.hostnames: - return self.hostnames[0] - - @computed_field(examples=[1]) - @cached_property - def hostname_count(self) -> int: - return len(self.hostnames) - - @computed_field(examples=["totalplay.net"]) - @cached_property - def primary_domain(self) -> str | None: - if self.primary_hostname: - return tldextract.extract( - self.primary_hostname - ).top_domain_under_public_suffix - - def model_dump_postgres(self): - # Writes for the network_rdnsresult table - d = self.model_dump( - mode="json", - include={"primary_hostname", "primary_domain", "hostname_count"}, - ) - d["hostnames"] = json.dumps(self.hostnames) - return d diff --git a/generalresearch/models/network/tool_run.py b/generalresearch/models/network/tool_run.py deleted file mode 100644 index 9088fe3..0000000 --- a/generalresearch/models/network/tool_run.py +++ /dev/null @@ -1,121 +0,0 @@ -from __future__ import annotations - -from enum import StrEnum -from typing import TYPE_CHECKING, Literal -from uuid import uuid4 - -from pydantic import BaseModel, Field, PositiveInt - -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - UUIDStr, -) - -if TYPE_CHECKING: - - from generalresearch.models.network.mtr.result import MTRResult - from generalresearch.models.network.nmap.result import NmapResult - from generalresearch.models.network.rdns.result import RDNSResult - from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - NmapRunCommand, - RDNSRunCommand, - ToolRunCommand, - ) - - -class ToolClass(StrEnum): - PORT_SCAN = "port_scan" - RDNS = "rdns" - PING = "ping" - TRACEROUTE = "traceroute" - - -class ToolName(StrEnum): - NMAP = "nmap" - RUSTMAP = "rustmap" - DIG = "dig" - PING = "ping" - TRACEROUTE = "traceroute" - MTR = "mtr" - - -class Status(StrEnum): - SUCCESS = "success" - FAILED = "failed" - TIMEOUT = "timeout" - ERROR = "error" - - -class ToolRun(BaseModel): - """ - A run of a networking tool against one host/ip. - """ - - id: PositiveInt | None = Field(default=None) - - ip: IPvAnyAddressStr = Field() - scan_group_id: UUIDStr = Field(default_factory=lambda: uuid4().hex) - tool_class: ToolClass = Field() - tool_name: ToolName = Field() - tool_version: str = Field() - - started_at: AwareDatetimeISO = Field() - finished_at: AwareDatetimeISO | None = Field(default=None) - status: Status | None = Field(default=None) - - raw_command: str = Field() - - config: ToolRunCommand = Field() - - def model_dump_postgres(self): - d = self.model_dump(mode="json", exclude={"config"}) - d["config"] = self.config.model_dump_json() - return d - - -class NmapRun(ToolRun): - tool_class: Literal[ToolClass.PORT_SCAN] = Field(default=ToolClass.PORT_SCAN) - tool_name: Literal[ToolName.NMAP] = Field(default=ToolName.NMAP) - config: NmapRunCommand = Field() - - parsed: NmapResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d.update(self.parsed.model_dump_postgres()) - return d - - -class RDNSRun(ToolRun): - tool_class: Literal[ToolClass.RDNS] = Field(default=ToolClass.RDNS) - tool_name: Literal[ToolName.DIG] = Field(default=ToolName.DIG) - config: RDNSRunCommand = Field() - - parsed: RDNSResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d.update(self.parsed.model_dump_postgres()) - return d - - -class MTRRun(ToolRun): - tool_class: Literal[ToolClass.TRACEROUTE] = Field(default=ToolClass.TRACEROUTE) - tool_name: Literal[ToolName.MTR] = Field(default=ToolName.MTR) - config: MTRRunCommand = Field() - - facility_id: int = Field(default=1) - source_ip: IPvAnyAddressStr = Field() - parsed: MTRResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d["source_ip"] = self.source_ip - d["facility_id"] = self.facility_id - d.update(self.parsed.model_dump_postgres()) - return d diff --git a/generalresearch/models/network/tool_run_command.py b/generalresearch/models/network/tool_run_command.py deleted file mode 100644 index b07b811..0000000 --- a/generalresearch/models/network/tool_run_command.py +++ /dev/null @@ -1,68 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Literal - -from pydantic import BaseModel, Field - -from generalresearch.models.network.definitions import IPProtocol - -if TYPE_CHECKING: - from generalresearch.models.custom_types import IPvAnyAddressStr - - -class ToolRunCommand(BaseModel): - command: str = Field() - options: dict[str, str | int | None] = Field(default_factory=dict) - - -class NmapRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr - top_ports: int | None = Field(default=1000) - ports: str | None = Field(default=None) - no_ping: bool = Field(default=True) - enable_advanced: bool = Field(default=True) - timing: int = Field(default=4) - - -class NmapRunCommand(ToolRunCommand): - command: Literal["nmap"] = Field(default="nmap") - options: NmapRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.nmap.command import build_nmap_command - - options = self.options - return build_nmap_command(**options.model_dump()) - - -class RDNSRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr - - -class RDNSRunCommand(ToolRunCommand): - command: Literal["dig"] = Field(default="dig") - options: RDNSRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.rdns.command import build_rdns_command - - options = self.options - return build_rdns_command(**options.model_dump()) - - -class MTRRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr = Field() - protocol: IPProtocol = Field(default=IPProtocol.ICMP) - port: int | None = Field(default=None) - report_cycles: int = Field(default=10) - - -class MTRRunCommand(ToolRunCommand): - command: Literal["mtr"] = Field(default="mtr") - options: MTRRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.mtr.command import build_mtr_command - - options = self.options - return build_mtr_command(**options.model_dump()) diff --git a/generalresearch/models/network/utils.py b/generalresearch/models/network/utils.py deleted file mode 100644 index fee9b80..0000000 --- a/generalresearch/models/network/utils.py +++ /dev/null @@ -1,5 +0,0 @@ -import requests - - -def get_source_ip(): - return requests.get("https://icanhazip.com?").text.strip() diff --git a/generalresearch/models/pollfish/question.py b/generalresearch/models/pollfish/question.py index f0c733c..aadf71e 100644 --- a/generalresearch/models/pollfish/question.py +++ b/generalresearch/models/pollfish/question.py @@ -4,17 +4,15 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import BaseModel, Field, model_validator from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index 3e39124..97f6ca1 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -4,22 +4,20 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.precision import PrecisionQuestionID from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.precision import PrecisionQuestionID - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index b77a365..cebe155 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from datetime import UTC from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -23,7 +23,7 @@ from generalresearch.models.custom_types import ( UUIDStrCoerce, ) from generalresearch.models.definitions import Source -from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision import PrecisionQuestionID, PrecisionStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -31,9 +31,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.precision import PrecisionQuestionID - class PrecisionCondition(MarketplaceCondition): question_id: PrecisionQuestionID | None = Field() diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py index 71241ad..c8db2af 100644 --- a/generalresearch/models/precision/task_collection.py +++ b/generalresearch/models/precision/task_collection.py @@ -1,18 +1,16 @@ -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision.survey import PrecisionSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.precision.survey import PrecisionSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index b963785..c0160bd 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -18,15 +18,13 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.prodege import ProdegeQuestionIdType from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.prodege import ProdegeQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 26898d0..27034b0 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -7,7 +7,7 @@ from collections import defaultdict from datetime import UTC, datetime from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -34,7 +34,9 @@ from generalresearch.models.definitions import ( ) from generalresearch.models.prodege import ( ProdegePastParticipationType, + ProdegeQuestionIdType, ProdegeStatus, + ProdgeRedirectStatus, ) from generalresearch.models.prodege.definitions import PG_COUNTRY_TO_ISO from generalresearch.models.thl.demographics import Gender @@ -44,13 +46,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.prodege import ( - ProdegeQuestionIdType, - ProdgeRedirectStatus, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 774fc7b..9f6a81b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -1,18 +1,16 @@ -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.prodege import ProdegeStatus +from generalresearch.models.prodege.survey import ProdegeSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.prodege.survey import ProdegeSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 0ec102b..5115426 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -4,7 +4,7 @@ import json import logging from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from uuid import UUID from pydantic import ( @@ -20,11 +20,9 @@ from pydantic import ( from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index f2cb63b..aa591a9 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,14 +8,12 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.repdata import RepDataStatus +from generalresearch.models.repdata.survey import RepDataSurveyHashed from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.repdata.survey import RepDataSurveyHashed - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 216b278..148b015 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -6,7 +6,7 @@ import json import logging from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -18,15 +18,13 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index c9bf431..5330ace 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -18,6 +18,14 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + IPLikeStrSet, +) from generalresearch.models.definitions import LogicalOperator, Source from generalresearch.models.sago import SagoStatus from generalresearch.models.thl.demographics import Gender @@ -27,16 +35,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - IPLikeStrSet, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/sago/task_collection.py b/generalresearch/models/sago/task_collection.py index 490f7a0..2879d9c 100644 --- a/generalresearch/models/sago/task_collection.py +++ b/generalresearch/models/sago/task_collection.py @@ -1,20 +1,18 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.sago import SagoStatus +from generalresearch.models.sago.survey import SagoSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.sago.survey import SagoSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index 4b854fb..839f7a8 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from enum import IntEnum, StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import ( @@ -18,18 +18,16 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.spectrum import SpectrumQuestionIdType from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.spectrum import SpectrumQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index 6689d27..5b330a8 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -55,7 +55,7 @@ class SpectrumCondition(MarketplaceCondition): try: values = [tuple(map(int, v.split("-"))) for v in self.values] assert all(len(x) == 2 for x in values) - except (ValueError, AssertionError): + except ValueError, AssertionError: return self self.values = sorted( {str(val) for tupl in values for val in range(tupl[0], tupl[1] + 1)} @@ -75,7 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] + f"{rs['from'] or 'inf'}-{rs['to'] or 'inf'}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index 8e49434..609114e 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -1,21 +1,17 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.spectrum import SpectrumStatus +from generalresearch.models.spectrum.survey import SpectrumSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.spectrum.survey import SpectrumSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index c8342b3..0d7ace5 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -12,12 +12,10 @@ from pydantic import ( model_validator, ) +from generalresearch.currency import USDCent from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ContestPrizeKind - -if TYPE_CHECKING: - from generalresearch.currency import USDCent - from generalresearch.models.thl.user import User +from generalresearch.models.thl.user import User class ContestEntryRule(BaseModel): @@ -88,9 +86,9 @@ class ContestPrize(BaseModel): @model_validator(mode="after") def validate_cash_value(self) -> Self: if self.kind == ContestPrizeKind.CASH: - assert ( - self.estimated_cash_value == self.cash_amount - ), "if kind is CASH, cash_amount must equal estimated_cash_value" + assert self.estimated_cash_value == self.cash_amount, ( + "if kind is CASH, cash_amount must equal estimated_cash_value" + ) return self diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 5e30778..6fc60f6 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from abc import ABC, abstractmethod from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -19,18 +19,14 @@ from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, + ContestWinner, ) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, ContestType, ) - -if TYPE_CHECKING: - from generalresearch.models.thl.contest import ( - ContestWinner, - ) - from generalresearch.models.thl.locales import CountryISOs +from generalresearch.models.thl.locales import CountryISOs class ContestBase(BaseModel, ABC): diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 261b3fc..4e90eb5 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any +from typing import Any from uuid import uuid4 from pydantic import ( @@ -16,9 +16,7 @@ from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ( ContestEntryType, ) - -if TYPE_CHECKING: - from generalresearch.models.thl.user import User +from generalresearch.models.thl.user import User class ContestEntryCreate(BaseModel): diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index efbbd0a..064c0d1 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -16,6 +16,9 @@ from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.managers.leaderboard import country_timezone from generalresearch.managers.leaderboard.manager import LeaderboardManager +from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, +) from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, @@ -39,11 +42,6 @@ from generalresearch.models.thl.leaderboard import ( LeaderboardFrequency, ) -if TYPE_CHECKING: - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) - class LeaderboardContestCreate(ContestBase): model_config = ConfigDict( @@ -75,9 +73,9 @@ class LeaderboardContestCreate(ContestBase): ranks = {x.leaderboard_rank for x in self.prizes} assert None not in ranks, "Must have leaderboard_rank defined" assert min(ranks) == 1, "Must start with rank 1" - assert ranks == set( - range(min(ranks), max(ranks) + 1) - ), "cannot skip prize leaderboard_ranks" + assert ranks == set(range(min(ranks), max(ranks) + 1)), ( + "cannot skip prize leaderboard_ranks" + ) return self @model_validator(mode="after") @@ -88,9 +86,9 @@ class LeaderboardContestCreate(ContestBase): @model_validator(mode="after") def check_end_condition(self) -> Self: - assert ( - not self.end_condition.target_entry_amount - ), "target_entry_amount not valid in leaderboard contest" + assert not self.end_condition.target_entry_amount, ( + "target_entry_amount not valid in leaderboard contest" + ) # the ends_at will get set automatically from the leaderboard_key return self @@ -174,13 +172,13 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): @model_validator(mode="after") def validate_product_lb_key(self) -> Self: - assert ( - self.product_id == self.leaderboard_key_parts["product_id"] - ), "leaderboard_key product_id is invalid" + assert self.product_id == self.leaderboard_key_parts["product_id"], ( + "leaderboard_key product_id is invalid" + ) if self.country_isos: - assert ( - len(self.country_isos) == 1 - ), "Can only set 1 country_iso in a leaderboard contest" + assert len(self.country_isos) == 1, ( + "Can only set 1 country_iso in a leaderboard contest" + ) assert ( next(iter(self.country_isos)) == self.leaderboard_key_parts["country_iso"] @@ -192,9 +190,9 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): @model_validator(mode="after") def validate_tie_break(self) -> Self: if self.tie_break_strategy == LeaderboardTieBreakStrategy.SPLIT_PRIZE_POOL: - assert all( - p.kind == ContestPrizeKind.CASH for p in self.prizes - ), "All prizes must be cash due to the tie-break strategy" + assert all(p.kind == ContestPrizeKind.CASH for p in self.prizes), ( + "All prizes must be cash due to the tie-break strategy" + ) return self @model_validator(mode="after") diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 21bc481..7e84e89 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -4,7 +4,7 @@ import logging import random from collections import defaultdict from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -20,12 +20,16 @@ from generalresearch.models.thl.contest import ( ContestEndCondition, ContestEntryRule, ContestPrize, + ContestWinner, ) from generalresearch.models.thl.contest.contest import ( Contest, ContestBase, ContestUserView, ) +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestEntryType, @@ -34,14 +38,6 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) -if TYPE_CHECKING: - from generalresearch.models.thl.contest import ( - ContestWinner, - ) - from generalresearch.models.thl.contest.contest_entry import ( - ContestEntry, - ) - logging.basicConfig() LOG = logging.getLogger() LOG.setLevel(logging.INFO) @@ -114,9 +110,9 @@ class RaffleContest(RaffleContestCreate, Contest): @model_validator(mode="after") def validate_entry_type(self): - assert all( - entry.entry_type == self.entry_type for entry in self.entries - ), f"all entries must be of type {self.entry_type}" + assert all(entry.entry_type == self.entry_type for entry in self.entries), ( + f"all entries must be of type {self.entry_type}" + ) return self @field_validator("current_amount", mode="before") diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py index 96f1b4c..e46e00a 100644 --- a/generalresearch/models/thl/profiling/upk_property.py +++ b/generalresearch/models/thl/profiling/upk_property.py @@ -2,17 +2,14 @@ from __future__ import annotations from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from generalresearch.models.custom_types import CountryISOLike, UUIDStr +from generalresearch.models.thl.category import Category from generalresearch.utils.enum import ReprEnumMeta -if TYPE_CHECKING: - from generalresearch.models.thl.category import Category - class PropertyType(StrEnum, metaclass=ReprEnumMeta): # UserProfileKnowledge Item diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py index 5124d17..46733d4 100644 --- a/generalresearch/models/thl/profiling/user_info.py +++ b/generalresearch/models/thl/profiling/user_info.py @@ -1,18 +1,14 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from pydantic import BaseModel, ConfigDict, Field from pydantic.json_schema import SkipJsonSchema from generalresearch.models.custom_types import AwareDatetimeISO - -if TYPE_CHECKING: - from generalresearch.models.definitions import Source - from generalresearch.models.thl.profiling.user_question_answer import ( - MarketplaceResearchProfileQuestion, - ) - from generalresearch.models.thl.user import User +from generalresearch.models.definitions import Source +from generalresearch.models.thl.profiling.user_question_answer import ( + MarketplaceResearchProfileQuestion, +) +from generalresearch.models.thl.user import User class UserProfileKnowledgeAnswer(BaseModel): diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index 76f819e..7749bdb 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -3,32 +3,27 @@ from __future__ import annotations from abc import ABC, abstractmethod from decimal import Decimal from itertools import product -from typing import TYPE_CHECKING from more_itertools import flatten from pydantic import BaseModel, Field +from generalresearch.models.definitions import Source from generalresearch.models.thl.demographics import ( AgeGroup, DemographicTarget, Gender, ) +from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISO, + LanguageISOs, +) from generalresearch.models.thl.survey.condition import ( ConditionValueType, + MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.definitions import Source - from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISO, - LanguageISOs, - ) - from generalresearch.models.thl.survey.condition import ( - MarketplaceCondition, - ) - class MarketplaceTask(BaseModel, ABC): """This is called a "Task" even though generally it represents a survey diff --git a/test_utils/managers/network/__init__.py b/test_utils/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/test_utils/managers/network/__init__.py +++ /dev/null diff --git a/test_utils/managers/network/conftest.py b/test_utils/managers/network/conftest.py deleted file mode 100644 index e69de29..0000000 --- a/test_utils/managers/network/conftest.py +++ /dev/null diff --git a/test_utils/models/network/__init__.py b/test_utils/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/test_utils/models/network/__init__.py +++ /dev/null diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py deleted file mode 100644 index c7bcc7e..0000000 --- a/test_utils/models/network/conftest.py +++ /dev/null @@ -1,145 +0,0 @@ -import os -from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING -from uuid import uuid4 - -import pytest -from pytest import FixtureRequest as 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=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=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/tests/managers/network/__init__.py b/tests/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/tests/managers/network/__init__.py +++ /dev/null diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py deleted file mode 100644 index abdd28f..0000000 --- a/tests/managers/network/test_label.py +++ /dev/null @@ -1,209 +0,0 @@ -import ipaddress -from datetime import datetime -from typing import TYPE_CHECKING - -import faker -import pytest -from psycopg.errors import UniqueViolation -from pydantic import ValidationError - -from generalresearch.models.network.label import ( - IPLabel, - IPLabelKind, - IPLabelMetadata, - IPLabelSource, -) -from generalresearch.models.thl.ipinfo import normalize_ip - -if TYPE_CHECKING: - from generalresearch.managers.network.label import IPLabelManager - -fake = faker.Faker() - - -@pytest.fixture -def ip_label(utc_now: datetime) -> IPLabel: - ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - return IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - metadata=IPLabelMetadata(services=["RDP"]), - ) - - -def test_model(utc_now: datetime): - ip = fake.ipv4_public() - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - assert lbl.ip.prefixlen == 32 - print(f"{lbl.ip=}") - - ip = ipaddress.IPv4Network((ip, 24), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - with pytest.raises(ValidationError, match="IPv6 network must be /64 or larger"): - IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=fake.ipv6(), - ) - - ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - ip = ipaddress.IPv6Network((ip.network_address, 48), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - -def test_create(iplabel_manager: IPLabelManager, ip_label: IPLabel): - iplabel_manager.create(ip_label) - - with pytest.raises( - UniqueViolation, match="duplicate key value violates unique constraint" - ): - iplabel_manager.create(ip_label) - - -def test_filter(iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago): - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 0 - - iplabel_manager.create(ip_label) - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 1 - - out = res[0] - assert out == ip_label - - res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago) - assert len(res) == 1 - - ip_label2 = ip_label.model_copy() - ip_label2.ip = fake.ipv4_public() - iplabel_manager.create(ip_label2) - res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip]) - assert len(res) == 2 - - -def test_filter_network( - iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago -): - print(ip_label) - ip_label = ip_label.model_copy() - ip_label.ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - - iplabel_manager.create(ip_label) - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 1 - - out = res[0] - assert out == ip_label - - res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago) - assert len(res) == 1 - - ip_label2 = ip_label.model_copy() - ip_label2.ip = fake.ipv4_public() - iplabel_manager.create(ip_label2) - res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip]) - assert len(res) == 2 - - -def test_network(iplabel_manager: IPLabelManager, utc_now: datetime): - # This is a fully-specific /128 ipv6 address. - # e.g. '51b7:b38d:8717:6c5b:cd3e:f5c3:3aba:17d' - ip = fake.ipv6() - # Generally, we'd want to annotate the /64 network - # e.g. '51b7:b38d:8717:6c5b::/64' - ip_64 = ipaddress.IPv6Network((ip, 64), strict=False) - - label = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip_64, - ) - iplabel_manager.create(label) - - # If I query for the /128 directly, I won't find it - res = iplabel_manager.filter(ips=[ip]) - assert len(res) == 0 - - # If I query for the /64 network I will - res = iplabel_manager.filter(ips=[ip_64]) - assert len(res) == 1 - - # Or, I can query for the /128 ip IN a network - res = iplabel_manager.filter(ip_in_network=ip) - assert len(res) == 1 - - -def test_label_cidr_and_ipinfo( - iplabel_manager: IPLabelManager, - ip_information_factory, - ip_geoname, - utc_now: datetime, -): - # We have network_iplabel.ip as a cidr col and - # thl_ipinformation.ip as a inet col. Make sure we can join appropriately - ip = fake.ipv6() - ip_information_factory(ip=ip, geoname=ip_geoname) - # We normalize for storage into ipinfo table - ip_norm, _ = normalize_ip(ip) - - # Test with a larger network - ip_48 = ipaddress.IPv6Network((ip, 48), strict=False) - print(f"{ip=}") - print(f"{ip_norm=}") - print(f"{ip_48=}") - label = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip_48, - ) - iplabel_manager.create(label) - - res = iplabel_manager.test_join(ip_norm) - print(res) diff --git a/tests/managers/network/test_tool_run.py b/tests/managers/network/test_tool_run.py deleted file mode 100644 index a815809..0000000 --- a/tests/managers/network/test_tool_run.py +++ /dev/null @@ -1,25 +0,0 @@ -def test_create_tool_run_from_nmap_run(nmap_run, toolrun_manager): - - toolrun_manager.create_nmap_run(nmap_run) - - run_out = toolrun_manager.get_nmap_run(nmap_run.id) - - assert nmap_run == run_out - - -def test_create_tool_run_from_rdns_run(rdns_run, toolrun_manager): - - toolrun_manager.create_rdns_run(rdns_run) - - run_out = toolrun_manager.get_rdns_run(rdns_run.id) - - assert rdns_run == run_out - - -def test_create_tool_run_from_mtr_run(mtr_run, toolrun_manager): - - toolrun_manager.create_mtr_run(mtr_run) - - run_out = toolrun_manager.get_mtr_run(mtr_run.id) - - assert mtr_run == run_out diff --git a/tests/models/network/__init__.py b/tests/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 --- a/tests/models/network/__init__.py +++ /dev/null diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py deleted file mode 100644 index 5d136c4..0000000 --- a/tests/models/network/test_mtr.py +++ /dev/null @@ -1,33 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import faker - -from generalresearch.models.network.mtr.execute import execute_mtr -from generalresearch.models.network.tool_run import ToolClass, ToolName - -if TYPE_CHECKING: - from generalresearch.managers.network.tool_run import ToolRunManager - -fake = faker.Faker() - - -def test_execute_mtr(toolrun_manager: ToolRunManager): - ip = "65.19.129.53" - - run = execute_mtr(ip=ip, report_cycles=3) - assert run.tool_name == ToolName.MTR - assert run.tool_class == ToolClass.TRACEROUTE - assert run.ip == ip - result = run.parsed - - last_hop = result.hops[-1] - assert last_hop.asn == 6939 - assert last_hop.domain == "grlengine.com" - - last_hop_1 = result.hops[-2] - assert last_hop_1.asn == 6939 - assert last_hop_1.domain == "he.net" - - toolrun_manager.create_mtr_run(run) diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py deleted file mode 100644 index 6adc9e4..0000000 --- a/tests/models/network/test_nmap.py +++ /dev/null @@ -1,39 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -import faker - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.nmap.execute import execute_nmap -from generalresearch.models.network.nmap.result import NmapResult, PortState -from generalresearch.models.network.tool_run import ToolClass, ToolName - -if TYPE_CHECKING: - from generalresearch.managers.network.tool_run import ToolRunManager - from generalresearch.models.network.tool_run import NmapRun - -fake = faker.Faker() - - -def resolve(host: str): - return subprocess.check_output(["dig", host, "+short"]).decode().strip() - - -def test_execute_nmap_scanme(toolrun_manager: ToolRunManager): - ip = resolve("scanme.nmap.org") - - run: NmapRun = execute_nmap( - ip=ip, top_ports=None, ports="20-30", enable_advanced=False - ) - assert run.tool_name == ToolName.NMAP - assert run.tool_class == ToolClass.PORT_SCAN - assert run.ip == ip - assert isinstance(run.parsed, NmapResult) - result = run.parsed - - port22 = result._port_index[(IPProtocol.TCP, 22)] - assert port22.state == PortState.OPEN - - toolrun_manager.create_nmap_run(run) diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py deleted file mode 100644 index fc9884b..0000000 --- a/tests/models/network/test_nmap_parser.py +++ /dev/null @@ -1,32 +0,0 @@ -from __future__ import annotations - -import os -from typing import TYPE_CHECKING - -import pytest - -from generalresearch.models.network.nmap.parser import parse_nmap_xml -from generalresearch.models.network.nmap.result import NmapTrace - -if TYPE_CHECKING: - from generalresearch.models.network.nmap.result import NmapResult - - -@pytest.fixture -def nmap_raw_output_2(request) -> str: - fp = os.path.join(request.config.rootpath, "data/nmaprun2.xml") - with open(fp) as f: - data = f.read() - return data - - -def test_nmap_xml_parser(nmap_raw_output: str, nmap_raw_output_2: str): - n: NmapResult = parse_nmap_xml(nmap_raw_output) - assert n.tcp_open_ports == [61232] - - assert isinstance(n.trace, NmapTrace) - assert len(n.trace.hops) == 18 - - n = parse_nmap_xml(nmap_raw_output_2) - assert n.tcp_open_ports == [22, 80, 9929, 31337] - assert n.trace is None diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py deleted file mode 100644 index 82126dd..0000000 --- a/tests/models/network/test_rdns.py +++ /dev/null @@ -1,40 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import faker - -from generalresearch.models.network.rdns.execute import execute_rdns -from generalresearch.models.network.tool_run import ToolClass, ToolName - -if TYPE_CHECKING: - from generalresearch.managers.network.tool_run import ToolRunManager - -fake = faker.Faker() - - -def test_execute_rdns_grl(toolrun_manager: ToolRunManager): - ip = "65.19.129.53" - run = execute_rdns(ip=ip) - assert run.tool_name == ToolName.DIG - assert run.tool_class == ToolClass.RDNS - assert run.ip == ip - result = run.parsed - assert result.primary_hostname == "in1-smtp.grlengine.com" - assert result.primary_domain == "grlengine.com" - assert result.hostname_count == 1 - - toolrun_manager.create_rdns_run(run) - - -def test_execute_rdns_none(toolrun_manager: ToolRunManager): - ip = fake.ipv6() - run = execute_rdns(ip) - result = run.parsed - - assert result.primary_hostname is None - assert result.primary_domain is None - assert result.hostname_count == 0 - assert result.hostnames == [] - - toolrun_manager.create_rdns_run(run) |
