diff --git a/AGENTS.md b/AGENTS.md index dccc2061c..353352638 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -695,3 +695,30 @@ Entry format: - Details: `ThHfModelBase` keeps Transformers/PT as the default GPU and fallback path, but CPU-only `HF_RUNTIME=auto` now loads `artifact_manifest.json`, selects a declared ONNX Runtime artifact, downloads only safe allow-patterns, loads schema and contract decoder from HF artifacts, and exposes the decoded artifact contract through the existing text-classifier flow. Business API response shaping now passes through generic model/runtime metadata emitted by serving. - Verification: `python3 -m unittest extensions.serving.test_th_hf_model_base extensions.serving.test_th_text_classifier extensions.serving.test_th_privacy_filter extensions.business.edge_inference_api.test_text_classifier_inference_api extensions.business.edge_inference_api.test_privacy_filter_inference_api`; `python3 -m py_compile extensions/serving/default_inference/nlp/th_hf_model_base.py extensions/business/edge_inference_api/text_classifier_inference_api.py`; required serving gate `python3 -m unittest extensions.serving.model_testing.test_llm_servings` currently fails at import with `ImportError: cannot import name 'Logger' from 'naeural_core'`. - Links: `extensions/serving/default_inference/nlp/th_hf_model_base.py`, `extensions/business/edge_inference_api/text_classifier_inference_api.py`, `extensions/serving/test_th_hf_model_base.py` + +- ID: `ML-20260723-001` +- Timestamp: `2026-07-23T13:45:20Z` +- Type: `change` +- Summary: dAuth job-secret requests now require signed 120-second timestamp nonces, and GET responses encrypt secret bundles to the authorized runner. +- Criticality: Security protocol change preventing indefinite signed-request/response replay and removing plaintext job secrets from HTTP responses. +- Details: `/add_secrets` and `/get_secrets` validate signed hex-millisecond timestamp nonces and echo them in successful signed responses. `/get_secrets` encrypts the serialized bundle to the signed requester address; clients must verify the response signer and echoed nonce before decrypting. +- Verification: `python -m unittest discover -s extensions/business/dauth -p 'test_*.py'`; cross-repo SDK dAuth client tests. +- Links: `extensions/business/dauth/dauth_mixin.py`, `extensions/business/dauth/dauth_manager.py` + +- ID: `ML-20260731-001` +- Timestamp: `2026-07-31T16:48:11Z` +- Type: `change` +- Summary: dAuth job-secret ChainStore writes and minute syncs now target only startup-cached dAuth registry peers. +- Criticality: Secret-replication boundary and recovery behavior across every dAuth server. +- Details: The dAuth manager reads registry ETH addresses once at startup, keeps local service eligibility fixed until restart, and refreshes only ETH-to-internal mappings from local NetMon state. `DAUTH_JOB_SECRETS` writes and 60-second hsync calls disable default/configured ChainStore peers. Known deferred risks: generic ChainStore does not authorize inbound operations by hash namespace, and first-response hsync has no freshness arbitration; production hardening requires an inbound ACL or dedicated authenticated replication protocol plus version-aware merges. +- Verification: `python3 -m unittest discover -s extensions/business/dauth -p 'test_*.py'`; `python3 -m py_compile extensions/business/dauth/dauth_registry.py extensions/business/dauth/dauth_manager.py extensions/business/dauth/dauth_mixin.py extensions/business/dauth/test_dauth_registry_gating.py extensions/business/dauth/test_dauth_secret_routing.py`; `git diff --check` +- Links: `extensions/business/dauth/dauth_registry.py`, `extensions/business/dauth/dauth_manager.py`, `extensions/business/dauth/dauth_mixin.py` + +- ID: `ML-20260803-001` +- Timestamp: `2026-08-03T16:09:26Z` +- Type: `change` +- Summary: dAuth server eligibility and secret-replication peers now refresh from the on-chain registry every hour; secret hsync runs every 10 minutes. +- Criticality: Authorization revocation and secret-replication routing across every dAuth server. +- Details: Lifecycle pause/resume predicates perform the rate-limited registry refresh without adding RPC calls to endpoint request paths. Successful reads remain cached for one hour; failed reads clear cached peers, fail closed, and retry after one minute. Registry reads are synchronous and rely on the SDK Web3 provider to return or time out. A removed local node causes the web app to pause and become unready; readiness returns only after a resumed Uvicorn process reports startup. Remaining dAuth nodes replace their cached peer set on their next hourly refresh. The inbound namespace authorization and version-aware hsync limitations from `ML-20260731-001` remain open. +- Verification: `python3 -m unittest discover -s extensions/business/dauth -p 'test_*.py'`; `python3 -m py_compile extensions/business/dauth/dauth_registry.py extensions/business/dauth/dauth_manager.py extensions/business/dauth/dauth_mixin.py extensions/business/dauth/test_dauth_registry_gating.py extensions/business/dauth/test_dauth_secret_routing.py`; `git diff --check` +- Links: `extensions/business/dauth/dauth_manager.py`, `extensions/business/dauth/test_dauth_registry_gating.py` diff --git a/README.md b/README.md index 9fc818682..15a5c2530 100644 --- a/README.md +++ b/README.md @@ -223,7 +223,7 @@ For further information, visit our website at [https://ratio1.ai](https://ratio1 ## Project Financing Disclaimer -This project incorporates open-source components developed with the support of financing grants **SMIS 143488** and **SMIS 156084**, provided by the Romanian Competitiveness Operational Programme. We extend our gratitude for this support, which has been instrumental in advancing our work and enabling us to share these resources with the community. +This project incorporates open-source components developed with the support of financing grants **SOLIS SMIS 143488** and **ReDeN SMIS 156084**, provided by the Romanian Competitiveness Operational Programme. We extend our gratitude for this support, which has been instrumental in advancing our work and enabling us to share these resources with the community. The content and information within this repository reflect the authors' views and do not necessarily represent those of the funding agencies. The grants have specifically supported certain aspects of this open-source project, facilitating broader dissemination and collaborative development. diff --git a/extensions/business/cybersec/red_mesh/connection_metrics.py b/extensions/business/cybersec/red_mesh/connection_metrics.py new file mode 100644 index 000000000..f8485600e --- /dev/null +++ b/extensions/business/cybersec/red_mesh/connection_metrics.py @@ -0,0 +1,98 @@ +"""Connection-window aggregation and signal semantics shared across scan levels.""" + +RESPONSIVE_CONNECTION_OUTCOMES = frozenset(("connected", "refused", "reset")) +MIN_QUALIFIED_WINDOW_ATTEMPTS = 5 +BLOCKING_BASELINE_RATE = 0.8 +BLOCKING_RESPONSE_RATE = 0.2 +THROTTLING_DROP_RATIO = 0.7 + + +def detect_connection_signals(windows: list | None) -> dict: + """Derive blocking/throttling signals from sufficiently sampled windows.""" + qualified = [] + for window in windows or []: + attempts = window.get("attempts") + if not isinstance(attempts, (int, float)) or attempts < MIN_QUALIFIED_WINDOW_ATTEMPTS: + continue + responsive_count = window.get("responsive_count") + response_rate = window.get("response_rate") + if response_rate is None and responsive_count is not None: + response_rate = responsive_count / attempts + if responsive_count is None and response_rate is not None: + responsive_count = response_rate * attempts + if response_rate is None or responsive_count is None: + continue + qualified.append({ + "attempts": attempts, + "responsive_count": responsive_count, + "response_rate": response_rate, + }) + + blocking = any( + previous["response_rate"] >= BLOCKING_BASELINE_RATE + and current["response_rate"] <= BLOCKING_RESPONSE_RATE + for previous, current in zip(qualified, qualified[1:]) + ) + + throttling = False + if len(qualified) >= 4: + first_attempts = sum(window["attempts"] for window in qualified[:2]) + last_attempts = sum(window["attempts"] for window in qualified[-2:]) + first_responsive = sum(window["responsive_count"] for window in qualified[:2]) + last_responsive = sum(window["responsive_count"] for window in qualified[-2:]) + baseline_rate = first_responsive / first_attempts + later_rate = last_responsive / last_attempts + throttling = ( + later_rate > BLOCKING_RESPONSE_RATE + and later_rate < baseline_rate * THROTTLING_DROP_RATIO + ) + + return { + "rate_limiting_detected": throttling, + "blocking_detected": blocking, + } + + +def merge_connection_windows(metrics_list: list) -> list | None: + """Merge aligned count-bearing windows, excluding unverifiable legacy samples.""" + grouped = {} + legacy_fallback = None + for metrics in metrics_list: + windows = metrics.get("success_rate_over_time") or [] + if legacy_fallback is None or len(windows) > len(legacy_fallback): + legacy_fallback = windows + for window in windows: + attempts = window.get("attempts") + responsive_count = window.get("responsive_count") + if attempts is None or attempts <= 0: + continue + if responsive_count is None: + response_rate = window.get("response_rate") + if response_rate is None: + continue + responsive_count = round(response_rate * attempts) + key = (window.get("window_start", 0), window.get("window_end", 0)) + bucket = grouped.setdefault(key, { + "attempts": 0, + "responsive_count": 0, + "connected_weight": 0.0, + }) + bucket["attempts"] += attempts + bucket["responsive_count"] += responsive_count + bucket["connected_weight"] += window.get("success_rate", 0) * attempts + + if not grouped: + return legacy_fallback or None + + merged = [] + for (window_start, window_end), counts in sorted(grouped.items()): + attempts = counts["attempts"] + merged.append({ + "window_start": window_start, + "window_end": window_end, + "success_rate": round(counts["connected_weight"] / attempts, 3), + "attempts": attempts, + "responsive_count": counts["responsive_count"], + "response_rate": round(counts["responsive_count"] / attempts, 3), + }) + return merged diff --git a/extensions/business/cybersec/red_mesh/constants.py b/extensions/business/cybersec/red_mesh/constants.py index 3b5bf78f1..162b02cad 100644 --- a/extensions/business/cybersec/red_mesh/constants.py +++ b/extensions/business/cybersec/red_mesh/constants.py @@ -182,6 +182,25 @@ class ScanType(str, Enum): PORT_ORDER_SHUFFLE = "SHUFFLE" PORT_ORDER_SEQUENTIAL = "SEQUENTIAL" +# Network target-response timeout profiles. Standard preserves every existing +# call-site timeout; Thorough expands ordinary waits without changing probe +# breadth, pacing, or timing-sensitive detection thresholds. +TIMEOUT_PROFILE_STANDARD = "STANDARD" +TIMEOUT_PROFILE_THOROUGH = "THOROUGH" +TIMEOUT_PROFILES = frozenset({TIMEOUT_PROFILE_STANDARD, TIMEOUT_PROFILE_THOROUGH}) + + +def normalize_timeout_profile(value): + normalized = str(value or TIMEOUT_PROFILE_STANDARD).strip().upper() + return normalized if normalized in TIMEOUT_PROFILES else TIMEOUT_PROFILE_STANDARD + + +def resolve_target_response_timeout(timeout_profile, standard_timeout): + """Resolve an ordinary network target-response maximum wait in seconds.""" + if normalize_timeout_profile(timeout_profile) != TIMEOUT_PROFILE_THOROUGH: + return standard_timeout + return round(min(float(standard_timeout) * 3, 15.0), 3) + # LLM Agent API status constants LLM_API_STATUS_OK = "ok" LLM_API_STATUS_ERROR = "error" @@ -298,6 +317,28 @@ class ScanType(str, Enum): ALL_PORTS = list(range(1, 65536)) +# ===================================================================== +# Geographic vantage-point comparison mode +# ===================================================================== +# When comparison mode is enabled every selected node runs the SAME "comparison +# tier" of ports (so results can be compared across countries). The distribution +# choice controls what the tier is: +# - SLICE (default): the tier is the standard COMMON_PORTS bundle; the +# operator's chosen range is split across nodes for coverage (not compared). +# - MIRROR: the tier is the whole chosen range (plus COMMON_PORTS), so every +# port is compared across countries, at N x the work. + +# Standard webapp/graybox feature bundle always run (mirrored to every node) in +# comparison mode so cross-country response divergence is meaningful even if the +# operator narrowed the selection. These are safe, unauthenticated checks; their +# methods are force-enabled (removed from excluded_features) when comparison +# mode is on. Referenced by feature id in FEATURE_CATALOG. +COMPARISON_GRAYBOX_BUNDLE_FEATURE_IDS = [ + "web_discovery", + "web_hardening", + "web_api_exposure", +] + # ===================================================================== # Risk score computation # ===================================================================== diff --git a/extensions/business/cybersec/red_mesh/mixins/live_progress.py b/extensions/business/cybersec/red_mesh/mixins/live_progress.py index 6450581f2..7e9039b2a 100644 --- a/extensions/business/cybersec/red_mesh/mixins/live_progress.py +++ b/extensions/business/cybersec/red_mesh/mixins/live_progress.py @@ -8,6 +8,7 @@ from ..graybox.models import GrayboxCredentialSet from ..models import WorkerProgress from ..constants import PHASE_ORDER, GRAYBOX_PHASE_ORDER +from ..connection_metrics import detect_connection_signals, merge_connection_windows DEFAULT_PROGRESS_PUBLISH_INTERVAL = 30.0 @@ -175,7 +176,6 @@ def _status_rank(v): all_phases[phase] = max(all_phases.get(phase, 0), dur) if all_phases: merged["phase_durations"] = all_phases - longest = max(metrics_list, key=lambda m: m.get("total_duration", 0)) # Merge stats distributions (response_times, port_scan_delays) # Use weighted mean, global min/max, approximate p95/p99 from max of per-thread values for stats_field in ("response_times", "port_scan_delays"): @@ -193,12 +193,12 @@ def _status_rank(v): "p99": round(max(s.get("p99", 0) for s in stats_list), 4), "count": total_count, } - # Success rate over time: take from the longest-running thread - if longest.get("success_rate_over_time"): - merged["success_rate_over_time"] = longest["success_rate_over_time"] - # Detection flags (any thread detecting = True) - merged["rate_limiting_detected"] = any(m.get("rate_limiting_detected") for m in metrics_list) - merged["blocking_detected"] = any(m.get("blocking_detected") for m in metrics_list) + # Merge aligned traffic evidence, then derive node signals from the combined + # sample counts. Legacy windows without counts remain visible but unverified. + connection_windows = merge_connection_windows(metrics_list) + if connection_windows: + merged["success_rate_over_time"] = connection_windows + merged.update(detect_connection_signals(connection_windows)) # Open port details: union, deduplicate by port all_details = [] seen_ports = set() diff --git a/extensions/business/cybersec/red_mesh/mixins/report.py b/extensions/business/cybersec/red_mesh/mixins/report.py index 67b7ca3f1..997ca0474 100644 --- a/extensions/business/cybersec/red_mesh/mixins/report.py +++ b/extensions/business/cybersec/red_mesh/mixins/report.py @@ -5,7 +5,10 @@ the UI aggregate view for the frontend. """ +import hashlib as _hashlib import json as _json +import math as _math +import struct as _struct from ..worker import PentestLocalWorker from ..models import UiAggregate @@ -201,6 +204,76 @@ def _dedup_in_dict_at_findings(container): return aggregated +def _iter_report_findings(report): + """Yield every raw finding record from the worker-report paths we publish.""" + if not isinstance(report, dict): + return + + def _items(value): + return value if isinstance(value, list) else () + + for section_name in ("service_info", "web_tests_info"): + for port_entry in (report.get(section_name) or {}).values(): + if not isinstance(port_entry, dict): + continue + for finding in _items(port_entry.get("findings")): + yield finding + for probe_entry in port_entry.values(): + if not isinstance(probe_entry, dict): + continue + for finding in _items(probe_entry.get("findings")): + yield finding + + for port_probes in (report.get("graybox_results") or {}).values(): + if not isinstance(port_probes, dict): + continue + for probe_entry in port_probes.values(): + if not isinstance(probe_entry, dict): + continue + for finding in _items(probe_entry.get("findings")): + yield finding + + for section_name in ("correlation_findings", "findings"): + for finding in _items(report.get(section_name)): + yield finding + + +def _compact_finding_signature(finding): + """Return a stable compact type signature without worker-attribution fields.""" + if isinstance(finding, dict): + explicit = finding.get("finding_signature") or finding.get("finding_id") + if explicit: + return str(explicit) + + def normalize(value): + if isinstance(value, (int, float)) and not isinstance(value, bool): + numeric = float(value) + if not _math.isfinite(numeric): + return None + if numeric == 0: + numeric = 0.0 + return "__redmesh_number__:" + _struct.pack(">d", numeric).hex() + if isinstance(value, dict): + return { + key: normalize(item) + for key, item in value.items() + if key not in _DEDUP_EXCLUDE_FIELDS + } + if isinstance(value, (list, tuple)): + return [normalize(item) for item in value] + return value + + try: + canonical = _json.dumps( + normalize(finding), sort_keys=True, default=str, + ensure_ascii=False, separators=(",", ":"), + ) + except (TypeError, ValueError): + canonical = repr(normalize(finding)) + stable = canonical.encode("utf-8", errors="replace") + return "sha256:" + _hashlib.sha256(stable).hexdigest() + + class _ReportMixin: """Report aggregation and UI view methods for PentesterApi01Plugin.""" @@ -221,14 +294,28 @@ def _count_nested_findings(section): def _count_all_findings(self, report): """Count all findings emitted by network and graybox reporting sections.""" - if not isinstance(report, dict): - return 0 - return ( - self._count_nested_findings(report.get("service_info")) + - self._count_nested_findings(report.get("web_tests_info")) + - len(report.get("correlation_findings") or []) + - self._count_nested_findings(report.get("graybox_results")) - ) + return sum(1 for _finding in _iter_report_findings(report)) + + @staticmethod + def _summarize_worker_findings(report): + """Count raw records and unique types before aggregate cross-worker dedup.""" + counts = {} + signatures = [] + seen_signatures = set() + nr_findings = 0 + for finding in _iter_report_findings(report): + nr_findings += 1 + severity = "INFO" + if isinstance(finding, dict): + severity = str(finding.get("severity") or "INFO").upper() + if severity not in ("CRITICAL", "HIGH", "MEDIUM", "LOW", "INFO"): + severity = "INFO" + counts[severity] = counts.get(severity, 0) + 1 + signature = _compact_finding_signature(finding) + if signature not in seen_signatures: + seen_signatures.add(signature) + signatures.append(signature) + return nr_findings, counts, signatures @staticmethod def _dedupe_items(items): @@ -648,6 +735,115 @@ def _redact_job_config(config_dict): redacted.pop("secret_ref", None) return redacted + def _resolve_node_country_tag(self, addr): + """Resolve a node's ISO-2 country from its netmon ``CT:`` tag. + + Fallback for nodes that produced no report (e.g. a fully-timed-out vantage + point), so every participating node still carries a country. Returns "" when + netmon is unavailable or the node has no country tag. + """ + netmon = getattr(self, "netmon", None) + if netmon is None: + return "" + try: + tags = netmon.get_network_node_tags(addr) or [] + except Exception: + return "" + for tag in tags: + if isinstance(tag, str) and tag.startswith("CT:"): + return tag[3:].strip().upper() + return "" + + def _compute_node_comparison(self, latest, job_config): + """Per-node vantage-point comparison for comparison-mode jobs. + + One entry per participating node — including nodes that never reported — + carrying country, reachability status, open ports, findings, and per-node + metrics so the UI/report can compare results across countries. + + Open ports and per-node metrics are node-accurate. Findings are attributed + via each finding's ``_source_node_addr`` stamp (single-node attribution + after cross-node dedup); the UI derives per-country "seen-by" sets from + these per-node finding lists. + """ + cfg = job_config or {} + worker_reports = latest.get("worker_reports") or {} + worker_scan_metrics = latest.get("worker_scan_metrics") or {} + findings = latest.get("findings") or [] + + findings_by_node = {} + for f in findings: + addr = f.get("_source_node_addr") + if not addr: + continue + findings_by_node.setdefault(addr, []).append({ + "signature": f.get("finding_signature") or f.get("finding_id"), + "severity": f.get("severity", "INFO"), + "title": f.get("title", ""), + "port": f.get("port"), + }) + + # Participating nodes = union of report authors, metric authors, and the + # originally selected peers (the latter surfaces nodes that never reported). + participating = list(dict.fromkeys( + list(worker_reports.keys()) + + list(worker_scan_metrics.keys()) + + list(cfg.get("selected_peers") or []) + )) + + comparison = [] + for addr in participating: + wr = worker_reports.get(addr) or {} + worker_metric_entry = worker_scan_metrics.get(addr) or {} + sm = worker_metric_entry.get("scan_metrics") or {} + outcomes = sm.get("connection_outcomes") or {} + response_times = sm.get("response_times") or {} + has_report = addr in worker_reports + country = (wr.get("country") or "").upper() or self._resolve_node_country_tag(addr) or "UN" + + # Status reflects REACHABILITY only. Blocking / rate-limiting are advisory + # detection signals carried in metrics (a node can be reached AND flagged), + # so they must not override "reached" when the node returned results. + if not has_report and addr not in worker_scan_metrics: + status = "failed" + elif outcomes and outcomes.get("connected", 0) == 0 and outcomes.get("timeout", 0) > 0: + status = "timeout" + else: + status = "reached" + + p95 = response_times.get("p95") + comparison.append({ + "address": addr, + "country": country, + "node_ip": wr.get("node_ip", ""), + "status": status, + "open_ports": wr.get("open_ports", []), + "nr_findings": wr.get("nr_findings", len(findings_by_node.get(addr, []))), + "finding_counts": wr.get("finding_counts"), + "finding_signatures": wr.get("finding_signatures"), + "response_evidence": wr.get("response_evidence"), + "findings": findings_by_node.get(addr, []), + "metrics": { + "connected": outcomes.get("connected", 0), + "timeout": outcomes.get("timeout", 0), + "refused": outcomes.get("refused", 0), + "reset": outcomes.get("reset", 0), + "error": outcomes.get("error", 0), + "response_p95_ms": round(p95 * 1000, 1) if isinstance(p95, (int, float)) else None, + "coverage": sm.get("coverage"), + "probes_attempted": sm.get("probes_attempted"), + "probes_completed": sm.get("probes_completed"), + "probes_failed": sm.get("probes_failed"), + "phase_durations": sm.get("phase_durations"), + "total_duration": sm.get("total_duration"), + "traffic_windows": sm.get("success_rate_over_time"), + "threads": worker_metric_entry.get("threads"), + "rate_limited": bool(sm.get("rate_limiting_detected")), + "blocked": bool(sm.get("blocking_detected")), + }, + }) + return comparison + def _compute_ui_aggregate(self, passes, latest_aggregated, job_config=None): """Compute pre-aggregated view for frontend from pass reports. @@ -702,6 +898,24 @@ def _compute_ui_aggregate(self, passes, latest_aggregated, job_config=None): finding_timeline[fid]["last_seen"] = pass_nr finding_timeline[fid]["pass_count"] += 1 + # Origin-country breakdown for the latest pass: count participating worker + # nodes per ISO-2 country (empty country grouped under "UN"/Unknown in the UI). + worker_reports = latest.get("worker_reports") or {} + country_counter = Counter( + (w.get("country") or "UN").upper() for w in worker_reports.values() + ) + country_breakdown = [ + {"code": code, "count": count} + for code, count in sorted(country_counter.items(), key=lambda kv: (-kv[1], kv[0])) + ] + + # Geographic vantage-point comparison (comparison_mode jobs only): durable + # per-node divergence — reachability, open ports, findings, and latency — + # including nodes that never reported (fully-timed-out vantage points). + node_comparison = None + if (job_config or {}).get("comparison_mode"): + node_comparison = self._compute_node_comparison(latest, job_config) or None + return UiAggregate( total_open_ports=sorted(set(agg.get("open_ports", []))), total_services=self._count_services(agg.get("service_info", {})), @@ -718,9 +932,12 @@ def _compute_ui_aggregate(self, passes, latest_aggregated, job_config=None): "start_port": w["start_port"], "end_port": w["end_port"], "open_ports": w.get("open_ports", []), + "country": (w.get("country") or "").upper(), } - for addr, w in (latest.get("worker_reports") or {}).items() + for addr, w in worker_reports.items() ] or None, + country_breakdown=country_breakdown or None, + node_comparison=node_comparison, scan_type=scan_type, total_routes_discovered=graybox_stats["total_routes_discovered"], total_forms_discovered=graybox_stats["total_forms_discovered"], diff --git a/extensions/business/cybersec/red_mesh/models/archive.py b/extensions/business/cybersec/red_mesh/models/archive.py index 7d91820d0..249164791 100644 --- a/extensions/business/cybersec/red_mesh/models/archive.py +++ b/extensions/business/cybersec/red_mesh/models/archive.py @@ -15,6 +15,7 @@ from extensions.business.cybersec.red_mesh.models.shared import _strip_none from extensions.business.cybersec.red_mesh.constants import ( DISTRIBUTION_SLICE, PORT_ORDER_SEQUENTIAL, RUN_MODE_SINGLEPASS, JOB_ARCHIVE_VERSION, + TIMEOUT_PROFILE_STANDARD, normalize_timeout_profile, ) @@ -35,6 +36,7 @@ class JobConfig: enabled_features: list # [str] excluded_features: list # [str] run_mode: str # SINGLEPASS | CONTINUOUS_MONITORING + timeout_profile: str = TIMEOUT_PROFILE_STANDARD # STANDARD | THOROUGH (network scans) scan_min_delay: float = 0 scan_max_delay: float = 0 ics_safe_mode: bool = False @@ -45,6 +47,9 @@ class JobConfig: task_description: str = "" monitor_interval: int = 0 selected_peers: list = None # [str] or None + # ── geographic vantage-point comparison mode ── + comparison_mode: bool = False # tiered mirror+slice; per-country comparison + comparison_ports: list = None # [int] ports mirrored to every node (comparison tier) created_by_name: str = "" created_by_id: str = "" authorized: bool = False @@ -135,6 +140,7 @@ def from_dict(cls, d: dict) -> JobConfig: enabled_features=d.get("enabled_features", []), excluded_features=d.get("excluded_features", []), run_mode=d.get("run_mode", RUN_MODE_SINGLEPASS), + timeout_profile=normalize_timeout_profile(d.get("timeout_profile")), scan_min_delay=d.get("scan_min_delay", 0), scan_max_delay=d.get("scan_max_delay", 0), ics_safe_mode=d.get("ics_safe_mode", False), @@ -145,6 +151,8 @@ def from_dict(cls, d: dict) -> JobConfig: task_description=d.get("task_description", ""), monitor_interval=d.get("monitor_interval", 0), selected_peers=d.get("selected_peers"), + comparison_mode=d.get("comparison_mode", False), + comparison_ports=d.get("comparison_ports"), created_by_name=d.get("created_by_name", ""), created_by_id=d.get("created_by_id", ""), authorized=d.get("authorized", False), @@ -267,12 +275,16 @@ class WorkerReportMeta: open_ports: list = None # [int] nr_findings: int = 0 node_ip: str = "" # worker node's IP address + country: str = "" # worker node's ISO-2 country (from location_data); "" when unknown + finding_counts: dict = None # compact raw record counts by severity + finding_signatures: list = None # unique raw finding-type signatures + response_evidence: dict = None # per-vantage target response fingerprint (comparison mode) def to_dict(self) -> dict: d = asdict(self) if d["open_ports"] is None: d["open_ports"] = [] - return d + return _strip_none(d) @classmethod def from_dict(cls, d: dict) -> WorkerReportMeta: @@ -284,6 +296,10 @@ def from_dict(cls, d: dict) -> WorkerReportMeta: open_ports=d.get("open_ports", []), nr_findings=d.get("nr_findings", 0), node_ip=d.get("node_ip", ""), + country=d.get("country", ""), + finding_counts=d.get("finding_counts"), + finding_signatures=d.get("finding_signatures"), + response_evidence=d.get("response_evidence"), ) @@ -376,7 +392,9 @@ class UiAggregate: findings_count: dict = None # { CRITICAL: int, HIGH: int, MEDIUM: int, LOW: int, INFO: int } top_findings: list = None # top 10 CRITICAL+HIGH findings for dashboard display finding_timeline: dict = None # { finding_id: { first_seen, last_seen, pass_count } } - worker_activity: list = None # [ { id, start_port, end_port, open_ports } ] + worker_activity: list = None # [ { id, start_port, end_port, open_ports, country } ] + country_breakdown: list = None # [ { code, count } ] origin countries the pass ran from + node_comparison: list = None # per-node vantage-point comparison (comparison_mode only); see report.py # ── graybox-aware ── scan_type: str = "network" total_routes_discovered: int = 0 # webapp: discovered routes @@ -400,6 +418,8 @@ def from_dict(cls, d: dict) -> UiAggregate: top_findings=d.get("top_findings"), finding_timeline=d.get("finding_timeline"), worker_activity=d.get("worker_activity"), + country_breakdown=d.get("country_breakdown"), + node_comparison=d.get("node_comparison"), scan_type=d.get("scan_type", "network"), total_routes_discovered=d.get("total_routes_discovered", 0), total_forms_discovered=d.get("total_forms_discovered", 0), diff --git a/extensions/business/cybersec/red_mesh/models/shared.py b/extensions/business/cybersec/red_mesh/models/shared.py index b565e31fa..0cc12dcd7 100644 --- a/extensions/business/cybersec/red_mesh/models/shared.py +++ b/extensions/business/cybersec/red_mesh/models/shared.py @@ -99,8 +99,11 @@ class ScanMetrics: # ── Detection indicators ── success_rate_over_time: list = None # [ { "window_start": 0, "window_end": 60, - # "success_rate": 0.98 }, ... ] - # degrading rate = scan likely detected + # "success_rate": 0.12, "attempts": 100, + # "responsive_count": 98, + # "response_rate": 0.98 }, ... ] + # success_rate remains connected-only; + # response_rate includes refused/reset. rate_limiting_detected: bool = False blocking_detected: bool = False diff --git a/extensions/business/cybersec/red_mesh/pentester_api_01.py b/extensions/business/cybersec/red_mesh/pentester_api_01.py index 1fd6fe013..3a84c46c0 100644 --- a/extensions/business/cybersec/red_mesh/pentester_api_01.py +++ b/extensions/business/cybersec/red_mesh/pentester_api_01.py @@ -1403,6 +1403,7 @@ def _maybe_launch_jobs(self, nr_local_workers=None): end_port=end_port, job_config=job_config, nr_local_workers_override=nr_local_workers, + target_ports=worker_entry.get("target_ports"), ) except ValueError as exc: self.P(f"Skipping job {job_id}: {exc}", color='r') @@ -3057,7 +3058,10 @@ def _close_job(self, job_id, canceled=False): location_data = self.global_shmem.get('location_data') or {} public_ip = location_data.get('ip') report["node_ip"] = public_ip or self.log.get_localhost_ip() - self.P(f"[CLOSE_JOB] Stamped node_ip={report['node_ip']} on report for job {job_id} (source={'location_data' if public_ip else 'localhost'})") + # Stamp this node's ISO-2 country (from location_data) for geographic attribution. + # Sits beside node_ip; empty when geoloc is unavailable (grouped as "Unknown" in the UI). + report["country_code"] = (location_data.get('country_code') or '').upper() + self.P(f"[CLOSE_JOB] Stamped node_ip={report['node_ip']} country={report['country_code'] or '?'} on report for job {job_id} (source={'location_data' if public_ip else 'localhost'})") # Redact credentials before persisting job_config = self._get_job_config(job_specs) @@ -3736,6 +3740,8 @@ def launch_network_scan( authorization: dict = None, unsafe_launch_confirmations: list[str] = None, blockchain_attestation_enabled: bool = False, + comparison_mode: bool = False, + timeout_profile: str = "STANDARD", ): """Launch a network scan using network-specific validation and worker slicing.""" return launch_network_scan( @@ -3772,6 +3778,8 @@ def launch_network_scan( authorization=authorization, unsafe_launch_confirmations=unsafe_launch_confirmations, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, + timeout_profile=timeout_profile, ) @BasePlugin.endpoint(method="post") @@ -3827,6 +3835,7 @@ def launch_webapp_scan( authorization: dict = None, unsafe_launch_confirmations: list[str] = None, blockchain_attestation_enabled: bool = False, + comparison_mode: bool = False, ): """Launch a graybox webapp scan using webapp-specific validation and (by default) SLICE worker assignment.""" return launch_webapp_scan( @@ -3881,6 +3890,7 @@ def launch_webapp_scan( authorization=authorization, unsafe_launch_confirmations=unsafe_launch_confirmations, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, ) @BasePlugin.endpoint(method="post") @@ -3943,6 +3953,8 @@ def launch_test( authorization: dict = None, unsafe_launch_confirmations: list[str] = None, blockchain_attestation_enabled: bool = False, + comparison_mode: bool = False, + timeout_profile: str = "STANDARD", ): """Compatibility shim that routes to scan-type-specific launch endpoints.""" return launch_test( @@ -4005,6 +4017,8 @@ def launch_test( authorization=authorization, unsafe_launch_confirmations=unsafe_launch_confirmations, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, + timeout_profile=timeout_profile, ) @BasePlugin.endpoint(method="post", require_token=True) diff --git a/extensions/business/cybersec/red_mesh/services/finalization.py b/extensions/business/cybersec/red_mesh/services/finalization.py index ec67d5b6f..11709dc18 100644 --- a/extensions/business/cybersec/red_mesh/services/finalization.py +++ b/extensions/business/cybersec/red_mesh/services/finalization.py @@ -296,6 +296,10 @@ def maybe_finalize_pass(owner): job_specs = _write_job_record(owner, job_key, job_specs, context="finalize_collecting") node_reports = owner._collect_node_reports(workers) + worker_finding_summaries = { + addr: owner._summarize_worker_findings(report) + for addr, report in node_reports.items() + } # Audit #4: resolve the worker class from scan_type so # graybox-specific aggregation fields (graybox_results, # completed_tests, aborted/abort_reason/abort_phase) merge @@ -387,7 +391,7 @@ def maybe_finalize_pass(owner): worker_metas = {} for addr, report in node_reports.items(): - nr_findings = owner._count_all_findings(report) + nr_findings, finding_counts, finding_signatures = worker_finding_summaries[addr] worker_metas[addr] = WorkerReportMeta( report_cid=workers[addr].get("report_cid", ""), start_port=report.get("start_port", 0), @@ -396,6 +400,10 @@ def maybe_finalize_pass(owner): open_ports=report.get("open_ports", []), nr_findings=nr_findings, node_ip=report.get("node_ip", ""), + country=(report.get("country_code") or "").upper(), + finding_counts=finding_counts, + finding_signatures=finding_signatures, + response_evidence=report.get("response_evidence"), ).to_dict() aggregated_report_cid = None diff --git a/extensions/business/cybersec/red_mesh/services/launch.py b/extensions/business/cybersec/red_mesh/services/launch.py index 9832d63c0..d4c9969b8 100644 --- a/extensions/business/cybersec/red_mesh/services/launch.py +++ b/extensions/business/cybersec/red_mesh/services/launch.py @@ -20,6 +20,7 @@ def _launch_network_jobs( end_port, job_config, nr_local_workers_override=None, + target_ports=None, ): exceptions = job_config.get("exceptions", []) if not isinstance(exceptions, list): @@ -32,6 +33,10 @@ def _launch_network_jobs( ics_safe_mode = job_config.get("ics_safe_mode", owner.cfg_ics_safe_mode) scanner_identity = job_config.get("scanner_identity", owner.cfg_scanner_identity) scanner_user_agent = job_config.get("scanner_user_agent", owner.cfg_scanner_user_agent) + timeout_profile = job_config.get("timeout_profile") + # The comparison tier is node-wide evidence, not per-thread work: it goes to + # a single worker so the tier is probed exactly once from this vantage. + comparison_ports = job_config.get("comparison_ports") or [] workers_from_spec = job_config.get("nr_local_workers") if nr_local_workers_override is not None: workers_requested = nr_local_workers_override @@ -42,7 +47,13 @@ def _launch_network_jobs( owner.P("Using {} local workers for job {}".format(workers_requested, job_id)) - ports = list(range(start_port, end_port + 1)) + # Comparison mode supplies an explicit, possibly non-contiguous port list + # (mirrored comparison tier + this node's coverage slice). Otherwise scan the + # contiguous assigned range. + if target_ports: + ports = [int(p) for p in target_ports] + else: + ports = list(range(start_port, end_port + 1)) batches = [] if port_order == PORT_ORDER_SEQUENTIAL: ports = sorted(ports) @@ -89,6 +100,8 @@ def _launch_network_jobs( ics_safe_mode=ics_safe_mode, scanner_identity=scanner_identity, scanner_user_agent=scanner_user_agent, + timeout_profile=timeout_profile, + comparison_ports=comparison_ports if index == 0 else None, ) batch_job.start() local_jobs[batch_job.local_worker_id] = batch_job @@ -139,6 +152,7 @@ def launch_local_jobs( end_port, job_config, nr_local_workers_override=None, + target_ports=None, ): strategy = get_scan_strategy(job_config.get("scan_type", ScanType.NETWORK.value)) if strategy.scan_type == ScanType.WEBAPP: @@ -160,4 +174,5 @@ def launch_local_jobs( end_port=end_port, job_config=job_config, nr_local_workers_override=nr_local_workers_override, + target_ports=target_ports, ) diff --git a/extensions/business/cybersec/red_mesh/services/launch_api.py b/extensions/business/cybersec/red_mesh/services/launch_api.py index a9abf482a..ed4d0d518 100644 --- a/extensions/business/cybersec/red_mesh/services/launch_api.py +++ b/extensions/business/cybersec/red_mesh/services/launch_api.py @@ -2,8 +2,11 @@ from urllib.parse import urlparse from ..constants import ( + COMMON_PORTS, + COMPARISON_GRAYBOX_BUNDLE_FEATURE_IDS, DISTRIBUTION_MIRROR, DISTRIBUTION_SLICE, + FEATURE_CATALOG, JOB_STATUS_RUNNING, LOCAL_WORKERS_MAX, LOCAL_WORKERS_MIN, @@ -11,6 +14,8 @@ PORT_ORDER_SHUFFLE, RUN_MODE_CONTINUOUS_MONITORING, RUN_MODE_SINGLEPASS, + TIMEOUT_PROFILES, + TIMEOUT_PROFILE_STANDARD, ScanType, ) from ..models import ( @@ -101,6 +106,15 @@ def validation_error(message: str): return {"error": "validation_error", "message": message} +def normalize_network_timeout_profile(value): + normalized = str(value or TIMEOUT_PROFILE_STANDARD).strip().upper() + if normalized not in TIMEOUT_PROFILES: + return None, validation_error( + "timeout_profile must be STANDARD or THOROUGH for network scans" + ) + return normalized, None + + def _parse_confirmation_ids(value): if value in (None, ""): return set(), None @@ -828,12 +842,82 @@ def required_unsafe_confirmation_ids( return required -def build_network_workers(owner, active_peers, start_port, end_port, distribution_strategy): +def comparison_bundle_methods(): + """Expand COMPARISON_GRAYBOX_BUNDLE_FEATURE_IDS to their probe method names.""" + ids = set(COMPARISON_GRAYBOX_BUNDLE_FEATURE_IDS) + methods = [] + for item in FEATURE_CATALOG: + if item.get("id") in ids: + methods.extend(item.get("methods", [])) + return methods + + +def compute_comparison_port_tier(start_port, end_port, full_mirror=False): + """Ports mirrored to every node in comparison mode (compared across countries). + + The comparison tier is the set every node scans, so it is what can be compared + across countries. By default (SLICE) it is the standard ``COMMON_PORTS`` bundle; + the operator's chosen range is NOT mirrored — it is sliced across nodes for + coverage. With ``full_mirror`` (the operator chose MIRROR) the whole chosen + range is added to the tier, so every port is compared, at N x the work. + Returns a sorted list of valid ports. + """ + tier = set(COMMON_PORTS) + if full_mirror: + tier |= set(range(start_port, end_port + 1)) + return sorted(p for p in tier if 1 <= p <= 65535) + + +def build_comparison_workers(active_peers, start_port, end_port, full_mirror=False): + """Comparison-mode assignment: mirror the comparison tier + slice the rest. + + Every node scans the same comparison tier (for cross-country comparison). Under + SLICE (default) the tier is just the standard ``COMMON_PORTS``, and the + operator's chosen range is split across nodes for efficient coverage; under + ``full_mirror`` (MIRROR) the tier is the whole range, so every node scans an + identical full set and there is no coverage slice. Each worker carries an + explicit, possibly non-contiguous ``target_ports`` list. + """ + comparison_tier = compute_comparison_port_tier(start_port, end_port, full_mirror=full_mirror) + comparison_set = set(comparison_tier) + # Coverage ports = the operator's chosen range minus whatever is already + # mirrored in the tier. Empty under full mirror (the whole range is the tier). + if full_mirror: + coverage_ports = [] + else: + coverage_ports = [p for p in range(start_port, end_port + 1) if p not in comparison_set] + + num_workers = len(active_peers) + base_count = len(coverage_ports) // num_workers + rem_count = len(coverage_ports) % num_workers + workers = {} + idx = 0 + for i, address in enumerate(active_peers): + size = base_count + 1 if i < rem_count else base_count + node_slice = coverage_ports[idx:idx + size] + idx += size + target_ports = sorted(comparison_set | set(node_slice)) + workers[address] = { + "start_port": target_ports[0], + "end_port": target_ports[-1], + "target_ports": target_ports, + "finished": False, + "result": None, + } + return workers + + +def build_network_workers(owner, active_peers, start_port, end_port, distribution_strategy, + comparison_mode=False): """Build peer assignments for network scans.""" num_workers = len(active_peers) if num_workers == 0: return None, validation_error("No workers available for job execution.") + if comparison_mode: + full_mirror = distribution_strategy == DISTRIBUTION_MIRROR + return build_comparison_workers(active_peers, start_port, end_port, full_mirror=full_mirror), None + workers = {} if distribution_strategy == DISTRIBUTION_MIRROR: for address in active_peers: @@ -957,14 +1041,37 @@ def announce_launch( gateway_bearer_refresh_token="", target_config_secrets=None, blockchain_attestation_enabled=False, + comparison_mode=False, + timeout_profile=TIMEOUT_PROFILE_STANDARD, ): """Persist immutable config, announce job in CStore, and return launch response.""" + comparison_mode = bool(comparison_mode) excluded_features, enabled_features = resolve_enabled_features( owner, excluded_features, scan_type=scan_type, ) + # Comparison mode: mirror the graybox tests so every node runs the same webapp + # checks (comparable across countries) and force-enable the standard bundle. + comparison_ports = None + if comparison_mode: + if scan_type == ScanType.WEBAPP.value: + graybox_assignment_strategy = GRAYBOX_ASSIGNMENT_MIRROR + bundle_methods = comparison_bundle_methods() + all_features = set(owner._get_all_features(scan_type=scan_type)) + enable = [m for m in bundle_methods if m in all_features] + excluded_features = [f for f in excluded_features if f not in enable] + enabled_features = sorted(set(enabled_features) | set(enable)) + else: + # Network scans: record the mirrored comparison tier (ports scanned from + # every node) so the report can distinguish compared vs coverage ports. + # Under MIRROR the whole range is the tier; under SLICE it is bounded. + comparison_ports = compute_comparison_port_tier( + start_port, end_port, + full_mirror=(distribution_strategy == DISTRIBUTION_MIRROR), + ) + if not scanner_identity: scanner_identity = owner.cfg_scanner_identity if not scanner_user_agent: @@ -985,6 +1092,7 @@ def announce_launch( enabled_features=enabled_features, excluded_features=excluded_features, run_mode=run_mode, + timeout_profile=timeout_profile, scan_min_delay=scan_min_delay, scan_max_delay=scan_max_delay, ics_safe_mode=ics_safe_mode, @@ -995,6 +1103,8 @@ def announce_launch( task_description=task_description, monitor_interval=monitor_interval, selected_peers=active_peers, + comparison_mode=comparison_mode, + comparison_ports=comparison_ports, created_by_name=created_by_name or "", created_by_id=created_by_id or "", authorized=True, @@ -1215,11 +1325,17 @@ def launch_network_scan( authorization=None, unsafe_launch_confirmations=None, blockchain_attestation_enabled=False, + comparison_mode=False, + timeout_profile=TIMEOUT_PROFILE_STANDARD, ): """Launch a network scan using network-specific validation and worker slicing.""" if not target: return validation_error("target required for network scan") + comparison_mode = bool(comparison_mode) + timeout_profile, timeout_profile_error = normalize_network_timeout_profile(timeout_profile) + if timeout_profile_error: + return timeout_profile_error start_port = int(start_port) end_port = int(end_port) if start_port > end_port: @@ -1243,6 +1359,9 @@ def launch_network_scan( ) if "error" in options: return options + # Comparison mode keeps the operator's MIRROR/SLICE choice: MIRROR mirrors the + # whole range to every node (full comparison); SLICE uses the tiered scheme + # (mirror the comparison tier, slice the bulk). See build_comparison_workers. required_confirmation_ids = required_unsafe_confirmation_ids( scan_type=ScanType.NETWORK.value, options=options, @@ -1303,6 +1422,7 @@ def launch_network_scan( start_port, end_port, options["distribution_strategy"], + comparison_mode=comparison_mode, ) if worker_error: return worker_error @@ -1353,6 +1473,8 @@ def launch_network_scan( roe=typed_context["roe"], authorization=typed_context["authorization"], blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, + timeout_profile=timeout_profile, ) @@ -1415,6 +1537,7 @@ def launch_webapp_scan( allow_mirror_per_worker_budget=False, unsafe_launch_confirmations=None, blockchain_attestation_enabled=False, + comparison_mode=False, ): """Launch a graybox webapp scan using webapp-specific validation and mirrored worker assignment. @@ -1669,6 +1792,7 @@ def launch_webapp_scan( gateway_bearer_refresh_token=gateway_bearer_refresh_token, target_config_secrets=target_config_secrets, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, ) @@ -1733,6 +1857,8 @@ def launch_test( authorization=None, unsafe_launch_confirmations=None, blockchain_attestation_enabled=False, + comparison_mode=False, + timeout_profile=TIMEOUT_PROFILE_STANDARD, ): """Compatibility shim that routes to scan-type-specific launch endpoints.""" try: @@ -1792,6 +1918,7 @@ def launch_test( authorization=authorization, unsafe_launch_confirmations=unsafe_launch_confirmations, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, ) return owner.launch_network_scan( @@ -1827,4 +1954,6 @@ def launch_test( authorization=authorization, unsafe_launch_confirmations=unsafe_launch_confirmations, blockchain_attestation_enabled=blockchain_attestation_enabled, + comparison_mode=comparison_mode, + timeout_profile=timeout_profile, ) diff --git a/extensions/business/cybersec/red_mesh/tests/test_api.py b/extensions/business/cybersec/red_mesh/tests/test_api.py index 532e2328b..f188b60d6 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_api.py +++ b/extensions/business/cybersec/red_mesh/tests/test_api.py @@ -98,6 +98,25 @@ def test_config_strip_none(self): d = config.to_dict() self.assertNotIn("selected_peers", d) + def test_worker_report_meta_optional_finding_evidence_roundtrip(self): + from extensions.business.cybersec.red_mesh.models import WorkerReportMeta + + legacy = WorkerReportMeta.from_dict({ + "report_cid": "cid-legacy", "start_port": 1, "end_port": 10, + }) + self.assertIsNone(legacy.finding_counts) + self.assertIsNone(legacy.finding_signatures) + self.assertNotIn("finding_counts", legacy.to_dict()) + self.assertNotIn("finding_signatures", legacy.to_dict()) + + current = WorkerReportMeta( + report_cid="cid-current", start_port=1, end_port=10, nr_findings=2, + finding_counts={"HIGH": 2}, finding_signatures=["type-a"], + ) + restored = WorkerReportMeta.from_dict(current.to_dict()) + self.assertEqual(restored.finding_counts, {"HIGH": 2}) + self.assertEqual(restored.finding_signatures, ["type-a"]) + @classmethod def _mock_plugin_modules(cls): mock_plugin_modules() @@ -247,6 +266,20 @@ def test_launch_builds_job_config_and_stores_cid(self): self.assertIsNotNone(job_specs, "Expected chainstore_hset call for job_specs") self.assertEqual(job_specs["job_config_cid"], "QmFakeConfigCID123") + def test_network_timeout_profile_is_validated_and_persisted(self): + plugin = self._build_mock_plugin(job_id="test-job-thorough") + + result = self._launch_network(plugin, timeout_profile="thorough") + + self.assertNotIn("error", result) + self.assertEqual(self._latest_job_config(plugin)["timeout_profile"], "THOROUGH") + + invalid_plugin = self._build_mock_plugin(job_id="test-job-invalid-timeout") + invalid = self._launch_network(invalid_plugin, timeout_profile="patient") + self.assertEqual(invalid["error"], "validation_error") + self.assertIn("STANDARD or THOROUGH", invalid["message"]) + invalid_plugin.r1fs.add_json.assert_not_called() + def test_cstore_has_no_static_config(self): """After launch, CStore object has no exceptions, distribution_strategy, etc.""" plugin = self._build_mock_plugin(job_id="test-job-2") @@ -1801,6 +1834,7 @@ def fake_add_json(data, show_logs=True): Plugin = self._get_plugin_class() plugin._count_nested_findings = lambda section: Plugin._count_nested_findings(section) plugin._count_all_findings = lambda report: Plugin._count_all_findings(plugin, report) + plugin._summarize_worker_findings = lambda report: Plugin._summarize_worker_findings(report) return plugin, job_specs @@ -1939,8 +1973,8 @@ def test_pass_report_cid_in_r1fs(self): self.assertIn("date_started", pass_report_dict) self.assertIn("date_completed", pass_report_dict) - def test_pass_report_worker_meta_counts_graybox_findings(self): - """WorkerReportMeta.nr_findings includes graybox findings.""" + def test_pass_report_worker_meta_counts_findings_before_aggregation(self): + """WorkerReportMeta evidence is captured before aggregate mutation/dedup.""" PentesterApi01Plugin = self._get_plugin_class() plugin, job_specs = self._build_finalize_plugin() @@ -1957,11 +1991,20 @@ def test_pass_report_worker_meta_counts_graybox_findings(self): correlation_findings=[{"title": "corr"}], ) plugin._collect_node_reports = MagicMock(return_value={"worker-A": report_a}) - plugin._get_aggregated_report = MagicMock(return_value={ - "open_ports": [443], "service_info": {}, "web_tests_info": {}, - "completed_tests": [], "ports_scanned": 512, "nr_open_ports": 1, - "port_protocols": {"443": "https"}, "graybox_results": report_a["graybox_results"], - }) + def aggregate_after_mutating_raw_reports(*_args, **_kwargs): + for section_name in ("service_info", "web_tests_info", "graybox_results"): + for port_entry in report_a[section_name].values(): + for probe_entry in port_entry.values(): + if isinstance(probe_entry, dict): + probe_entry["findings"] = [] + report_a["correlation_findings"] = [] + return { + "open_ports": [443], "service_info": {}, "web_tests_info": {}, + "completed_tests": [], "ports_scanned": 512, "nr_open_ports": 1, + "port_protocols": {"443": "https"}, "graybox_results": {}, + } + + plugin._get_aggregated_report = MagicMock(side_effect=aggregate_after_mutating_raw_reports) plugin._normalize_job_record = MagicMock(return_value=(job_specs["job_id"], job_specs)) plugin._get_job_config = MagicMock(return_value={"target": "example.com", "scan_type": "webapp"}) plugin._compute_risk_and_findings = MagicMock(return_value=({"score": 10, "breakdown": {"findings_score": 5}}, [])) @@ -1972,7 +2015,10 @@ def test_pass_report_worker_meta_counts_graybox_findings(self): PentesterApi01Plugin._maybe_finalize_pass(plugin) pass_report_dict = plugin.r1fs.add_json.call_args_list[1][0][0] - self.assertEqual(pass_report_dict["worker_reports"]["worker-A"]["nr_findings"], 5) + worker_meta = pass_report_dict["worker_reports"]["worker-A"] + self.assertEqual(worker_meta["nr_findings"], 5) + self.assertEqual(worker_meta["finding_counts"], {"INFO": 5}) + self.assertEqual(len(worker_meta["finding_signatures"]), 5) def test_aggregated_report_separate_cid(self): """aggregated_report_cid is a separate R1FS write from the PassReport.""" @@ -4948,6 +4994,7 @@ def test_close_job_audit_counts_graybox_findings(self): plugin._log_audit_event = MagicMock() plugin._count_nested_findings = lambda section: Plugin._count_nested_findings(section) plugin._count_all_findings = lambda report: Plugin._count_all_findings(plugin, report) + plugin._summarize_worker_findings = lambda report: Plugin._summarize_worker_findings(report) report = { "start_port": 443, diff --git a/extensions/business/cybersec/red_mesh/tests/test_connection_metrics.py b/extensions/business/cybersec/red_mesh/tests/test_connection_metrics.py new file mode 100644 index 000000000..705c1f682 --- /dev/null +++ b/extensions/business/cybersec/red_mesh/tests/test_connection_metrics.py @@ -0,0 +1,165 @@ +import unittest + +from extensions.business.cybersec.red_mesh.connection_metrics import ( + detect_connection_signals, + merge_connection_windows, +) +from extensions.business.cybersec.red_mesh.mixins.live_progress import _LiveProgressMixin +from extensions.business.cybersec.red_mesh.worker.metrics_collector import MetricsCollector + + +def _window(index, responsive_count, attempts=5, success_rate=None): + if success_rate is None: + success_rate = responsive_count / attempts + return { + "window_start": index * 60.0, + "window_end": (index + 1) * 60.0, + "success_rate": success_rate, + "attempts": attempts, + "responsive_count": responsive_count, + "response_rate": responsive_count / attempts, + } + + +class TestConnectionWindows(unittest.TestCase): + + def test_refused_and_reset_are_responsive_but_not_connected(self): + collector = MetricsCollector() + collector._connection_log = [ + (100.0, "connected"), + (101.0, "refused"), + (102.0, "reset"), + (103.0, "timeout"), + (104.0, "error"), + ] + + windows = collector._compute_success_windows() + + self.assertEqual(windows, [{ + "window_start": 0.0, + "window_end": 60.0, + "success_rate": 0.2, + "attempts": 5, + "responsive_count": 3, + "response_rate": 0.6, + }]) + self.assertFalse(detect_connection_signals(windows)["blocking_detected"]) + + def test_single_connection_still_produces_a_window(self): + collector = MetricsCollector() + collector._connection_log = [(100.0, "refused")] + + self.assertEqual(collector._compute_success_windows()[0]["attempts"], 1) + + def test_closed_ports_remain_responsive_across_windows(self): + collector = MetricsCollector() + collector._connection_log = ( + [(100.0 + index, "refused") for index in range(5)] + + [(160.0 + index, "reset") for index in range(5)] + ) + + windows = collector._compute_success_windows() + signals = detect_connection_signals(windows) + + self.assertEqual([window["success_rate"] for window in windows], [0.0, 0.0]) + self.assertEqual([window["response_rate"] for window in windows], [1.0, 1.0]) + self.assertFalse(signals["blocking_detected"]) + + def test_blocking_requires_qualified_high_to_low_response_transition(self): + windows = [ + _window(0, 4), + _window(1, 0, attempts=2), # Unqualified evidence is ignored. + _window(2, 1), + ] + + signals = detect_connection_signals(windows) + + self.assertTrue(signals["blocking_detected"]) + self.assertFalse(signals["rate_limiting_detected"]) + + def test_insufficient_samples_never_raise_a_signal(self): + windows = [ + _window(0, 4, attempts=4), + _window(1, 0, attempts=4), + _window(2, 4, attempts=4), + _window(3, 1, attempts=4), + ] + + self.assertEqual(detect_connection_signals(windows), { + "rate_limiting_detected": False, + "blocking_detected": False, + }) + + def test_optional_response_rate_can_supply_missing_responsive_count(self): + windows = [ + {"attempts": 5, "response_rate": 0.8}, + {"attempts": 5, "response_rate": 0.2}, + ] + + self.assertTrue(detect_connection_signals(windows)["blocking_detected"]) + + def test_throttling_uses_attempt_weighted_first_and_last_pairs(self): + windows = [ + _window(0, 5, attempts=5), + _window(1, 15, attempts=15), + _window(2, 5, attempts=5), + _window(3, 40, attempts=100), + ] + + signals = detect_connection_signals(windows) + + self.assertTrue(signals["rate_limiting_detected"]) + self.assertFalse(signals["blocking_detected"]) + + def test_throttling_requires_four_qualified_windows(self): + windows = [_window(0, 5), _window(1, 5), _window(2, 3)] + + self.assertFalse(detect_connection_signals(windows)["rate_limiting_detected"]) + + def test_blocking_threshold_is_not_reported_as_throttling(self): + windows = [_window(0, 5), _window(1, 5), _window(2, 1), _window(3, 1)] + + signals = detect_connection_signals(windows) + + self.assertTrue(signals["blocking_detected"]) + self.assertFalse(signals["rate_limiting_detected"]) + + +class TestConnectionWindowMerge(unittest.TestCase): + + def test_aligned_windows_sum_counts_and_recompute_signals(self): + worker_one = { + "success_rate_over_time": [_window(0, 5), _window(1, 0)], + "blocking_detected": True, + } + worker_two = { + "success_rate_over_time": [_window(0, 0), _window(1, 5)], + "blocking_detected": False, + } + + merged = _LiveProgressMixin._merge_worker_metrics([worker_one, worker_two]) + windows = merged["success_rate_over_time"] + + self.assertEqual([window["attempts"] for window in windows], [10, 10]) + self.assertEqual([window["responsive_count"] for window in windows], [5, 5]) + self.assertEqual([window["response_rate"] for window in windows], [0.5, 0.5]) + self.assertFalse(merged["blocking_detected"]) + self.assertFalse(merged["rate_limiting_detected"]) + + def test_legacy_windows_remain_visible_but_cannot_raise_a_signal(self): + legacy_windows = [ + {"window_start": 0.0, "window_end": 60.0, "success_rate": 1.0}, + {"window_start": 60.0, "window_end": 120.0, "success_rate": 0.0}, + ] + + merged = merge_connection_windows([{ + "success_rate_over_time": legacy_windows, + "blocking_detected": True, + }]) + + self.assertEqual(merged, legacy_windows) + self.assertFalse(detect_connection_signals(merged)["blocking_detected"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/extensions/business/cybersec/red_mesh/tests/test_finalization_aggregation.py b/extensions/business/cybersec/red_mesh/tests/test_finalization_aggregation.py index 88facaace..e33a2ed77 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_finalization_aggregation.py +++ b/extensions/business/cybersec/red_mesh/tests/test_finalization_aggregation.py @@ -131,6 +131,42 @@ def test_stamping_is_idempotent(self): self.assertEqual(f["_source_node_addr"], "0xaddr") +class TestWorkerFindingEvidence(unittest.TestCase): + + def test_raw_record_counts_and_unique_signatures_precede_cross_worker_dedup(self): + host = _Host() + findings = [ + { + "finding_signature": f"type-{index % 188}", + "severity": "HIGH" if index < 28 else "MEDIUM", + "title": f"record-{index}", + "_source_node_addr": "0xworker-a", + } + for index in range(228) + ] + report = { + "service_info": {"443": {"probe": {"findings": findings}}}, + } + + nr_findings, counts, signatures = host._summarize_worker_findings(report) + + self.assertEqual(nr_findings, 228) + self.assertEqual(counts, {"HIGH": 28, "MEDIUM": 200}) + self.assertEqual(len(signatures), 188) + self.assertEqual(signatures[0], "type-0") + + def test_fallback_signature_ignores_worker_attribution(self): + host = _Host() + base = {"title": "Same issue", "severity": "LOW", "port": 80} + report_a = {"findings": [{**base, "_source_node_addr": "0xa"}]} + report_b = {"findings": [{**base, "_source_node_addr": "0xb"}]} + + self.assertEqual( + host._summarize_worker_findings(report_a)[2], + host._summarize_worker_findings(report_b)[2], + ) + + class TestGrayboxMultiWorkerAggregation(unittest.TestCase): def test_graybox_results_merge_across_workers(self): @@ -388,5 +424,116 @@ def test_network_aggregation_still_works_without_worker_cls(self): self.assertIn("443", agg["service_info"]) +class _AggHost(_Host): + """_Host plus the service-count helper _compute_ui_aggregate depends on.""" + + def _count_services(self, service_info): + return len(service_info or {}) + + +class TestOriginCountryAndComparisonAggregate(unittest.TestCase): + """Origin-country breakdown and comparison-mode per-node vantage comparison.""" + + def _latest_pass(self): + return { + "pass_nr": 1, + "findings": [ + {"finding_signature": "sig1", "severity": "HIGH", "title": "XSS", + "port": 443, "_source_node_addr": "0xUS"}, + ], + "worker_reports": { + "0xUS": {"start_port": 1, "end_port": 443, "open_ports": [80, 443], + "node_ip": "1.1.1.1", "country": "US", "nr_findings": 1, + "finding_counts": {"HIGH": 1}, "finding_signatures": ["sig1"]}, + "0xIN": {"start_port": 1, "end_port": 443, "open_ports": [80], + "node_ip": "2.2.2.2", "country": "in", "nr_findings": 0}, + }, + "worker_scan_metrics": { + "0xUS": {"scan_metrics": { + "connection_outcomes": {"connected": 5, "timeout": 0, "refused": 2, "reset": 1}, + "response_times": {"p95": 0.12}, + "coverage": 0.75, + "probes_attempted": 4, + "probes_completed": 3, + "probes_failed": 1, + "phase_durations": {"port_scan": 12.5}, + "total_duration": 19.0, + "success_rate_over_time": [{ + "window_start": 0, "window_end": 60, "success_rate": 0.625, + "attempts": 8, "responsive_count": 8, "response_rate": 1.0, + }], + }, "threads": [{"local_worker_id": "thread-1"}]}, + "0xBR": {"scan_metrics": {"connection_outcomes": {"connected": 0, "timeout": 9}, + "response_times": {"p95": 2.0}}}, + }, + } + + def test_compact_fallback_signature_uses_cross_client_canonical_json(self): + report = { + "findings": [{ + "title": "TLS issue", "severity": "HIGH", "port": 443, + "cvss_score": 7.0, "negative_zero": -0.0, "not_a_number": float("nan"), + "evidence": {"z": True, "a": "é", "_source_worker_id": "thread-a"}, + "_source_node_addr": "0xUS", + }], + } + + _, _, signatures = _AggHost()._summarize_worker_findings(report) + + self.assertEqual(signatures, [ + "sha256:34f703bc05595ea95796eda463b738b5e7bca69b2a6a611575d2f22ffcee8f8d", + ]) + + def test_country_breakdown_and_per_worker_country(self): + host = _AggHost() + ui = host._compute_ui_aggregate( + [self._latest_pass()], + {"open_ports": [80, 443], "service_info": {}}, + job_config={"scan_type": "network"}, + ) + d = ui.to_dict() + # Counts per ISO-2 (uppercased), sorted by (-count, code) — tie sorts IN before US. + self.assertEqual(d.get("country_breakdown"), [{"code": "IN", "count": 1}, {"code": "US", "count": 1}]) + by_id = {w["id"]: w["country"] for w in d["worker_activity"]} + self.assertEqual(by_id["0xUS"], "US") + self.assertEqual(by_id["0xIN"], "IN") + # node_comparison is only computed in comparison mode. + self.assertNotIn("node_comparison", d) + + def test_node_comparison_includes_failed_and_timed_out_nodes(self): + host = _AggHost() + cfg = { + "scan_type": "network", + "comparison_mode": True, + "selected_peers": ["0xUS", "0xIN", "0xBR", "0xCN"], + "comparison_ports": [80, 443], + } + ui = host._compute_ui_aggregate( + [self._latest_pass()], + {"open_ports": [80, 443], "service_info": {}}, + job_config=cfg, + ) + comp = {e["address"]: e for e in ui.to_dict()["node_comparison"]} + # Reached node: open ports + finding attribution + latency. + self.assertEqual(comp["0xUS"]["status"], "reached") + self.assertEqual(comp["0xUS"]["open_ports"], [80, 443]) + self.assertEqual(comp["0xUS"]["metrics"]["response_p95_ms"], 120.0) + self.assertEqual(comp["0xUS"]["findings"][0]["signature"], "sig1") + self.assertEqual(comp["0xUS"]["finding_counts"], {"HIGH": 1}) + self.assertEqual(comp["0xUS"]["finding_signatures"], ["sig1"]) + self.assertEqual(comp["0xUS"]["metrics"]["refused"], 2) + self.assertEqual(comp["0xUS"]["metrics"]["reset"], 1) + self.assertEqual(comp["0xUS"]["metrics"]["coverage"], 0.75) + self.assertEqual(comp["0xUS"]["metrics"]["probes_completed"], 3) + self.assertEqual(comp["0xUS"]["metrics"]["total_duration"], 19.0) + self.assertEqual(comp["0xUS"]["metrics"]["traffic_windows"][0]["attempts"], 8) + self.assertEqual(comp["0xUS"]["metrics"]["threads"][0]["local_worker_id"], "thread-1") + # Metrics-only node with all-timeout connections -> timeout. + self.assertEqual(comp["0xBR"]["status"], "timeout") + # Selected peer that never reported at all -> failed (China-timeout case). + self.assertEqual(comp["0xCN"]["status"], "failed") + self.assertEqual(comp["0xCN"]["country"], "UN") + + if __name__ == '__main__': unittest.main() diff --git a/extensions/business/cybersec/red_mesh/tests/test_integration.py b/extensions/business/cybersec/red_mesh/tests/test_integration.py index 14d705628..71274dee5 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_integration.py +++ b/extensions/business/cybersec/red_mesh/tests/test_integration.py @@ -2657,7 +2657,7 @@ def test_scan_metrics_strip_none(self): self.assertNotIn("probe_breakdown", d) def test_merge_worker_metrics(self): - """_merge_worker_metrics sums outcomes, coverage, findings; maxes duration; ORs flags.""" + """_merge_worker_metrics sums outcomes, coverage, findings, and recomputes flags.""" mock_plugin_modules() from extensions.business.cybersec.red_mesh.pentester_api_01 import PentesterApi01Plugin m1 = { @@ -2734,8 +2734,8 @@ def test_merge_worker_metrics(self): self.assertEqual(rt["p99"], 0.7) # max of per-thread p99 # Max duration self.assertEqual(merged["total_duration"], 75.0) - # OR flags - self.assertTrue(merged["rate_limiting_detected"]) + # Flags without count-bearing windows are not treated as verified evidence. + self.assertFalse(merged["rate_limiting_detected"]) self.assertFalse(merged["blocking_detected"]) # Open port details: deduplicated by port, sorted opd = merged["open_port_details"] @@ -2818,8 +2818,8 @@ def capture_add_json(data, show_logs=False): # Probes summed self.assertEqual(sm["probes_attempted"], 4) self.assertEqual(sm["probes_completed"], 3) - # OR flags - self.assertTrue(sm["rate_limiting_detected"]) + # Flags are recomputed from windows, not ORed from thread booleans. + self.assertFalse(sm["rate_limiting_detected"]) live_writes = [ call.kwargs["value"] @@ -2906,6 +2906,7 @@ def capture_add_json(data, show_logs=False): plugin.r1fs.add_json.side_effect = capture_add_json plugin._compute_risk_and_findings = MagicMock(return_value=({"score": 25, "breakdown": {}}, [])) + plugin._summarize_worker_findings = MagicMock(return_value=(0, {}, [])) plugin._get_job_config = MagicMock(return_value={}) plugin._submit_redmesh_test_attestation = MagicMock(return_value=None) plugin._build_job_archive = MagicMock() @@ -2932,6 +2933,6 @@ def capture_add_json(data, show_logs=False): self.assertEqual(sm["probes_attempted"], 6) self.assertEqual(sm["probes_completed"], 5) self.assertEqual(sm["probes_failed"], 1) - # OR flags + # Flags are recomputed from windows, not ORed from node booleans. self.assertFalse(sm["rate_limiting_detected"]) - self.assertTrue(sm["blocking_detected"]) + self.assertFalse(sm["blocking_detected"]) diff --git a/extensions/business/cybersec/red_mesh/tests/test_launch_service.py b/extensions/business/cybersec/red_mesh/tests/test_launch_service.py index 39491f7aa..1c51e79c0 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_launch_service.py +++ b/extensions/business/cybersec/red_mesh/tests/test_launch_service.py @@ -119,3 +119,129 @@ def test_launch_local_jobs_uses_webapp_strategy_dispatch(self): self.assertTrue(worker.started) self.assertEqual(worker.target_url, "https://example.com/app") self.assertEqual(worker.job_config.scan_type, "webapp") + + def test_explicit_target_ports_override_contiguous_range(self): + """Comparison mode supplies an explicit, non-contiguous port list.""" + owner = DummyOwner() + strategy = ScanStrategy( + scan_type=ScanType.NETWORK, + worker_cls=DummyNetworkWorker, + catalog_categories=("service",), + ) + with patch("extensions.business.cybersec.red_mesh.services.launch.get_scan_strategy", return_value=strategy): + local_jobs = launch_local_jobs( + owner, + job_id="job-cmp", + target="10.0.0.10", + launcher="0xlauncher", + start_port=1, + end_port=2, # ignored when target_ports is provided + job_config={ + "scan_type": "network", + "nr_local_workers": 1, + "port_order": PORT_ORDER_SEQUENTIAL, + }, + target_ports=[22, 443, 8080], + ) + scanned = sorted( + p for worker in local_jobs.values() for p in worker.worker_target_ports + ) + self.assertEqual(scanned, [22, 443, 8080]) + + def test_network_timeout_profile_reaches_each_local_worker(self): + owner = DummyOwner() + strategy = ScanStrategy( + scan_type=ScanType.NETWORK, + worker_cls=DummyNetworkWorker, + catalog_categories=("service",), + ) + with patch("extensions.business.cybersec.red_mesh.services.launch.get_scan_strategy", return_value=strategy): + local_jobs = launch_local_jobs( + owner, + job_id="job-thorough", + target="10.0.0.10", + launcher="0xlauncher", + start_port=80, + end_port=81, + job_config={ + "scan_type": "network", + "nr_local_workers": 2, + "port_order": PORT_ORDER_SEQUENTIAL, + "timeout_profile": "THOROUGH", + }, + ) + + self.assertEqual( + {worker.kwargs["timeout_profile"] for worker in local_jobs.values()}, + {"THOROUGH"}, + ) + + +class TestComparisonTieredAssignment(unittest.TestCase): + """Tiered mirror+slice port assignment for geographic comparison mode.""" + + def test_slice_mirrors_common_ports_and_splits_the_chosen_range(self): + """SLICE (default): only COMMON_PORTS are mirrored/compared; the operator's + chosen range is split across nodes for coverage (not mirrored).""" + from extensions.business.cybersec.red_mesh.constants import COMMON_PORTS + from extensions.business.cybersec.red_mesh.services.launch_api import ( + build_comparison_workers, + compute_comparison_port_tier, + ) + # The comparison tier is exactly COMMON_PORTS — the chosen range is NOT in it. + tier = set(compute_comparison_port_tier(1, 33)) + self.assertEqual(tier, {p for p in COMMON_PORTS if 1 <= p <= 65535}) + self.assertNotIn(1, tier) + self.assertIn(443, tier) + + workers = build_comparison_workers(["0xA", "0xB", "0xC"], 1, 33) + common = set(COMMON_PORTS) + coverage_slices = [] + for w in workers.values(): + target = set(w["target_ports"]) + self.assertTrue(common.issubset(target)) # standard ports mirrored to all + coverage_slices.append(target - common) + # The chosen range (minus common ports already mirrored) is split disjointly + # across nodes and together covers 1..33 — i.e. sliced, not mirrored. + union = set().union(*coverage_slices) | common + self.assertTrue(set(range(1, 34)).issubset(union)) + for i in range(len(coverage_slices)): + for j in range(i + 1, len(coverage_slices)): + self.assertTrue(coverage_slices[i].isdisjoint(coverage_slices[j])) + # Not every node scans the same set (slicing actually happened). + self.assertGreater(len({tuple(w["target_ports"]) for w in workers.values()}), 1) + + def test_large_range_mirrors_common_ports_and_slices_bulk(self): + from extensions.business.cybersec.red_mesh.constants import COMMON_PORTS + from extensions.business.cybersec.red_mesh.services.launch_api import ( + build_comparison_workers, + ) + workers = build_comparison_workers(["0xA", "0xB", "0xC"], 1, 5000) + common = set(COMMON_PORTS) + union = set() + coverage_slices = [] + for w in workers.values(): + target = set(w["target_ports"]) + self.assertTrue(common.issubset(target)) # comparison tier mirrored to all + union |= target + coverage_slices.append(target - common) + # Coverage slices are disjoint and together cover the whole range. + self.assertTrue(set(range(1, 5001)).issubset(union)) + for i in range(len(coverage_slices)): + for j in range(i + 1, len(coverage_slices)): + self.assertTrue(coverage_slices[i].isdisjoint(coverage_slices[j])) + + def test_full_mirror_gives_every_node_the_full_range(self): + """MIRROR choice in comparison mode: every node scans the identical full + range (plus standard ports), with no coverage split.""" + from extensions.business.cybersec.red_mesh.constants import COMMON_PORTS + from extensions.business.cybersec.red_mesh.services.launch_api import ( + build_comparison_workers, + ) + workers = build_comparison_workers(["0xA", "0xB", "0xC"], 1, 5000, full_mirror=True) + expected = sorted(set(range(1, 5001)) | set(COMMON_PORTS)) + port_sets = [w["target_ports"] for w in workers.values()] + for target in port_sets: + self.assertEqual(target, expected) # identical full set on every node + # All nodes scan the same set (no disjoint coverage slices). + self.assertEqual(len({tuple(t) for t in port_sets}), 1) diff --git a/extensions/business/cybersec/red_mesh/tests/test_normalization.py b/extensions/business/cybersec/red_mesh/tests/test_normalization.py index 99ba7bde6..1e1e2bb66 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_normalization.py +++ b/extensions/business/cybersec/red_mesh/tests/test_normalization.py @@ -409,8 +409,8 @@ class MockHost(_ReportMixin): class TestFindingCounting(unittest.TestCase): - def test_count_all_findings_walks_all_sections(self): - """_count_all_findings counts service, web, correlation, and graybox findings.""" + def test_count_all_findings_walks_all_published_paths(self): + """Counting covers nested/flat service, web, graybox, correlation, and top-level paths.""" from extensions.business.cybersec.red_mesh.mixins.report import _ReportMixin class MockHost(_ReportMixin): @@ -420,15 +420,18 @@ class MockHost(_ReportMixin): report = { "service_info": { "80": { + "findings": [{"title": "legacy-flat-service"}], "_service_info_http": {"findings": [{"title": "svc-1"}, {"title": "svc-2"}]}, }, }, "web_tests_info": { "80": { + "findings": [{"title": "legacy-flat-web"}], "_web_test_xss": {"findings": [{"title": "web-1"}]}, }, }, "correlation_findings": [{"title": "corr-1"}], + "findings": [{"title": "top-1"}], "graybox_results": { "443": { "_graybox_test": {"findings": [{"title": "gb-1"}, {"title": "gb-2"}]}, @@ -436,7 +439,7 @@ class MockHost(_ReportMixin): }, } - self.assertEqual(host._count_all_findings(report), 6) + self.assertEqual(host._count_all_findings(report), 9) class TestLaunchValidation(unittest.TestCase): diff --git a/extensions/business/cybersec/red_mesh/tests/test_probes.py b/extensions/business/cybersec/red_mesh/tests/test_probes.py index 770ba0021..144f231e5 100644 --- a/extensions/business/cybersec/red_mesh/tests/test_probes.py +++ b/extensions/business/cybersec/red_mesh/tests/test_probes.py @@ -522,6 +522,31 @@ def close(self): self.assertNotIn(81, worker.state["open_ports"]) self.assertIn("scan_ports_step_completed", worker.state["completed_tests"]) + def test_port_scan_does_not_treat_unreachable_errno_as_responsive(self): + import errno + + owner, worker = self._build_worker(ports=[81]) + + class DummySocket: + def settimeout(self, timeout): + return None + + def connect_ex(self, address): + return errno.EHOSTUNREACH + + def close(self): + return None + + with patch( + "extensions.business.cybersec.red_mesh.worker.pentest_worker.socket.socket", + return_value=DummySocket(), + ): + worker._scan_ports_step() + + outcomes = worker.metrics.build().connection_outcomes + self.assertEqual(outcomes["error"], 1) + self.assertEqual(outcomes["refused"], 0) + def test_service_telnet_banner(self): owner, worker = self._build_worker(ports=[23]) diff --git a/extensions/business/cybersec/red_mesh/tests/test_response_fingerprint.py b/extensions/business/cybersec/red_mesh/tests/test_response_fingerprint.py new file mode 100644 index 000000000..695525794 --- /dev/null +++ b/extensions/business/cybersec/red_mesh/tests/test_response_fingerprint.py @@ -0,0 +1,200 @@ +""" +Per-vantage response fingerprint capture (RM-050). + +Covers the excerpt safety rules, DNS resolution semantics, the comparison-tier +scoping that keeps the phase off non-comparison jobs, and the aggregation +contract that keeps each vantage's evidence attributable to that vantage. +""" + +import unittest +from unittest.mock import MagicMock, patch + +from extensions.business.cybersec.red_mesh.worker import PentestLocalWorker +from extensions.business.cybersec.red_mesh.worker.response_fingerprint import ( + EXCERPT_MAX_BYTES, + certificate_identity, + excerpt_allowed, + normalize_content_type, + resolve_host, + sanitize_excerpt, +) +from .conftest import DummyOwner + + +def _make_worker(**overrides): + defaults = dict( + owner=DummyOwner(), + target="127.0.0.1", + job_id="test-job", + initiator="test-addr", + local_id_prefix="1", + worker_target_ports=[80, 443], + ) + defaults.update(overrides) + return PentestLocalWorker(**defaults) + + +class TestExcerptSanitization(unittest.TestCase): + + def test_returns_none_for_empty_body(self): + self.assertIsNone(sanitize_excerpt("")) + self.assertIsNone(sanitize_excerpt(None)) + + def test_caps_at_byte_limit(self): + excerpt = sanitize_excerpt("a" * 5000) + self.assertLessEqual(len(excerpt.encode("utf-8")), EXCERPT_MAX_BYTES) + + def test_truncates_on_utf8_character_boundary(self): + # Three-byte characters do not divide evenly into the byte cap, so a + # naive slice would emit a partial character. + excerpt = sanitize_excerpt("中" * 400) + self.assertLessEqual(len(excerpt.encode("utf-8")), EXCERPT_MAX_BYTES) + self.assertNotIn("�", excerpt) + excerpt.encode("utf-8").decode("utf-8") + + def test_redacts_credential_key_values(self): + excerpt = sanitize_excerpt('{"api_key": "sk-abc123def456", "ok": 1}') + self.assertNotIn("sk-abc123def456", excerpt) + self.assertIn("REDACTED", excerpt) + + def test_redacts_bearer_tokens(self): + excerpt = sanitize_excerpt("Authorization: Bearer eyJhbGciOiJIUzI1NiJ9.payload") + self.assertNotIn("eyJhbGciOiJIUzI1NiJ9", excerpt) + + def test_redacts_email_addresses(self): + excerpt = sanitize_excerpt("contact operator@example.com for access") + self.assertNotIn("operator@example.com", excerpt) + self.assertIn("[REDACTED_EMAIL]", excerpt) + + def test_redacts_long_hex_runs(self): + digest = "a" * 64 + excerpt = sanitize_excerpt(f"csrf={digest}") + self.assertNotIn(digest, excerpt) + + def test_redacts_long_base64_runs(self): + blob = "QUJDREVGR0hJSktMTU5PUFFSU1RVVldYWVoxMjM0NTY3ODkw" + excerpt = sanitize_excerpt(f"state {blob} end") + self.assertNotIn(blob, excerpt) + + def test_redaction_precedes_truncation(self): + # A credential straddling the byte cap must not survive as a prefix. + padding = "x" * (EXCERPT_MAX_BYTES - 20) + excerpt = sanitize_excerpt(f"{padding} password=supersecretvalue123") + self.assertNotIn("supersecretvalue", excerpt) + + +class TestContentTypeGate(unittest.TestCase): + + def test_allows_textual_types(self): + self.assertTrue(excerpt_allowed("text/html; charset=utf-8")) + self.assertTrue(excerpt_allowed("text/plain")) + self.assertTrue(excerpt_allowed("application/json")) + + def test_rejects_binary_and_missing_types(self): + self.assertFalse(excerpt_allowed("application/octet-stream")) + self.assertFalse(excerpt_allowed("image/png")) + self.assertFalse(excerpt_allowed("")) + self.assertFalse(excerpt_allowed(None)) + + def test_normalizes_content_type(self): + self.assertEqual(normalize_content_type("TEXT/HTML; charset=utf-8"), "text/html") + self.assertIsNone(normalize_content_type(None)) + + +class TestHostResolution(unittest.TestCase): + + def test_literal_ipv4_target_yields_no_dns_answer(self): + addresses, error = resolve_host("93.184.216.34") + self.assertEqual(addresses, []) + self.assertIsNone(error) + + def test_literal_ipv6_target_yields_no_dns_answer(self): + addresses, error = resolve_host("2001:4860:4860::8888") + self.assertEqual(addresses, []) + self.assertIsNone(error) + + def test_resolution_failure_is_recorded_not_raised(self): + with patch( + "extensions.business.cybersec.red_mesh.worker.response_fingerprint.socket.getaddrinfo", + side_effect=OSError("Name or service not known"), + ): + addresses, error = resolve_host("nonexistent.invalid") + self.assertEqual(addresses, []) + self.assertIn("Name or service not known", error) + + def test_addresses_are_sorted_and_deduplicated(self): + infos = [ + (2, 1, 6, "", ("93.184.216.34", 0)), + (2, 1, 6, "", ("93.184.216.34", 0)), + (2, 1, 6, "", ("1.2.3.4", 0)), + ] + with patch( + "extensions.business.cybersec.red_mesh.worker.response_fingerprint.socket.getaddrinfo", + return_value=infos, + ): + addresses, error = resolve_host("example.test") + self.assertEqual(addresses, ["1.2.3.4", "93.184.216.34"]) + self.assertIsNone(error) + + +class TestCertificateIdentity(unittest.TestCase): + + def test_absent_certificate_yields_no_identity(self): + self.assertIsNone(certificate_identity(None)) + self.assertIsNone(certificate_identity(b"")) + + def test_unparseable_certificate_still_fingerprints(self): + # An unparseable certificate must not silently match another one. + identity = certificate_identity(b"not-a-certificate") + self.assertIn("cert_sha256", identity) + self.assertEqual(len(identity["cert_sha256"]), 64) + + +class TestComparisonTierScoping(unittest.TestCase): + + def test_phase_skipped_without_comparison_tier(self): + worker = _make_worker() + self.assertEqual(worker.comparison_ports, []) + worker._fingerprint_port = MagicMock() + worker._capture_response_fingerprint() + worker._fingerprint_port.assert_not_called() + self.assertNotIn("response_evidence", worker.state) + + def test_phase_probes_only_the_comparison_tier(self): + worker = _make_worker(worker_target_ports=[8080, 9090], comparison_ports=[443, 80, 443]) + worker._check_stopped = MagicMock(return_value=False) + worker._fingerprint_port = MagicMock(return_value={"reachable": False}) + with patch( + "extensions.business.cybersec.red_mesh.worker.response_fingerprint.resolve_host", + return_value=([], None), + ): + worker._capture_response_fingerprint() + + probed = sorted(call.args[0] for call in worker._fingerprint_port.call_args_list) + self.assertEqual(probed, [80, 443]) + self.assertEqual(sorted(worker.state["response_evidence"]["ports"]), ["443", "80"]) + + def test_evidence_absent_from_status_without_capture(self): + worker = _make_worker() + self.assertNotIn("response_evidence", worker.get_status()) + + def test_evidence_present_in_status_after_capture(self): + worker = _make_worker(comparison_ports=[443]) + worker.state["response_evidence"] = {"target_host": "example.test", "ports": {}} + self.assertIn("response_evidence", worker.get_status()) + + +class TestAggregationAttribution(unittest.TestCase): + """ + Response evidence describes one vantage. Registering it as a cross-worker + aggregated field would deep-merge every vantage's ports into one colliding + structure — the defect RM-049 had to repair for findings. + """ + + def test_response_evidence_is_not_a_cross_worker_aggregated_field(self): + fields = PentestLocalWorker.get_worker_specific_result_fields() + self.assertNotIn("response_evidence", fields) + + +if __name__ == "__main__": + unittest.main() diff --git a/extensions/business/cybersec/red_mesh/tests/test_timeout_profile.py b/extensions/business/cybersec/red_mesh/tests/test_timeout_profile.py new file mode 100644 index 000000000..fccad1666 --- /dev/null +++ b/extensions/business/cybersec/red_mesh/tests/test_timeout_profile.py @@ -0,0 +1,94 @@ +import re +import unittest +from pathlib import Path + +from extensions.business.cybersec.red_mesh.constants import ( + TIMEOUT_PROFILE_STANDARD, + TIMEOUT_PROFILE_THOROUGH, + resolve_target_response_timeout, +) +from extensions.business.cybersec.red_mesh.models.archive import JobConfig +from extensions.business.cybersec.red_mesh.services.launch_api import ( + normalize_network_timeout_profile, +) + + +def _minimal_job_config(**overrides): + values = { + "target": "10.0.0.10", + "start_port": 1, + "end_port": 443, + "exceptions": [], + "distribution_strategy": "SLICE", + "port_order": "SEQUENTIAL", + "nr_local_workers": 1, + "enabled_features": [], + "excluded_features": [], + "run_mode": "SINGLEPASS", + } + values.update(overrides) + return values + + +class TestTimeoutProfileContract(unittest.TestCase): + def test_standard_preserves_every_supported_wait(self): + for wait in (0.3, 2, 3, 4, 5): + self.assertEqual(resolve_target_response_timeout(TIMEOUT_PROFILE_STANDARD, wait), wait) + + def test_thorough_uses_exact_capped_mapping(self): + self.assertEqual( + [resolve_target_response_timeout(TIMEOUT_PROFILE_THOROUGH, wait) for wait in (0.3, 2, 3, 4, 5)], + [0.9, 6.0, 9.0, 12.0, 15.0], + ) + + def test_launch_normalization_defaults_and_rejects_invalid_values(self): + self.assertEqual(normalize_network_timeout_profile(None), (TIMEOUT_PROFILE_STANDARD, None)) + self.assertEqual(normalize_network_timeout_profile("thorough"), (TIMEOUT_PROFILE_THOROUGH, None)) + value, error = normalize_network_timeout_profile("patient") + self.assertIsNone(value) + self.assertEqual(error["error"], "validation_error") + self.assertIn("STANDARD or THOROUGH", error["message"]) + + def test_job_config_roundtrip_and_legacy_default(self): + thorough = JobConfig(**_minimal_job_config(timeout_profile=TIMEOUT_PROFILE_THOROUGH)) + self.assertEqual(JobConfig.from_dict(thorough.to_dict()).timeout_profile, TIMEOUT_PROFILE_THOROUGH) + self.assertEqual(JobConfig.from_dict(_minimal_job_config()).timeout_profile, TIMEOUT_PROFILE_STANDARD) + + +class TestTimeoutProfileCallSiteAudit(unittest.TestCase): + """Keep ordinary network waits profiled and explicit exclusions unchanged.""" + + def test_network_worker_waits_use_resolver_except_timing_probe(self): + root = Path(__file__).resolve().parents[1] + sources = [root / "worker" / "pentest_worker.py"] + sources.extend(sorted((root / "worker" / "service").glob("*.py"))) + sources.extend(sorted((root / "worker" / "web").glob("*.py"))) + + ordinary_literal = re.compile(r"(?= 2.0:", injection_text) + + +if __name__ == "__main__": + unittest.main() diff --git a/extensions/business/cybersec/red_mesh/worker/metrics_collector.py b/extensions/business/cybersec/red_mesh/worker/metrics_collector.py index 8bbb27922..5bf42c859 100644 --- a/extensions/business/cybersec/red_mesh/worker/metrics_collector.py +++ b/extensions/business/cybersec/red_mesh/worker/metrics_collector.py @@ -1,6 +1,7 @@ import time import statistics +from ..connection_metrics import RESPONSIVE_CONNECTION_OUTCOMES, detect_connection_signals from ..models.shared import ScanMetrics @@ -24,8 +25,8 @@ def __init__(self): self._banner_confirmed = 0 self._banner_guessed = 0 self._finding_counts = {} - # For success rate over time windows - self._connection_log = [] # [(timestamp, success_bool)] + # For connection behavior over time windows. + self._connection_log = [] # [(timestamp, outcome)] # Aborts: fatal safety/policy gate failures that stop the scan. # Tracked separately from probe_failed because the abort is the # reason the scan stopped, not a per-probe outcome. @@ -45,7 +46,7 @@ def record_connection(self, outcome: str, response_time: float): self._connection_outcomes[outcome] = self._connection_outcomes.get(outcome, 0) + 1 if response_time >= 0: self._response_times.append(response_time) - self._connection_log.append((time.time(), outcome == "connected")) + self._connection_log.append((time.time(), outcome)) self._ports_scanned += 1 def record_port_scan_delay(self, delay: float): @@ -109,41 +110,40 @@ def _compute_phase_durations(self) -> dict | None: def _compute_success_windows(self, window_size: float = 60.0) -> list | None: if not self._connection_log: return None + events = sorted(self._connection_log, key=lambda entry: entry[0]) + start_time = self._scan_start if self._scan_start is not None else events[0][0] + buckets = {} + for timestamp, outcome in events: + bucket_index = int((timestamp - start_time) // window_size) + counts = buckets.setdefault(bucket_index, { + "attempts": 0, + "connected": 0, + "responsive": 0, + }) + counts["attempts"] += 1 + counts["connected"] += int(outcome == "connected") + counts["responsive"] += int(outcome in RESPONSIVE_CONNECTION_OUTCOMES) + windows = [] - start_time = self._connection_log[0][0] - end_time = self._connection_log[-1][0] - t = start_time - while t < end_time: - w_end = t + window_size - entries = [(ts, ok) for ts, ok in self._connection_log if t <= ts < w_end] - if entries: - rate = sum(1 for _, ok in entries if ok) / len(entries) - windows.append({ - "window_start": round(t - start_time, 1), - "window_end": round(w_end - start_time, 1), - "success_rate": round(rate, 3), - }) - t = w_end - return windows if windows else None + for bucket_index, counts in sorted(buckets.items()): + attempts = counts["attempts"] + window_start = bucket_index * window_size + windows.append({ + "window_start": round(window_start, 1), + "window_end": round(window_start + window_size, 1), + # Retain the historical connected-only meaning for compatibility. + "success_rate": round(counts["connected"] / attempts, 3), + "attempts": attempts, + "responsive_count": counts["responsive"], + "response_rate": round(counts["responsive"] / attempts, 3), + }) + return windows def _detect_rate_limiting(self) -> bool: - windows = self._compute_success_windows() - if not windows or len(windows) < 3: - return False - # Detect: last 2 windows have significantly lower success rate than first 2 - first = sum(w["success_rate"] for w in windows[:2]) / 2 - last = sum(w["success_rate"] for w in windows[-2:]) / 2 - return first > 0.5 and last < first * 0.7 + return detect_connection_signals(self._compute_success_windows())["rate_limiting_detected"] def _detect_blocking(self) -> bool: - windows = self._compute_success_windows() - if not windows or len(windows) < 2: - return False - # Detect: any window with 0% success rate after a window with >50% success - for i in range(1, len(windows)): - if windows[i - 1]["success_rate"] > 0.5 and windows[i]["success_rate"] == 0: - return True - return False + return detect_connection_signals(self._compute_success_windows())["blocking_detected"] def _compute_port_distribution(self) -> dict | None: if not self._open_ports: @@ -178,6 +178,8 @@ def build(self) -> ScanMetrics: probes_failed = sum(1 for v in self._probe_results.values() if v == "failed" or v.startswith("failed:")) banner_total = self._banner_confirmed + self._banner_guessed + connection_windows = self._compute_success_windows() + connection_signals = detect_connection_signals(connection_windows) return ScanMetrics( phase_durations=self._compute_phase_durations(), total_duration=round(time.time() - self._scan_start, 2) if self._scan_start else 0, @@ -185,9 +187,9 @@ def build(self) -> ScanMetrics: connection_outcomes=outcomes if total_connections > 0 else None, response_times=self._compute_stats(self._response_times), slow_ports=None, - success_rate_over_time=self._compute_success_windows(), - rate_limiting_detected=self._detect_rate_limiting(), - blocking_detected=self._detect_blocking(), + success_rate_over_time=connection_windows, + rate_limiting_detected=connection_signals["rate_limiting_detected"], + blocking_detected=connection_signals["blocking_detected"], coverage=self._compute_coverage(), probes_attempted=probes_attempted, probes_completed=probes_completed, diff --git a/extensions/business/cybersec/red_mesh/worker/pentest_worker.py b/extensions/business/cybersec/red_mesh/worker/pentest_worker.py index 1e9ca617f..c074ba04e 100644 --- a/extensions/business/cybersec/red_mesh/worker/pentest_worker.py +++ b/extensions/business/cybersec/red_mesh/worker/pentest_worker.py @@ -10,6 +10,7 @@ from .base import BaseLocalWorker from .service import _ServiceInfoMixin from .correlation import _CorrelationMixin +from .response_fingerprint import _ResponseFingerprintMixin from ..constants import ( PROBE_PROTOCOL_MAP, WEB_PROTOCOLS, WELL_KNOWN_PORTS as _WELL_KNOWN_PORTS, @@ -17,6 +18,8 @@ FINGERPRINT_NUDGE_TIMEOUT, SCAN_PORT_TIMEOUT, COMMON_PORTS, ALL_PORTS, NETWORK_FEATURE_METHODS, NETWORK_FEATURE_REGISTRY, + TIMEOUT_PROFILE_STANDARD, normalize_timeout_profile, + resolve_target_response_timeout, ) from .web import _WebTestsMixin from ..cve_db import reset_dynamic_reference_cache, set_dynamic_reference_cache @@ -29,6 +32,7 @@ class PentestLocalWorker( _ServiceInfoMixin, _WebTestsMixin, _CorrelationMixin, + _ResponseFingerprintMixin, BaseLocalWorker, ): FEATURE_CATEGORIES = ("service", "web", "correlation") @@ -39,6 +43,12 @@ class PentestLocalWorker( } PHASE_EXECUTION_PLAN = ( {"phase": "port_scan", "runner": "_scan_ports_step"}, + # Comparison-only, and node-wide rather than per-thread: it probes the + # mirrored comparison tier, which only one local worker per node holds. + # Deliberately carries no completion marker — it is evidence capture, not + # a feature test, and the progress denominator counts feature markers. + {"phase": "response_fingerprint", "runner": "_capture_response_fingerprint", + "skip_without_comparison_tier": True}, {"phase": "fingerprint", "runner": "_active_fingerprint_ports", "completion_marker": "fingerprint_completed"}, {"phase": "service_probes", "runner": "_gather_service_info", "completion_marker": "service_info_completed"}, {"phase": "web_tests", "runner": "_run_web_tests", "completion_marker": "web_tests_completed", "skip_on_ics": True}, @@ -81,6 +91,8 @@ def __init__( ics_safe_mode: bool = True, scanner_identity: str = "probe.redmesh.local", scanner_user_agent: str = "", + timeout_profile: str = TIMEOUT_PROFILE_STANDARD, + comparison_ports=None, dynamic_reference_cache=None, ): """ @@ -112,6 +124,10 @@ def __init__( Maximum random delay (seconds) between operations (Dune sand walking). ics_safe_mode : bool, optional Halt probing when ICS/SCADA indicators are detected. + comparison_ports : list[int], optional + Mirrored comparison-tier ports for geographic response comparison. + Supplied to exactly one worker per node; empty elsewhere, which skips + the response-fingerprint phase. scanner_identity : str, optional EHLO domain for SMTP probes. scanner_user_agent : str, optional @@ -149,6 +165,11 @@ def __init__( self._ics_detected = False self.scanner_identity = scanner_identity self.scanner_user_agent = scanner_user_agent + self.timeout_profile = normalize_timeout_profile(timeout_profile) + # Mirrored comparison tier. Only the node's designated worker receives it; + # for every other worker this stays empty and the response-fingerprint + # phase is skipped, so the tier is probed exactly once per vantage. + self.comparison_ports = [int(p) for p in (comparison_ports or [])] self.dynamic_reference_cache = dynamic_reference_cache or getattr(owner, "dynamic_reference_cache", None) self.P(f"Initializing pentest worker {self.local_worker_id} for target {self.target}...") @@ -358,6 +379,13 @@ def get_status(self, for_aggregations=False): dct_status["port_protocols"] = self.state.get("port_protocols", {}) dct_status["port_banners"] = self.state.get("port_banners", {}) + # Emitted only by the worker that holds the comparison tier. Absent + # elsewhere so the first-wins merge in report aggregation cannot overwrite + # real evidence with an empty placeholder from a sibling worker. + response_evidence = self.state.get("response_evidence") + if response_evidence: + dct_status["response_evidence"] = response_evidence + dct_status["scan_metadata"] = self.state.get("scan_metadata", {}) dct_status["correlation_findings"] = self.state.get("correlation_findings", []) dct_status["reference_data"] = self.state.get("reference_data", {}) @@ -370,6 +398,10 @@ def get_status(self, for_aggregations=False): # start(), stop(), _check_stopped(), P() are ALL inherited from # BaseLocalWorker. Not redefined here. + def _target_timeout(self, standard_timeout): + """Maximum wait for an ordinary response from the network target.""" + return resolve_target_response_timeout(self.timeout_profile, standard_timeout) + def _interruptible_sleep(self): """ Sleep for a random interval (Dune sand walking). @@ -394,6 +426,8 @@ def _execute_phase(self, phase_config): return if phase_config.get("skip_on_ics") and self._ics_detected: return + if phase_config.get("skip_without_comparison_tier") and not self.comparison_ports: + return phase_name = phase_config["phase"] runner = getattr(self, phase_config["runner"]) @@ -512,7 +546,7 @@ def _scan_ports_step(self, batch_size=None, batch_nr=1): if self.stop_event.is_set(): break sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(SCAN_PORT_TIMEOUT) + sock.settimeout(self._target_timeout(SCAN_PORT_TIMEOUT)) t0 = time.time() try: result = sock.connect_ex((target, port)) @@ -526,7 +560,7 @@ def _scan_ports_step(self, batch_size=None, batch_nr=1): protocol = None banner_text = "" try: - sock.settimeout(FINGERPRINT_TIMEOUT) + sock.settimeout(self._target_timeout(FINGERPRINT_TIMEOUT)) raw = sock.recv(FINGERPRINT_MAX_BANNER) except (socket.timeout, OSError): raw = b"" @@ -598,7 +632,7 @@ def _scan_ports_step(self, batch_size=None, batch_nr=1): elif result == errno.ECONNRESET: self.metrics.record_connection("reset", conn_time) else: - self.metrics.record_connection("refused", conn_time) + self.metrics.record_connection("error", conn_time) except Exception as e: self.metrics.record_connection("error", time.time() - t0) self.P(f"Exception scanning port {port} on {target}: {e}") @@ -676,7 +710,7 @@ def _active_fingerprint_ports(self): nudge_sock = None try: nudge_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - nudge_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + nudge_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) nudge_sock.connect((target, port)) nudge_sock.sendall(b"\r\n") try: @@ -725,7 +759,7 @@ def _active_fingerprint_ports(self): http_sock = None try: http_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - http_sock.settimeout(FINGERPRINT_HTTP_TIMEOUT) + http_sock.settimeout(self._target_timeout(FINGERPRINT_HTTP_TIMEOUT)) http_sock.connect((target, port)) http_sock.sendall(f"HEAD / HTTP/1.0\r\nHost: {target}\r\n\r\n".encode()) try: @@ -757,7 +791,7 @@ def _active_fingerprint_ports(self): mb_sock = None try: mb_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - mb_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + mb_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) mb_sock.connect((target, port)) mb_sock.sendall(b'\x00\x01\x00\x00\x00\x05\x01\x2b\x0e\x01\x00') try: @@ -797,7 +831,7 @@ def _active_fingerprint_ports(self): dns_sock = None try: dns_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - dns_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + dns_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) dns_sock.connect((target, port)) dns_query = ( b'\x12\x34' @@ -831,7 +865,7 @@ def _active_fingerprint_ports(self): r_sock = None try: r_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - r_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + r_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) r_sock.connect((target, port)) r_sock.sendall(b"PING\r\n") try: @@ -853,7 +887,7 @@ def _active_fingerprint_ports(self): pg_sock = None try: pg_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - pg_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + pg_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) pg_sock.connect((target, port)) pg_sock.sendall(b'\x00\x00\x00\x08\x04\xd2\x16\x2f') try: @@ -877,7 +911,7 @@ def _active_fingerprint_ports(self): mg_sock = None try: mg_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - mg_sock.settimeout(FINGERPRINT_NUDGE_TIMEOUT) + mg_sock.settimeout(self._target_timeout(FINGERPRINT_NUDGE_TIMEOUT)) mg_sock.connect((target, port)) _mg_field = b'\x10isMaster\x00' + struct.pack('(.*?)", re.IGNORECASE | re.DOTALL) + +# Order matters: +# - bearer precedes the generic key/value rule, which would otherwise match +# "Authorization: Bearer" and consume the scheme as the value, leaving the +# token itself in the clear; +# - email precedes the hex/base64 rules, which would otherwise swallow a long +# local part; +# - hex precedes base64, because hex is a strict subset of that alphabet. +_REDACTIONS = ( + (re.compile(r"(?i)\bbearer\s+[A-Za-z0-9._~+/\-]+=*"), "bearer [REDACTED]"), + ( + re.compile( + r"(?i)\b(authorization|api[-_]?key|apikey|access[-_]?token|token|" + r"session[-_]?id|sessionid|session|secret|password|passwd|pwd)\b" + # An optional closing quote covers JSON keys such as {"api_key": "..."}. + r"[\"']?(\s*[:=]\s*)[\"']?[^\s\"',;&<>]+" + ), + r"\1\2[REDACTED]", + ), + (re.compile(r"[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}"), "[REDACTED_EMAIL]"), + (re.compile(r"\b[A-Fa-f0-9]{32,}\b"), "[REDACTED_HEX]"), + (re.compile(r"\b[A-Za-z0-9+/]{32,}={0,2}"), "[REDACTED_B64]"), +) + + +def sanitize_excerpt(text): + """ + Scrub and bound a response body excerpt. + + Redaction runs before truncation so a credential straddling the byte cap + cannot survive as a leaked prefix. Truncation is byte-based but respects + UTF-8 character boundaries. + + Parameters + ---------- + text : str + Raw response body text. + + Returns + ------- + str or None + Scrubbed excerpt of at most ``EXCERPT_MAX_BYTES`` bytes, or None when + the input is empty. + """ + if not text: + return None + scrubbed = text[:EXCERPT_SCAN_CHARS] + for pattern, replacement in _REDACTIONS: + scrubbed = pattern.sub(replacement, scrubbed) + encoded = scrubbed.encode("utf-8")[:EXCERPT_MAX_BYTES] + # errors="ignore" drops a partial multi-byte character left by the slice. + excerpt = encoded.decode("utf-8", errors="ignore") + return excerpt or None + + +def excerpt_allowed(content_type): + """True when the content type is textual enough to excerpt safely.""" + if not content_type: + return False + base = content_type.split(";")[0].strip().lower() + return base in EXCERPT_CONTENT_TYPES + + +def normalize_content_type(content_type): + """Reduce a Content-Type header to its bare media type.""" + if not content_type: + return None + return content_type.split(";")[0].strip().lower() or None + + +def resolve_host(host): + """ + Resolve a target host to its sorted, deduplicated A/AAAA set. + + Each vantage resolves independently — this is what makes geo-DNS + divergence observable at all. + + Returns + ------- + tuple[list[str], str or None] + Resolved addresses and an error string when resolution failed. A + literal IP target resolves to an empty list with no error, because + there is no DNS answer to compare. + """ + if not host: + return [], "empty host" + try: + ipaddress.ip_address(host) + return [], None + except ValueError: + pass + try: + infos = socket.getaddrinfo(host, None) + except Exception as exc: + return [], str(exc) + addresses = {info[4][0] for info in infos if info[4]} + return sorted(addresses), None + + +def certificate_identity(cert_der): + """ + Extract comparable identity fields from a DER-encoded certificate. + + Returns None when the certificate is absent or unparseable — an + unparseable certificate must not masquerade as a matching one. + """ + if not cert_der: + return None + identity = {"cert_sha256": hashlib.sha256(cert_der).hexdigest()} + try: + from cryptography import x509 + from cryptography.x509.oid import NameOID + + cert = x509.load_der_x509_certificate(cert_der) + + def _common_name(name): + attributes = name.get_attributes_for_oid(NameOID.COMMON_NAME) + return attributes[0].value if attributes else None + + identity["subject_cn"] = _common_name(cert.subject) + identity["issuer_cn"] = _common_name(cert.issuer) + identity["not_before"] = cert.not_valid_before_utc.isoformat() + identity["not_after"] = cert.not_valid_after_utc.isoformat() + except Exception: + # The fingerprint alone still distinguishes one certificate from + # another, which is all clustering needs. + pass + return identity + + +class _ResponseFingerprintMixin: + """Captures per-vantage target response evidence for comparison jobs.""" + + def _capture_response_fingerprint(self): + """ + Probe the comparison port tier and record what the target answered. + + No-ops when this worker holds no comparison tier, which is how + non-comparison jobs and the non-designated local workers skip the phase. + """ + ports = sorted({int(port) for port in (self.comparison_ports or [])}) + if not ports: + return + + resolved_ips, resolver_error = resolve_host(self.target) + evidence = { + "target_host": self.target, + "resolved_ips": resolved_ips, + "ports": {}, + } + if resolver_error: + evidence["resolver_error"] = resolver_error + + for port in ports: + if self._check_stopped(): + break + evidence["ports"][str(port)] = self._fingerprint_port(port) + + self.state["response_evidence"] = evidence + + def _fingerprint_port(self, port): + """Capture reachability, TLS identity, and HTTP response for one port.""" + entry = {"reachable": False, "tls": None, "http": None, "excerpt": None} + + try: + with socket.create_connection((self.target, port), timeout=self._target_timeout(3)): + entry["reachable"] = True + except Exception: + # An unreachable port is itself a comparable result: a target that + # refuses one vantage and serves another is the divergence we are + # looking for, so this is recorded rather than treated as an error. + return entry + + _proto, _cipher, cert_der = self._tls_unverified_connect(self.target, port) + if cert_der: + identity = certificate_identity(cert_der) + if identity: + _dns_names, _ips = self._tls_parse_san_from_der(cert_der) + identity["san_count"] = len(_dns_names) + len(_ips) + entry["tls"] = identity + + scheme = "https" if entry["tls"] else "http" + http, excerpt = self._fingerprint_http(scheme, port) + entry["http"] = http + entry["excerpt"] = excerpt + return entry + + def _fingerprint_http(self, scheme, port): + """Issue one GET and reduce the response to comparable attributes.""" + url = f"{scheme}://{self.target}:{port}/" + try: + user_agent = getattr(self, "scanner_user_agent", "") + headers = {"User-Agent": user_agent} if user_agent else {} + resp = requests.get( + url, + timeout=self._target_timeout(5), + verify=False, + allow_redirects=True, + headers=headers, + ) + except Exception as exc: + self.P(f"Response fingerprint GET failed on {url}: {exc}", color='y') + return None, None + + content_type = normalize_content_type(resp.headers.get("Content-Type")) + title_match = _TITLE_RE.search(resp.text[:5000]) + http = { + "status": resp.status_code, + "final_url": resp.url, + "redirect_count": len(resp.history), + "title": title_match.group(1).strip()[:TITLE_MAX_CHARS] if title_match else None, + "content_type": content_type, + "body_length": len(resp.content), + "body_sha256": hashlib.sha256(resp.content).hexdigest(), + "headers": { + name: resp.headers.get(name) + for name in CAPTURED_HEADERS + if resp.headers.get(name) + }, + } + + excerpt = sanitize_excerpt(resp.text) if excerpt_allowed(content_type) else None + return http, excerpt diff --git a/extensions/business/cybersec/red_mesh/worker/service/_base.py b/extensions/business/cybersec/red_mesh/worker/service/_base.py index 383f55b2d..927171d03 100644 --- a/extensions/business/cybersec/red_mesh/worker/service/_base.py +++ b/extensions/business/cybersec/red_mesh/worker/service/_base.py @@ -1,5 +1,6 @@ from ...findings import Finding, Severity, probe_result, probe_error from ...cve_db import check_cves +from ...constants import resolve_target_response_timeout class _ServiceProbeBase: @@ -11,6 +12,12 @@ class _ServiceProbeBase: level imports. """ + def _target_timeout(self, standard_timeout): + """Resolve target waits, defaulting lightweight probe fixtures to Standard.""" + return resolve_target_response_timeout( + getattr(self, "timeout_profile", None), standard_timeout, + ) + def _emit_metadata(self, category, key_or_item, value=None): """Safely append to scan_metadata sub-dicts without crashing if state is uninitialized.""" meta = self.state.get("scan_metadata") diff --git a/extensions/business/cybersec/red_mesh/worker/service/common.py b/extensions/business/cybersec/red_mesh/worker/service/common.py index 0394c7a02..d4a7ccaeb 100644 --- a/extensions/business/cybersec/red_mesh/worker/service/common.py +++ b/extensions/business/cybersec/red_mesh/worker/service/common.py @@ -103,7 +103,7 @@ def _service_info_http(self, target, port): # default port: 80 self.P(f"Fetching {url} for banner...") ua = getattr(self, 'scanner_user_agent', '') headers = {'User-Agent': ua} if ua else {} - resp = requests.get(url, timeout=5, verify=False, allow_redirects=True, headers=headers) + resp = requests.get(url, timeout=self._target_timeout(5), verify=False, allow_redirects=True, headers=headers) result["banner"] = f"HTTP {resp.status_code} {resp.reason}" result["server"] = resp.headers.get("Server") @@ -158,7 +158,7 @@ def _service_info_http(self, target, port): # default port: 80 # (some servers like nginx drop requests with unrecognized Host values). try: _s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - _s.settimeout(3) + _s.settimeout(self._target_timeout(3)) _s.connect((target, port)) # Use HTTP/1.0 without Host — matches nmap's GetRequest probe _s.send(b"GET / HTTP/1.0\r\n\r\n") @@ -234,7 +234,7 @@ def _service_info_http(self, target, port): # default port: 80 dangerous = [] for method in ("TRACE", "PUT", "DELETE"): try: - r = requests.request(method, url, timeout=3, verify=False) + r = requests.request(method, url, timeout=self._target_timeout(3), verify=False) if r.status_code < 400: dangerous.append(method) except Exception: @@ -310,7 +310,7 @@ def _service_info_http_alt(self, target, port): # default port: 8080 raw = {"banner": None, "server": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(2) + sock.settimeout(self._target_timeout(2)) sock.connect((target, port)) ua = getattr(self, 'scanner_user_agent', '') ua_header = f"\r\nUser-Agent: {ua}" if ua else "" @@ -370,7 +370,7 @@ def _service_info_https(self, target, port): # default port: 443 self.P(f"Fetching {url} for banner...") ua = getattr(self, 'scanner_user_agent', '') headers = {'User-Agent': ua} if ua else {} - resp = requests.get(url, timeout=3, verify=False, headers=headers) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False, headers=headers) raw["banner"] = f"HTTPS {resp.status_code} {resp.reason}" raw["server"] = resp.headers.get("Server") if raw["server"]: @@ -445,7 +445,7 @@ def _service_info_http_basic_auth(self, target, port): realm = None for path in ("/", "/admin", "/manager"): try: - resp = requests.get(base_url + path, timeout=3, verify=False) + resp = requests.get(base_url + path, timeout=self._target_timeout(3), verify=False) if resp.status_code == 401: www_auth = resp.headers.get("WWW-Authenticate", "") if "Basic" in www_auth: @@ -466,7 +466,7 @@ def _service_info_http_basic_auth(self, target, port): consecutive_401 = 0 for username, password in self._HTTP_BASIC_CREDS: try: - resp = requests.get(auth_url, timeout=3, verify=False, auth=(username, password)) + resp = requests.get(auth_url, timeout=self._target_timeout(3), verify=False, auth=(username, password)) raw["tested"] += 1 if resp.status_code == 429: @@ -558,8 +558,8 @@ def _service_info_ftp(self, target, port): # default port: 21 def _ftp_connect(user=None, passwd=None): """Open a fresh FTP connection and optionally login.""" - ftp = ftplib.FTP(timeout=5) - ftp.connect(target, port, timeout=5) + ftp = ftplib.FTP(timeout=self._target_timeout(5)) + ftp.connect(target, port, timeout=self._target_timeout(5)) if user is not None: ftp.login(user, passwd or "") return ftp @@ -837,7 +837,7 @@ def _service_info_ssh(self, target, port): # default port: 22 # --- 1. Banner grab (raw socket) --- try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) banner = sock.recv(1024).decode("utf-8", errors="ignore").strip() sock.close() @@ -886,7 +886,7 @@ def _service_info_ssh(self, target, port): # default port: 22 client.connect( target, port=port, username=username, password=password, - timeout=3, auth_timeout=3, + timeout=self._target_timeout(3), auth_timeout=self._target_timeout(3), look_for_keys=False, allow_agent=False, ) accepted_creds.append(f"{username}:{password}") @@ -905,7 +905,7 @@ def _service_info_ssh(self, target, port): # default port: 22 client.connect( target, port=port, username=random_user, password=random_pass, - timeout=3, auth_timeout=3, + timeout=self._target_timeout(3), auth_timeout=self._target_timeout(3), look_for_keys=False, allow_agent=False, ) findings.append(Finding( @@ -1110,7 +1110,7 @@ def _ssh_check_libssh_bypass(self, target, port): msg.add_byte(b'\x34') transport._send_message(msg) try: - chan = transport.open_session(timeout=3) + chan = transport.open_session(timeout=self._target_timeout(3)) if chan is not None: chan.close() transport.close() @@ -1181,7 +1181,7 @@ def _service_info_smtp(self, target, port): # default port: 25 # --- 1. Connect and grab banner --- try: - smtp = smtplib.SMTP(timeout=5) + smtp = smtplib.SMTP(timeout=self._target_timeout(5)) code, msg = smtp.connect(target, port) result["banner"] = f"{code} {msg.decode(errors='replace')}" except Exception as e: @@ -1355,7 +1355,7 @@ def _service_info_smtp(self, target, port): # default port: 25 except Exception: pass try: - smtp = smtplib.SMTP(target, port, timeout=5) + smtp = smtplib.SMTP(target, port, timeout=self._target_timeout(5)) smtp.ehlo(identity) except Exception: smtp = None @@ -1467,7 +1467,7 @@ def _service_info_telnet(self, target, port): # default port: 23 # --- 1. Banner grab + IAC negotiation parsing --- try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(5) + sock.settimeout(self._target_timeout(5)) sock.connect((target, port)) raw = sock.recv(2048) sock.close() @@ -1509,7 +1509,7 @@ def _try_telnet_login(user, passwd): """Attempt Telnet login, return (success, uid_line, uname_line).""" try: s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - s.settimeout(5) + s.settimeout(self._target_timeout(5)) s.connect((target, port)) # Read until login prompt @@ -1707,7 +1707,7 @@ def _service_info_rsync(self, target, port): # default port: 873 # --- 1. Connect and receive banner --- try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) banner = sock.recv(256).decode("utf-8", errors="ignore").strip() except Exception as e: @@ -1794,7 +1794,7 @@ def _service_info_rsync(self, target, port): # default port: 873 for mod in raw["modules"]: try: sock2 = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock2.settimeout(3) + sock2.settimeout(self._target_timeout(3)) sock2.connect((target, port)) sock2.recv(256) # banner sock2.sendall(f"@RSYNCD: {proto_version}\n".encode()) diff --git a/extensions/business/cybersec/red_mesh/worker/service/database.py b/extensions/business/cybersec/red_mesh/worker/service/database.py index f924836c7..e56e16577 100644 --- a/extensions/business/cybersec/red_mesh/worker/service/database.py +++ b/extensions/business/cybersec/red_mesh/worker/service/database.py @@ -44,7 +44,7 @@ def _service_info_mysql(self, target, port): # default port: 3306 raw = {"version": None, "auth_plugin": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) data = sock.recv(256) sock.close() @@ -174,7 +174,7 @@ def _service_info_mysql_creds(self, target, port): # default port: 3306 for username, password in creds: try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) data = sock.recv(256) @@ -289,7 +289,7 @@ def _mysql_test_cve_2012_2122(self, target, port): # First, connect to get version try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) data = sock.recv(256) sock.close() @@ -322,7 +322,7 @@ def _mysql_test_cve_2012_2122(self, target, port): try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(5) + sock.settimeout(self._target_timeout(5)) sock.connect((target, port)) for _ in range(attempts): @@ -382,7 +382,7 @@ def _mysql_test_cve_2012_2122(self, target, port): if resp and len(resp) >= 5 and resp[4] == 0xFF: sock.close() sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) sock.close() @@ -451,7 +451,7 @@ def _redis_connect(self, target, port): """Open a TCP socket to Redis.""" try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) return sock except Exception as e: @@ -666,7 +666,7 @@ def _service_info_mssql(self, target, port): # default port: 1433 raw = {"banner": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) prelogin = bytes.fromhex( "1201001600000000000000000000000000000000000000000000000000000000" @@ -719,7 +719,7 @@ def _service_info_postgresql(self, target, port): # default port: 5432 raw = {"auth_type": None, "version": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) payload = b'user\x00postgres\x00database\x00postgres\x00\x00' startup = struct.pack('!I', len(payload) + 8) + struct.pack('!I', 196608) + payload @@ -890,7 +890,7 @@ def _service_info_postgresql_creds(self, target, port): # default port: 5432 for username, password in creds: try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) payload = f'user\x00{username}\x00database\x00postgres\x00\x00'.encode() startup = struct.pack('!I', len(payload) + 8) + struct.pack('!I', 196608) + payload @@ -1041,7 +1041,7 @@ def _service_info_memcached(self, target, port): # default port: 11211 raw = {"banner": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(2) + sock.settimeout(self._target_timeout(2)) sock.connect((target, port)) # Extract version @@ -1153,11 +1153,10 @@ def _service_info_mongodb(self, target, port): # default port: 27017 return probe_error(target, port, "MongoDB", e) return probe_result(raw_data=raw, findings=findings) - @staticmethod - def _mongodb_query(target, port, command_name): + def _mongodb_query(self, target, port, command_name): """Send a MongoDB OP_QUERY command and return the raw response bytes.""" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) # Build BSON: {: 1} field = b'\x10' + command_name + b'\x00' + struct.pack('HHHHHH', tid, 0x0100, 1, 0, 0, 0) qname = b'\x07version\x04bind\x00' @@ -604,7 +604,7 @@ def _dns_discover_zones(self, target, port): for domain in list(candidates): try: sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - sock.settimeout(2) + sock.settimeout(self._target_timeout(2)) tid = random.randint(0, 0xffff) header = struct.pack('>HHHHHH', tid, 0x0100, 1, 0, 0, 0) qname = b"" @@ -648,7 +648,7 @@ def _dns_test_axfr(self, target, port): for domain in test_domains[:4]: # Test at most 4 domains try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) # Build AXFR query @@ -711,7 +711,7 @@ def _dns_test_open_resolver(self, target, port): """ try: sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - sock.settimeout(2) + sock.settimeout(self._target_timeout(2)) tid = random.randint(0, 0xffff) # Standard recursive query for example.com A record header = struct.pack('>HHHHHH', tid, 0x0100, 1, 0, 0, 0) # RD=1 @@ -808,7 +808,7 @@ def _service_info_smb(self, target, port): # default port: 445 try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(4) + sock.settimeout(self._target_timeout(4)) sock.connect((target, port)) sock.sendall(netbios_header + smb_payload) @@ -1024,7 +1024,7 @@ def _smb_enum_shares(self, target, port): sock = None try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(4) + sock.settimeout(self._target_timeout(4)) sock.connect((target, port)) def _send_smb(payload): @@ -1479,7 +1479,7 @@ def _smb_try_null_session(self, target, port): """ try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) # --- Negotiate --- @@ -1665,7 +1665,7 @@ def _udp_nbns_probe(udp_port): sock = None try: sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.sendto(nbns_query, (target, udp_port)) data, _ = sock.recvfrom(1024) return _parse_nbns_response(data) @@ -1750,7 +1750,7 @@ def _add_nbns_findings(names, probe_label): wrepl_packet = wrepl_header + wrepl_body sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) sock.sendall(wrepl_packet) @@ -1887,7 +1887,7 @@ def _service_info_modbus(self, target, port): # default port: 502 raw = {"banner": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) request = b'\x00\x01\x00\x00\x00\x06\x01\x2b\x0e\x01\x00' sock.sendall(request) @@ -1963,7 +1963,7 @@ def _es_check_root(self, base_url, raw): """GET / — extract version, cluster name.""" findings = [] try: - resp = requests.get(base_url, timeout=3) + resp = requests.get(base_url, timeout=self._target_timeout(3)) if resp.ok: try: data = resp.json() @@ -2001,7 +2001,7 @@ def _es_check_indices(self, base_url, raw): """GET /_cat/indices — list accessible indices.""" findings = [] try: - resp = requests.get(f"{base_url}/_cat/indices?v", timeout=3) + resp = requests.get(f"{base_url}/_cat/indices?v", timeout=self._target_timeout(3)) if resp.ok and resp.text.strip(): lines = resp.text.strip().split("\n") index_count = max(0, len(lines) - 1) # subtract header @@ -2025,7 +2025,7 @@ def _es_check_nodes(self, base_url, raw): """GET /_nodes — extract transport/publish addresses, classify IPs, check JVM.""" findings = [] try: - resp = requests.get(f"{base_url}/_nodes", timeout=3) + resp = requests.get(f"{base_url}/_nodes", timeout=self._target_timeout(3)) if resp.ok: data = resp.json() nodes = data.get("nodes", {}) diff --git a/extensions/business/cybersec/red_mesh/worker/service/tls.py b/extensions/business/cybersec/red_mesh/worker/service/tls.py index 02c85d123..c88ef56d2 100644 --- a/extensions/business/cybersec/red_mesh/worker/service/tls.py +++ b/extensions/business/cybersec/red_mesh/worker/service/tls.py @@ -111,7 +111,7 @@ def _tls_unverified_connect(self, target, port): ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) ctx.check_hostname = False ctx.verify_mode = ssl.CERT_NONE - with socket.create_connection((target, port), timeout=3) as sock: + with socket.create_connection((target, port), timeout=self._target_timeout(3)) as sock: with ctx.wrap_socket(sock, server_hostname=target) as ssock: proto = ssock.version() cipher_info = ssock.cipher() @@ -157,7 +157,7 @@ def _tls_check_certificate(self, target, port, raw): findings = [] try: ctx = ssl.create_default_context() - with socket.create_connection((target, port), timeout=3) as sock: + with socket.create_connection((target, port), timeout=self._target_timeout(3)) as sock: with ctx.wrap_socket(sock, server_hostname=target) as ssock: cert = ssock.getpeercert() subj = dict(x[0] for x in cert.get("subject", ())) @@ -353,7 +353,7 @@ def _tls_check_heartbleed(self, target, port): ctx.minimum_version = ssl.TLSVersion.MINIMUM_SUPPORTED raw_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - raw_sock.settimeout(3) + raw_sock.settimeout(self._target_timeout(3)) raw_sock.connect((target, port)) tls_sock = ctx.wrap_socket(raw_sock, server_hostname=target) @@ -393,7 +393,7 @@ def _tls_check_heartbleed(self, target, port): try: raw_after = tls_sock.unwrap() raw_after.sendall(tls_record) - raw_after.settimeout(3) + raw_after.settimeout(self._target_timeout(3)) response = raw_after.recv(65536) raw_after.close() except (ssl.SSLError, OSError): @@ -443,7 +443,7 @@ def _tls_heartbleed_raw(self, target, port, tls_ver_bytes): """ try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(5) + sock.settimeout(self._target_timeout(5)) sock.connect((target, port)) # Minimal TLS 1.0 ClientHello with heartbeat extension @@ -506,7 +506,7 @@ def _tls_heartbleed_raw(self, target, port, tls_ver_bytes): sock.sendall(hb_record) # Read response - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) try: response = sock.recv(65536) except (socket.timeout, OSError): @@ -548,7 +548,7 @@ def _tls_check_downgrade(self, target, port): ctx.maximum_version = ssl.TLSVersion.SSLv3 ctx.minimum_version = ssl.TLSVersion.SSLv3 sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) tls_sock = ctx.wrap_socket(sock, server_hostname=target) negotiated = tls_sock.version() @@ -576,7 +576,7 @@ def _tls_check_downgrade(self, target, port): ctx.maximum_version = ssl.TLSVersion.TLSv1 ctx.minimum_version = ssl.TLSVersion.TLSv1 sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(3) + sock.settimeout(self._target_timeout(3)) sock.connect((target, port)) tls_sock = ctx.wrap_socket(sock, server_hostname=target) negotiated = tls_sock.version() @@ -660,7 +660,7 @@ def _service_info_generic(self, target, port): raw = {"banner": None} try: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(2) + sock.settimeout(self._target_timeout(2)) sock.connect((target, port)) raw_bytes = sock.recv(512) sock.close() diff --git a/extensions/business/cybersec/red_mesh/worker/web/api_exposure.py b/extensions/business/cybersec/red_mesh/worker/web/api_exposure.py index f2b038ae0..477f5f924 100644 --- a/extensions/business/cybersec/red_mesh/worker/web/api_exposure.py +++ b/extensions/business/cybersec/red_mesh/worker/web/api_exposure.py @@ -45,7 +45,7 @@ def _web_test_graphql_introspection(self, target, port): graphql_url = base_url.rstrip("/") + "/graphql" try: payload = {"query": "{__schema{types{name}}}"} - resp = requests.post(graphql_url, json=payload, timeout=5, verify=False) + resp = requests.post(graphql_url, json=payload, timeout=self._target_timeout(5), verify=False) if resp.status_code == 200 and "__schema" in resp.text: findings_list.append(Finding( severity=Severity.MEDIUM, @@ -112,7 +112,7 @@ def _web_test_metadata_endpoints(self, target, port): try: for path, provider, extra_headers in metadata_paths: url = base_url.rstrip("/") + path - resp = requests.get(url, timeout=3, verify=False, headers=extra_headers) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False, headers=extra_headers) if resp.status_code == 200: findings_list.append(Finding( severity=Severity.CRITICAL, @@ -187,7 +187,7 @@ def _web_test_ssrf_basic(self, target, port): break try: url = f"{base_url.rstrip('/')}{path}?{param}={ssrf_payload}" - resp = requests.get(url, timeout=4, verify=False) + resp = requests.get(url, timeout=self._target_timeout(4), verify=False) body_lower = resp.text.lower() if resp.status_code == 200 and any(m in body_lower for m in self._SSRF_MARKERS): findings_list.append(Finding( @@ -251,7 +251,7 @@ def _web_test_api_auth_bypass(self, target, port): url = base_url.rstrip("/") + path resp = requests.get( url, - timeout=3, + timeout=self._target_timeout(3), verify=False, headers={"Authorization": "Bearer invalid-token"}, ) diff --git a/extensions/business/cybersec/red_mesh/worker/web/discovery.py b/extensions/business/cybersec/red_mesh/worker/web/discovery.py index a8333245f..d9d53a438 100644 --- a/extensions/business/cybersec/red_mesh/worker/web/discovery.py +++ b/extensions/business/cybersec/red_mesh/worker/web/discovery.py @@ -51,7 +51,7 @@ def _web_test_common(self, target, port): # --- Catch-all detection: 200-for-all --- try: canary_path = f"/{_uuid.uuid4().hex}" - canary_resp = requests.get(base_url + canary_path, timeout=2, verify=False) + canary_resp = requests.get(base_url + canary_path, timeout=self._target_timeout(2), verify=False) if canary_resp.status_code == 200: findings_list.append(Finding( severity=Severity.HIGH, @@ -100,7 +100,7 @@ def _web_test_common(self, target, port): try: for path, (severity, cwe, owasp, desc) in _PATH_META.items(): url = base_url + path - resp = requests.get(url, timeout=2, verify=False) + resp = requests.get(url, timeout=self._target_timeout(2), verify=False) if resp.status_code == 200: findings_list.append(Finding( severity=severity, @@ -161,7 +161,7 @@ def _web_test_homepage(self, target, port): } try: - resp_main = requests.get(base_url, timeout=3, verify=False) + resp_main = requests.get(base_url, timeout=self._target_timeout(3), verify=False) text = resp_main.text[:10000] for marker, (severity, title, owasp) in _MARKER_META.items(): if marker in text: @@ -215,7 +215,7 @@ def _web_test_tech_fingerprint(self, target, port): base_url = f"{scheme}://{target}" if port in (80, 443) else f"{scheme}://{target}:{port}" try: - resp = requests.get(base_url, timeout=4, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(4), verify=False) # Server header server = resp.headers.get("Server") @@ -359,7 +359,7 @@ def _web_test_vpn_endpoints(self, target, port): for entry in vpn_checks: try: url = base_url.rstrip("/") + entry["path"] - resp = requests.get(url, timeout=3, verify=False, allow_redirects=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False, allow_redirects=False) if entry["check"](resp): raw["vpn_endpoints"].append({"product": entry["product"], "path": entry["path"]}) findings_list.append(Finding( @@ -429,7 +429,7 @@ def _web_test_cms_fingerprint(self, target, port): # --- WordPress detection --- wp_version = None try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) if resp.ok: gen_match = _re.search( r']*name=["\']generator["\'][^>]*content=["\']WordPress\s+([0-9.]+)', @@ -444,7 +444,7 @@ def _web_test_cms_fingerprint(self, target, port): if not wp_version: try: - resp = requests.get(base_url + "/wp-login.php", timeout=3, verify=False, allow_redirects=False) + resp = requests.get(base_url + "/wp-login.php", timeout=self._target_timeout(3), verify=False, allow_redirects=False) if resp.status_code in (200, 302) and ('wp-login' in resp.text.lower() or 'wordpress' in resp.text.lower()): wp_version = "unknown" except Exception: @@ -458,7 +458,7 @@ def _web_test_cms_fingerprint(self, target, port): ("/readme.html", r'Version\s+([0-9.]+)'), ]: try: - resp = requests.get(base_url + _wp_path, timeout=3, verify=False) + resp = requests.get(base_url + _wp_path, timeout=self._target_timeout(3), verify=False) if resp.ok: _wp_m = _re.search(_wp_re, resp.text, _re.IGNORECASE) if _wp_m: @@ -484,7 +484,7 @@ def _web_test_cms_fingerprint(self, target, port): findings_list += self._wp_detect_plugins(base_url) for path, desc in self._WP_SENSITIVE_PATHS: try: - resp = requests.get(base_url + path, timeout=3, verify=False) + resp = requests.get(base_url + path, timeout=self._target_timeout(3), verify=False) if resp.status_code == 200: findings_list.append(Finding( severity=Severity.MEDIUM, @@ -503,7 +503,7 @@ def _web_test_cms_fingerprint(self, target, port): # --- Drupal detection --- drupal_version = None try: - resp = requests.get(base_url + "/core/CHANGELOG.txt", timeout=3, verify=False) + resp = requests.get(base_url + "/core/CHANGELOG.txt", timeout=self._target_timeout(3), verify=False) if resp.ok and "Drupal" in resp.text: ver_match = _re.search(r'Drupal\s+([0-9.]+)', resp.text) drupal_version = ver_match.group(1) if ver_match else "unknown" @@ -511,7 +511,7 @@ def _web_test_cms_fingerprint(self, target, port): pass if not drupal_version: try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) if resp.ok: gen_match = _re.search( r']*name=["\']generator["\'][^>]*content=["\']Drupal\s+([0-9.]+)', @@ -531,7 +531,7 @@ def _web_test_cms_fingerprint(self, target, port): if drupal_version and (drupal_version == "unknown" or _re.match(r'^\d+$', drupal_version)): for _dp_path, _dp_re in _DRUPAL_VERSION_SOURCES: try: - resp = requests.get(base_url + _dp_path, timeout=3, verify=False) + resp = requests.get(base_url + _dp_path, timeout=self._target_timeout(3), verify=False) if resp.ok: _dp_m = _re.search(_dp_re, resp.text) if _dp_m: @@ -559,11 +559,11 @@ def _web_test_cms_fingerprint(self, target, port): # --- Joomla detection --- joomla_version = None try: - resp = requests.get(base_url + "/administrator/", timeout=3, verify=False, allow_redirects=False) + resp = requests.get(base_url + "/administrator/", timeout=self._target_timeout(3), verify=False, allow_redirects=False) if resp.status_code in (200, 302) and 'joomla' in resp.text.lower(): joomla_version = "unknown" try: - resp2 = requests.get(base_url + "/language/en-GB/en-GB.xml", timeout=3, verify=False) + resp2 = requests.get(base_url + "/language/en-GB/en-GB.xml", timeout=self._target_timeout(3), verify=False) if resp2.ok: ver_match = _re.search(r'([0-9.]+)', resp2.text) if ver_match: @@ -591,7 +591,7 @@ def _web_test_cms_fingerprint(self, target, port): try: resp = requests.get( base_url + "/api/index.php/v1/config/application?public=true", - timeout=3, verify=False, + timeout=self._target_timeout(3), verify=False, ) if resp.ok and ("password" in resp.text.lower() or '"db"' in resp.text.lower() or '"dbtype"' in resp.text.lower()): findings_list.append(Finding( @@ -615,7 +615,7 @@ def _web_test_cms_fingerprint(self, target, port): # Check /_ignition/health-check try: - resp = requests.get(base_url + "/_ignition/health-check", timeout=3, verify=False) + resp = requests.get(base_url + "/_ignition/health-check", timeout=self._target_timeout(3), verify=False) if resp.ok and ("can_execute_commands" in resp.text or "ok" in resp.text.lower()): ignition_detected = True laravel_detected = True @@ -638,7 +638,7 @@ def _web_test_cms_fingerprint(self, target, port): try: resp = requests.get( base_url + "/nonexistent_" + _uuid.uuid4().hex[:8], - timeout=3, verify=False, + timeout=self._target_timeout(3), verify=False, ) body = resp.text[:10000].lower() if "laravel" in body or "illuminate" in body: @@ -652,7 +652,7 @@ def _web_test_cms_fingerprint(self, target, port): resp = requests.post( base_url + "/_ignition/execute-solution", json={"solution": "test", "parameters": {}}, - timeout=3, verify=False, + timeout=self._target_timeout(3), verify=False, ) if resp.status_code != 404: findings_list.append(Finding( @@ -724,7 +724,7 @@ def _wp_detect_plugins(self, base_url): for slug, name in self._WP_PLUGIN_CHECKS: try: url = f"{base_url}/wp-content/plugins/{slug}/readme.txt" - resp = requests.get(url, timeout=3, verify=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False) if resp.status_code != 200: continue ver_match = _re.search(r'Stable tag:\s*([0-9.]+)', resp.text, _re.IGNORECASE) @@ -799,7 +799,7 @@ def _web_test_verbose_errors(self, target, port): # --- 1. Trigger a 404 and inspect the error page --- try: canary = f"/nonexistent_{_uuid.uuid4().hex[:8]}" - resp = requests.get(base_url + canary, timeout=3, verify=False) + resp = requests.get(base_url + canary, timeout=self._target_timeout(3), verify=False) body = resp.text[:10000] for marker, framework in self._STACK_TRACE_MARKERS: @@ -835,7 +835,7 @@ def _web_test_verbose_errors(self, target, port): # --- 2. Debug mode detection on homepage --- try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) body = resp.text[:10000] for marker, framework in self._DEBUG_MODE_MARKERS: if marker in body: @@ -856,7 +856,7 @@ def _web_test_verbose_errors(self, target, port): # --- 3. Django __debug__/ endpoint --- try: - resp = requests.get(base_url + "/__debug__/", timeout=3, verify=False) + resp = requests.get(base_url + "/__debug__/", timeout=self._target_timeout(3), verify=False) if resp.status_code == 200 and "djdt" in resp.text.lower(): findings_list.append(Finding( severity=Severity.HIGH, @@ -915,7 +915,7 @@ def _web_test_java_servers(self, target, port): weblogic_version = None # Console login page try: - resp = requests.get(base_url + "/console/login/LoginForm.jsp", timeout=4, verify=False, allow_redirects=True) + resp = requests.get(base_url + "/console/login/LoginForm.jsp", timeout=self._target_timeout(4), verify=False, allow_redirects=True) if resp.ok and "WebLogic" in resp.text: raw["java_server"] = "WebLogic" ver_m = _re.search(r'(?:WebLogic Server|footerVersion)[^0-9]*(\d+\.\d+\.\d+\.\d+)', resp.text) @@ -928,7 +928,7 @@ def _web_test_java_servers(self, target, port): # T3/IIOP banner on root (some WebLogic instances) if not weblogic_version: try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) if resp.ok: # Check for WebLogic error page patterns if "WebLogic" in resp.text or "BEA-" in resp.text: @@ -958,7 +958,7 @@ def _web_test_java_servers(self, target, port): findings_list += check_cves("weblogic", weblogic_version) # Check for console exposure try: - resp = requests.get(base_url + "/console/", timeout=3, verify=False, allow_redirects=True) + resp = requests.get(base_url + "/console/", timeout=self._target_timeout(3), verify=False, allow_redirects=True) if resp.ok and ("login" in resp.text.lower() or "WebLogic" in resp.text): findings_list.append(Finding( severity=Severity.HIGH, @@ -975,7 +975,7 @@ def _web_test_java_servers(self, target, port): # CVE-2020-14882: Console authentication bypass via double-encoded path try: bypass_url = base_url + "/console/css/%252e%252e%252fconsole.portal" - resp = requests.get(bypass_url, timeout=4, verify=False, allow_redirects=False) + resp = requests.get(bypass_url, timeout=self._target_timeout(4), verify=False, allow_redirects=False) if resp.status_code == 200 and len(resp.text) > 500 and ( "portal" in resp.text.lower() or "console" in resp.text.lower()): findings_list.append(Finding( @@ -997,7 +997,7 @@ def _web_test_java_servers(self, target, port): # --- 2. Tomcat detection --- tomcat_version = None try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) if resp.ok: # Tomcat default page or error page tc_m = _re.search(r'Apache Tomcat[/\s]*(\d+\.\d+\.\d+)', resp.text) @@ -1015,7 +1015,7 @@ def _web_test_java_servers(self, target, port): # Try 404 page which often reveals Tomcat version if not tomcat_version: try: - resp = requests.get(base_url + "/nonexistent_" + _uuid.uuid4().hex[:6], timeout=3, verify=False) + resp = requests.get(base_url + "/nonexistent_" + _uuid.uuid4().hex[:6], timeout=self._target_timeout(3), verify=False) tc_m = _re.search(r'Apache Tomcat[/\s]*(\d+\.\d+\.\d+)', resp.text) if tc_m: tomcat_version = tc_m.group(1) @@ -1038,7 +1038,7 @@ def _web_test_java_servers(self, target, port): # Manager app exposure for mgr_path in ["/manager/html", "/manager/status"]: try: - resp = requests.get(base_url + mgr_path, timeout=3, verify=False) + resp = requests.get(base_url + mgr_path, timeout=self._target_timeout(3), verify=False) if resp.status_code in (200, 401, 403): findings_list.append(Finding( severity=Severity.HIGH if resp.status_code == 200 else Severity.MEDIUM, @@ -1058,7 +1058,7 @@ def _web_test_java_servers(self, target, port): # --- 3. JBoss / WildFly detection --- jboss_version = None try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) xpb = resp.headers.get("X-Powered-By", "") jb_m = _re.search(r'JBossAS[- ]*(\d+)', xpb) if jb_m: @@ -1105,7 +1105,7 @@ def _web_test_java_servers(self, target, port): findings_list += check_cves("jboss", jboss_version) # JMX console exposure try: - resp = requests.get(base_url + "/jmx-console/", timeout=3, verify=False) + resp = requests.get(base_url + "/jmx-console/", timeout=self._target_timeout(3), verify=False) if resp.status_code in (200, 401): findings_list.append(Finding( severity=Severity.HIGH if resp.status_code == 200 else Severity.MEDIUM, @@ -1124,7 +1124,7 @@ def _web_test_java_servers(self, target, port): # --- 4. Spring Framework detection --- spring_detected = False try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) body = resp.text[:10000] # Spring Whitelabel Error Page if "Whitelabel Error Page" in body or "Spring" in resp.headers.get("X-Application-Context", ""): @@ -1136,7 +1136,7 @@ def _web_test_java_servers(self, target, port): pass if not spring_detected: try: - resp = requests.get(base_url + "/nonexistent_" + _uuid.uuid4().hex[:6], timeout=3, verify=False) + resp = requests.get(base_url + "/nonexistent_" + _uuid.uuid4().hex[:6], timeout=self._target_timeout(3), verify=False) if "Whitelabel Error Page" in resp.text: spring_detected = True elif "org.springframework" in resp.text or "DispatcherServlet" in resp.text: @@ -1146,7 +1146,7 @@ def _web_test_java_servers(self, target, port): # Spring MVC: POST to root returns 405 with Spring-specific message if not spring_detected: try: - resp = requests.post(base_url, data="", timeout=3, verify=False) + resp = requests.post(base_url, data="", timeout=self._target_timeout(3), verify=False) if resp.status_code == 405: body = resp.text if "Request method" in body and "not supported" in body: @@ -1173,7 +1173,7 @@ def _web_test_java_servers(self, target, port): struts_evidence = "" # 5a. Check /struts/utils.js — present in all Struts2 apps using tag try: - resp = requests.get(base_url + "/struts/utils.js", timeout=3, verify=False) + resp = requests.get(base_url + "/struts/utils.js", timeout=self._target_timeout(3), verify=False) if resp.ok and len(resp.text) > 50: struts_detected = True struts_evidence = "/struts/utils.js present" @@ -1182,7 +1182,7 @@ def _web_test_java_servers(self, target, port): # 5b. Check homepage for .action/.do URLs or Struts indicators if not struts_detected: try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) body = resp.text[:10000] struts_indicators = [".action", ".do", "struts", "Struts Problem Report"] if any(ind in body for ind in struts_indicators): @@ -1219,7 +1219,7 @@ def _web_test_java_servers(self, target, port): # --- 6. Jetty detection (from Server header) --- try: - resp = requests.get(base_url, timeout=3, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) srv = resp.headers.get("Server", "") jetty_m = _re.search(r'[Jj]etty\(?(\d+\.\d+\.\d+)', srv) if jetty_m: @@ -1293,7 +1293,7 @@ def _web_test_js_library_versions(self, target, port): base_url = f"{scheme}://{target}" if port in (80, 443) else f"{scheme}://{target}:{port}" try: - resp = requests.get(base_url, timeout=4, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(4), verify=False) if resp.status_code != 200: return probe_result(findings=findings_list) html = resp.text diff --git a/extensions/business/cybersec/red_mesh/worker/web/hardening.py b/extensions/business/cybersec/red_mesh/worker/web/hardening.py index 4f91eba29..a16e165bb 100644 --- a/extensions/business/cybersec/red_mesh/worker/web/hardening.py +++ b/extensions/business/cybersec/red_mesh/worker/web/hardening.py @@ -45,7 +45,7 @@ def _web_test_flags(self, target, port): base_url = f"{scheme}://{target}:{port}" try: - resp_main = requests.get(base_url, timeout=3, verify=False) + resp_main = requests.get(base_url, timeout=self._target_timeout(3), verify=False) # Check cookies for Secure/HttpOnly flags cookies_hdr = resp_main.headers.get("Set-Cookie", "") if cookies_hdr: @@ -150,7 +150,7 @@ def _web_test_security_headers(self, target, port): } try: - resp_main = requests.get(base_url, timeout=3, verify=False) + resp_main = requests.get(base_url, timeout=self._target_timeout(3), verify=False) for header, (severity, cwe, owasp, desc) in _HEADER_META.items(): if header not in resp_main.headers: findings_list.append(Finding( @@ -207,7 +207,7 @@ def _web_test_cors_misconfiguration(self, target, port): malicious_origin = "https://attacker.example" resp = requests.get( base_url, - timeout=3, + timeout=self._target_timeout(3), verify=False, headers={"Origin": malicious_origin} ) @@ -276,7 +276,7 @@ def _web_test_open_redirect(self, target, port): redirect_url = base_url.rstrip("/") + f"/login?next={quote(payload, safe=':/')}" resp = requests.get( redirect_url, - timeout=3, + timeout=self._target_timeout(3), verify=False, allow_redirects=False ) @@ -330,7 +330,7 @@ def _web_test_http_methods(self, target, port): if port not in (80, 443): base_url = f"{scheme}://{target}:{port}" try: - resp = requests.options(base_url, timeout=3, verify=False) + resp = requests.options(base_url, timeout=self._target_timeout(3), verify=False) allow = resp.headers.get("Allow", "") if allow: risky = [method for method in ("PUT", "DELETE", "TRACE", "CONNECT") if method in allow.upper()] @@ -408,7 +408,7 @@ def _web_test_csrf(self, target, port): for path in ("/", "/login", "/contact", "/register"): try: - resp = requests.get(base_url + path, timeout=3, verify=False) + resp = requests.get(base_url + path, timeout=self._target_timeout(3), verify=False) if resp.status_code != 200: continue @@ -500,7 +500,7 @@ def _web_test_account_enumeration(self, target, port): try: resp_fake = requests.post( url, data={"username": fake_user, "password": password}, - timeout=3, verify=False, allow_redirects=False, + timeout=self._target_timeout(3), verify=False, allow_redirects=False, ) if resp_fake.status_code == 404: continue @@ -508,7 +508,7 @@ def _web_test_account_enumeration(self, target, port): for real_user in real_candidates: resp_real = requests.post( url, data={"username": real_user, "password": password}, - timeout=3, verify=False, allow_redirects=False, + timeout=self._target_timeout(3), verify=False, allow_redirects=False, ) fake_lower = resp_fake.text.lower() real_lower = resp_real.text.lower() @@ -589,7 +589,7 @@ def _web_test_rate_limiting(self, target, port): for path in login_paths: url = base_url.rstrip("/") + path try: - probe_resp = requests.get(url, timeout=3, verify=False, allow_redirects=False) + probe_resp = requests.get(url, timeout=self._target_timeout(3), verify=False, allow_redirects=False) if probe_resp.status_code == 404: continue @@ -598,7 +598,7 @@ def _web_test_rate_limiting(self, target, port): resp = requests.post( url, data={"username": f"test_user_{i}", "password": password}, - timeout=3, verify=False, allow_redirects=False, + timeout=self._target_timeout(3), verify=False, allow_redirects=False, ) if resp.status_code == 429: rate_limited = True @@ -676,7 +676,7 @@ def _web_test_subresource_integrity(self, target, port): base_url = f"{scheme}://{target}" if port in (80, 443) else f"{scheme}://{target}:{port}" try: - resp = requests.get(base_url, timeout=4, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(4), verify=False) if resp.status_code != 200: return probe_result(findings=findings_list) html = resp.text @@ -765,7 +765,7 @@ def _web_test_mixed_content(self, target, port): base_url = f"https://{target}" if port == 443 else f"https://{target}:{port}" try: - resp = requests.get(base_url, timeout=4, verify=False) + resp = requests.get(base_url, timeout=self._target_timeout(4), verify=False) if resp.status_code != 200: return probe_result(findings=findings_list) html = resp.text diff --git a/extensions/business/cybersec/red_mesh/worker/web/injection.py b/extensions/business/cybersec/red_mesh/worker/web/injection.py index 610c034e7..f86086f23 100644 --- a/extensions/business/cybersec/red_mesh/worker/web/injection.py +++ b/extensions/business/cybersec/red_mesh/worker/web/injection.py @@ -26,7 +26,7 @@ def _run_injection_test(self, target, port, *, params, payloads, check_fn, for payload, needle in payloads: try: url = f"{base_url}?{param}={payload}" - resp = requests.get(url, timeout=3, verify=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False) if check_fn(resp, needle): findings.append(finding_factory(param, payload, resp, url)) break # Found for this param, next param @@ -85,7 +85,7 @@ def _web_test_path_traversal(self, target, port): break try: url = base_url.rstrip("/") + payload_path - resp = requests.get(url, timeout=2, verify=False) + resp = requests.get(url, timeout=self._target_timeout(2), verify=False) if any(n in resp.text for n in unix_needles): findings_list.append(Finding( severity=Severity.CRITICAL, @@ -127,7 +127,7 @@ def _web_test_path_traversal(self, target, port): for payload, needles in payloads_qs: try: url = f"{base_url}?{param}={payload}" - resp = requests.get(url, timeout=2, verify=False) + resp = requests.get(url, timeout=self._target_timeout(2), verify=False) if any(n in resp.text for n in needles): findings_list.append(Finding( severity=Severity.CRITICAL, @@ -189,7 +189,7 @@ def _web_test_xss(self, target, port): break try: url = base_url.rstrip("/") + f"/{payload}" - resp = requests.get(url, timeout=3, verify=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False) if needle in resp.text: findings_list.append(Finding( severity=Severity.HIGH, @@ -313,9 +313,9 @@ def _sqli_error_finding(param, payload, resp, url): url_false = f"{base_url}?{param}=1 AND 1=2" url_base = f"{base_url}?{param}=1" - resp_base = requests.get(url_base, timeout=3, verify=False) - resp_true = requests.get(url_true, timeout=3, verify=False) - resp_false = requests.get(url_false, timeout=3, verify=False) + resp_base = requests.get(url_base, timeout=self._target_timeout(3), verify=False) + resp_true = requests.get(url_true, timeout=self._target_timeout(3), verify=False) + resp_false = requests.get(url_false, timeout=self._target_timeout(3), verify=False) # Baseline should match true, differ from false if (resp_base.status_code == resp_true.status_code and @@ -406,7 +406,7 @@ def _web_test_ssti(self, target, port): # (e.g. "49" naturally appears in many pages) baseline_text = "" try: - baseline_resp = requests.get(base_url, timeout=3, verify=False) + baseline_resp = requests.get(base_url, timeout=self._target_timeout(3), verify=False) baseline_text = baseline_resp.text except Exception: pass @@ -423,12 +423,12 @@ def _web_test_ssti(self, target, port): # For short expected values (e.g. "49"), bracket the payload with # two control requests to catch incrementing counters/timestamps if len(expected) <= 3: - ctrl1 = requests.get(f"{base_url}?{param}=harmless1", timeout=3, verify=False) + ctrl1 = requests.get(f"{base_url}?{param}=harmless1", timeout=self._target_timeout(3), verify=False) url = f"{base_url}?{param}={quote(payload)}" - resp = requests.get(url, timeout=3, verify=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False) if expected in resp.text and payload not in resp.text: if len(expected) <= 3: - ctrl2 = requests.get(f"{base_url}?{param}=harmless2", timeout=3, verify=False) + ctrl2 = requests.get(f"{base_url}?{param}=harmless2", timeout=self._target_timeout(3), verify=False) if expected in ctrl1.text or expected in ctrl2.text: continue findings_list.append(Finding( @@ -455,7 +455,7 @@ def _web_test_ssti(self, target, port): continue try: url = base_url.rstrip("/") + "/" + quote(payload) - resp = requests.get(url, timeout=3, verify=False) + resp = requests.get(url, timeout=self._target_timeout(3), verify=False) if expected in resp.text and payload not in resp.text: findings_list.append(Finding( severity=Severity.CRITICAL, @@ -531,7 +531,7 @@ def _web_test_shellshock(self, target, port): "User-Agent": shellshock_payload, "Referer": shellshock_payload, }, - timeout=4, + timeout=self._target_timeout(4), verify=False, ) if marker in resp.text: @@ -591,7 +591,7 @@ def _web_test_php_cgi(self, target, port): resp = requests.get( base_url, headers={"User-Agentt": "zerodiumsystem('echo REDMESH_PHP_BACKDOOR');"}, - timeout=3, + timeout=self._target_timeout(3), verify=False, ) if "REDMESH_PHP_BACKDOOR" in resp.text: @@ -626,7 +626,7 @@ def _web_test_php_cgi(self, target, port): test_url, data="", headers={"Content-Type": "application/x-www-form-urlencoded"}, - timeout=3, + timeout=self._target_timeout(3), verify=False, ) # Guard: auto_prepend_file output appears at the very start of the @@ -652,7 +652,7 @@ def _web_test_php_cgi(self, target, port): # --- 3. PHP-CGI source disclosure via -s flag --- if not findings_list: try: - resp = requests.get(base_url + "/?%ADs", timeout=3, verify=False) + resp = requests.get(base_url + "/?%ADs", timeout=self._target_timeout(3), verify=False) if "" in resp.text and "= 1.0: + raise RuntimeError("Failed to stop tunnel engine while pausing") + self.sleep(0.01) + # endwhile + self.reset_tunnel_engine() + + nr_commands = len(self.get_start_commands()) + self.start_commands_started = [False] * nr_commands + self.start_commands_finished = [False] * nr_commands + self.start_commands_processes = [None] * nr_commands + self.start_commands_start_time = [None] * nr_commands + self._dauth_pause_teardown_succeeded = True + return + + def on_resume(self): + if not self._is_dauth_server_enabled(): + raise RuntimeError("Cannot resume an ineligible dAuth server") + if not self._dauth_pause_teardown_succeeded: + raise RuntimeError("Cannot resume after an incomplete web app teardown") + # endif + + self.failed = False + self._stop_request_monitor.clear() + self._start_request_monitor_thread() + return + + def on_log_handler(self, text, key=None): + super(DauthManagerPlugin, self).on_log_handler(text, key=key) + if self._is_dauth_server_enabled() and "Uvicorn running on " in text: + self.set_plugin_ready(True) + # endif + return def on_request(self, request): @@ -124,11 +297,34 @@ def on_response(self, method, response): return def process(self): + self._maybe_hsync_dauth_job_secrets() # TODO: this will be re-enabled in the future. if False: self._maybe_log_and_save_tracked_requests() return + def _maybe_hsync_dauth_job_secrets(self): + if not self._is_dauth_server_enabled(): + return None + + now = self.time() + last_sync = getattr(self, "_last_dauth_job_secrets_hsync", None) + if ( + last_sync is not None + and now - last_sync < self.cfg_dauth_job_secrets_hsync_interval + ): + return None + + self._last_dauth_job_secrets_hsync = now + try: + return self.chainstore_hsync( + hkey=DAUTH_JOB_SECRETS_CSTORE_HKEY, + **dauth_registry_write_kwargs(self), + ) + except Exception as exc: + self.P(f"Could not sync dAuth job secrets: {exc}", color="y") + return None + def __get_current_epoch(self): """ Get the current epoch of the node. @@ -212,6 +408,12 @@ def get_auth_data(self, body: dict): } } """ + if not self._is_dauth_server_enabled(): + response = self.__get_response({ + 'error': 'dAuth server is not registered as a dAuth oracle' + }) + return response + try: data = self.process_dauth_request(body) except Exception as e: @@ -224,3 +426,65 @@ def get_auth_data(self, body: dict): **data }) return response + + @BasePlugin.endpoint(method="post") + # /add_secrets + def add_secrets(self, body: dict): + """ + Store a full job secret bundle from a protocol oracle. + + The signed request must include a hex-millisecond timestamp nonce no older + than 120 seconds. + """ + request_nonce = body.get("nonce") if isinstance(body, dict) else None + if not self._is_dauth_server_enabled(): + response = self.__get_response({ + 'error': 'dAuth server is not registered as a dAuth oracle', + 'nonce': request_nonce, + }) + return response + + try: + data = self.process_dauth_add_secrets_request(body) + except Exception as e: + self.P("Error processing add_secrets request: {}".format(e), color='r') + data = { + 'error' : str(e) + } + + response = self.__get_response({ + 'nonce': request_nonce, + **data + }) + return response + + @BasePlugin.endpoint(method="post") + # /get_secrets + def get_secrets(self, body: dict): + """ + Return an encrypted job secret bundle to a current R1FS job runner. + + The signed request must include a hex-millisecond timestamp nonce no older + than 120 seconds. The signed response echoes that nonce. + """ + request_nonce = body.get("nonce") if isinstance(body, dict) else None + if not self._is_dauth_server_enabled(): + response = self.__get_response({ + 'error': 'dAuth server is not registered as a dAuth oracle', + 'nonce': request_nonce, + }) + return response + + try: + data = self.process_dauth_get_secret_request(body) + except Exception as e: + self.P("Error processing get_secrets request: {}".format(e), color='r') + data = { + 'error' : str(e) + } + + response = self.__get_response({ + 'nonce': request_nonce, + **data + }) + return response diff --git a/extensions/business/dauth/dauth_mixin.py b/extensions/business/dauth/dauth_mixin.py index 1c59689f4..7ece280c0 100644 --- a/extensions/business/dauth/dauth_mixin.py +++ b/extensions/business/dauth/dauth_mixin.py @@ -15,6 +15,14 @@ """ +from extensions.business.dauth.dauth_registry import dauth_registry_write_kwargs + + +DAUTH_JOB_SECRETS_CSTORE_HKEY = "DAUTH_JOB_SECRETS" +DEEPLOY_JOBS_CSTORE_HKEY = "DEEPLOY_DEPLOYED_JOBS" +DAUTH_SECRET_REQUEST_MAX_AGE_SECONDS = 120 + + def version_to_int(version): """ Convert a version string to an integer. @@ -159,8 +167,162 @@ def check_if_node_allowed( node_address_eth, self.evm_network, e ) return result, msg - - + + def _verify_signed_dauth_body(self, body): + if not isinstance(body, dict): + raise ValueError("Invalid request body.") + + bcct = self.const.BASE_CT.BCctbase + requester = body.get(bcct.SENDER) + requester_send_eth = body.get(bcct.ETH_SENDER) + if not requester: + raise ValueError("No sender address in request.") + if not requester_send_eth: + raise ValueError("No sender ETH address in request.") + + requester_eth = self.bc.node_address_to_eth_address(requester) + if requester_eth.lower() != requester_send_eth.lower(): + raise ValueError("Sender ETH address and recovered ETH address do not match.") + + verify_data = self.bc.verify(body, return_full_info=True) + if not verify_data.valid: + raise ValueError("Invalid request signature: {}".format(verify_data.message)) + + return requester, requester_eth + + def _is_protocol_oracle_eth(self, node_address_eth): + eth_oracles = self.bc.get_eth_oracles() + if len(eth_oracles) == 0: + raise ValueError("No oracles found - this is a critical issue!") + return node_address_eth.lower() in [addr.lower() for addr in eth_oracles] + + def _validate_dauth_secret_request_nonce(self, body): + """Validate the signed hex-millisecond timestamp nonce.""" + nonce = body.get(self.const.BASE_CT.dAuth.DAUTH_NONCE) + if not isinstance(nonce, str) or not nonce: + raise ValueError("dAuth request nonce is required.") + try: + request_time = int(nonce, 16) / 1000 + except (TypeError, ValueError) as exc: + raise ValueError("dAuth request nonce is invalid.") from exc + + request_age = self.time() - request_time + if request_age < 0: + raise ValueError("dAuth request nonce is from the future.") + if request_age > DAUTH_SECRET_REQUEST_MAX_AGE_SECONDS: + raise ValueError("dAuth request nonce is expired.") + return nonce + + def _normalize_dauth_job_id(self, job_id): + if job_id in [None, ""]: + raise ValueError("Job ID is required.") + return str(job_id) + + def _build_secret_bundle_from_request(self, body, job_id): + job_secrets = body.get("job_secrets") + if not isinstance(job_secrets, dict): + raise ValueError("job_secrets must be a dictionary.") + return { + "job_id": job_id, + "job_secrets": self.deepcopy(job_secrets), + } + + def _save_dauth_job_secret_bundle(self, job_id, secret_bundle): + result = self.chainstore_hset( + hkey=DAUTH_JOB_SECRETS_CSTORE_HKEY, + key=job_id, + value=secret_bundle, + **dauth_registry_write_kwargs(self), + ) + if not result: + raise ValueError(f"Failed to store dAuth secrets for job {job_id}.") + return result + + def _load_dauth_job_secret_bundle(self, job_id): + return self.chainstore_hget( + hkey=DAUTH_JOB_SECRETS_CSTORE_HKEY, + key=job_id, + ) + + def _load_dauth_job_pipeline(self, job_id): + cid = self.chainstore_hget( + hkey=DEEPLOY_JOBS_CSTORE_HKEY, + key=job_id, + ) + if not cid: + return None + return self.r1fs.get_json(cid, show_logs=False) + + def _pipeline_runner_nodes(self, pipeline): + if not isinstance(pipeline, dict): + return [] + specs = pipeline.get("DEEPLOY_SPECS") or pipeline.get("deeploy_specs") or {} + if not isinstance(specs, dict): + return [] + nodes = specs.get("current_target_nodes") or specs.get("CURRENT_TARGET_NODES") or [] + if isinstance(nodes, str): + nodes = [nodes] + if not isinstance(nodes, list): + return [] + return [node for node in nodes if isinstance(node, str) and len(node) > 0] + + def _normalize_node_address_for_compare(self, node_address): + try: + return self.bc.maybe_add_prefix(node_address) + except Exception: + return node_address + + def _is_node_running_dauth_job(self, job_id, node_address): + pipeline = self._load_dauth_job_pipeline(job_id) + runner_nodes = self._pipeline_runner_nodes(pipeline) + requester = self._normalize_node_address_for_compare(node_address) + runner_nodes = [ + self._normalize_node_address_for_compare(node) + for node in runner_nodes + ] + return requester in runner_nodes + + def process_dauth_add_secrets_request(self, body): + requester, requester_eth = self._verify_signed_dauth_body(body) + request_nonce = self._validate_dauth_secret_request_nonce(body) + if not self._is_protocol_oracle_eth(requester_eth): + raise ValueError(f"Sender {requester_eth} is not an oracle.") + + job_id = self._normalize_dauth_job_id(body.get("job_id")) + secret_bundle = self._build_secret_bundle_from_request(body, job_id) + self._save_dauth_job_secret_bundle(job_id, secret_bundle) + self.Pd(f"dAuth stored secret bundle for job {job_id} from oracle {requester}.") + return { + "status": "success", + "job_id": job_id, + self.const.BASE_CT.dAuth.DAUTH_NONCE: request_nonce, + } + + def process_dauth_get_secret_request(self, body): + requester, _ = self._verify_signed_dauth_body(body) + request_nonce = self._validate_dauth_secret_request_nonce(body) + job_id = self._normalize_dauth_job_id(body.get("job_id")) + + if not self._is_node_running_dauth_job(job_id, requester): + raise ValueError(f"Sender {requester} is not running job {job_id}.") + + secret_bundle = self._load_dauth_job_secret_bundle(job_id) + if not isinstance(secret_bundle, dict): + raise ValueError(f"No dAuth secret bundle found for job {job_id}.") + encrypted_secret_bundle = self.bc.encrypt_str( + str_data=self.json_dumps(secret_bundle), + str_recipient=requester, + ) + if not isinstance(encrypted_secret_bundle, str) or not encrypted_secret_bundle: + raise ValueError(f"Failed to encrypt dAuth secrets for job {job_id}.") + + return { + "status": "success", + "job_id": job_id, + self.const.BASE_CT.dAuth.DAUTH_NONCE: request_nonce, + "encrypted_secret_bundle": encrypted_secret_bundle, + } + def chainstore_store_dauth_request( self, node_address : str, @@ -298,11 +460,24 @@ def fill_dauth_data( # end if is_node # end set node tags - # set the supervisor flag if this is identified as an oracle + # set supervisor secrets if this is a protocol oracle; dAuth-only secrets + # are additionally gated by the dAuth oracle registry. if is_node and requester_node_address in oracles: + requester_is_dauth_oracle = False + try: + requester_is_dauth_oracle = self.bc.is_dauth_oracle(node_address_eth=sender_eth_address) + except Exception as e: + self.P( + f"dAuth oracle check failed for {sender_eth_address}; omitting supervisor keys: {e}", + color='r' + ) + # end try + dauth_data["EE_SUPERVISOR"] = True for key in self.cfg_supervisor_keys: if isinstance(key, str) and len(key) > 0: + if key in self.cfg_dauth_oracle_only_supervisor_keys and not requester_is_dauth_oracle: + continue dauth_data[key] = self.os_environ.get(key) # end if # end for @@ -556,4 +731,3 @@ def process_dauth_request(self, body): res = eng.process_dauth_request(request_sdk) # res = eng.process_dauth_request(request_bad) l.P(f"Result:\n{json.dumps(res, indent=2)}") - \ No newline at end of file diff --git a/extensions/business/dauth/dauth_registry.py b/extensions/business/dauth/dauth_registry.py new file mode 100644 index 000000000..a9951972f --- /dev/null +++ b/extensions/business/dauth/dauth_registry.py @@ -0,0 +1,77 @@ +"""dAuth registry lookup and ChainStore routing helpers.""" + + +def resolve_dauth_registry_internal_peers(plugin, eth_oracles): + """Resolve cached registry ETH addresses through current NetMon state.""" + current_eth = plugin.bc.eth_address.lower() + peers = [] + for eth_address in eth_oracles: + internal_address = plugin.bc.eth_addr_to_internal_addr(eth_address) + if internal_address is None and eth_address.lower() == current_eth: + internal_address = plugin.bc.address + if isinstance(internal_address, str) and internal_address: + peers.append(internal_address) + return list(dict.fromkeys(peers)) + + +def load_dauth_registry_snapshot(plugin): + """Load the current dAuth registry and resolve its currently known peers.""" + eth_oracles = plugin.bc.get_eth_dauth_oracles() + eth_oracles = list(dict.fromkeys( + address + for address in eth_oracles or [] + if isinstance(address, str) and address + )) + if not eth_oracles: + raise ValueError("No dAuth oracles are registered.") + + peers = resolve_dauth_registry_internal_peers(plugin, eth_oracles) + if not peers: + raise ValueError("No dAuth registry internal peers are available.") + return peers, eth_oracles + + +def get_cached_dauth_registry_internal_peers(plugin): + """Return the latest cached dAuth oracle internal addresses.""" + eth_oracles = getattr(plugin, "_dauth_registry_eth_oracles", None) + if eth_oracles: + peers = resolve_dauth_registry_internal_peers(plugin, eth_oracles) + if peers: + plugin._dauth_registry_internal_peers = peers + peers = getattr(plugin, "_dauth_registry_internal_peers", None) + if not peers: + raise ValueError("dAuth registry peers are not cached.") + return list(peers) + + +def get_dauth_registry_internal_peers(plugin): + """Return cached peers when available, otherwise load a registry snapshot.""" + if ( + getattr(plugin, "_dauth_registry_eth_oracles", None) + or hasattr(plugin, "_dauth_registry_internal_peers") + ): + return get_cached_dauth_registry_internal_peers(plugin) + peers, _ = load_dauth_registry_snapshot(plugin) + return peers + + +def dauth_registry_write_kwargs(plugin, peers=None): + """Route a ChainStore write exclusively to dAuth registry peers.""" + if peers is None: + peers = get_dauth_registry_internal_peers(plugin) + return { + "extra_peers": list(peers), + "include_default_peers": False, + "include_configured_peers": False, + } + + +def pipeline_registry_write_kwargs(plugin, peers=None): + """Add dAuth peers without disabling normal pipeline metadata peers.""" + if peers is None: + peers = get_dauth_registry_internal_peers(plugin) + return { + "extra_peers": list(peers), + "include_default_peers": True, + "include_configured_peers": True, + } diff --git a/extensions/business/dauth/test_dauth_registry_gating.py b/extensions/business/dauth/test_dauth_registry_gating.py new file mode 100644 index 000000000..eac591298 --- /dev/null +++ b/extensions/business/dauth/test_dauth_registry_gating.py @@ -0,0 +1,942 @@ +from collections import deque +import json +import queue +import threading +import unittest +from copy import deepcopy +from pathlib import Path + +from extensions.business.dauth.dauth_mixin import ( + DAUTH_JOB_SECRETS_CSTORE_HKEY, + DEEPLOY_JOBS_CSTORE_HKEY, + _DauthMixin, +) +from extensions.business.dauth.dauth_registry import load_dauth_registry_snapshot + + +ROOT = Path(__file__).resolve().parents[3] +REQUEST_TIME = 1_700_000_000 +REQUEST_NONCE = hex(REQUEST_TIME * 1000) + + +class _FakeProcess: + + def __init__(self): + self.running = True + + def poll(self): + return None if self.running else 0 + + +class _FakeThread: + + def __init__(self): + self.running = True + + def join(self, timeout=None): # pylint: disable=unused-argument + self.running = False + + def is_alive(self): + return self.running + + +class _FakeBasePlugin: + CONFIG = {"VALIDATION_RULES": {}} + + @staticmethod + def endpoint(method="get", require_token=False): # pylint: disable=unused-argument + def decorator(func): + return func + return decorator + + def on_init(self): + self._base_init_calls += 1 + self._lifecycle_events.append("base_init") + self._stop_request_monitor = threading.Event() + self._request_monitor_thread = _FakeThread() + self._incoming_lock = threading.Lock() + self._incoming_requests = deque(["incoming"]) + self.postponed_requests = deque(["postponed"]) + self._server_queue = queue.Queue() + self._server_queue.put("server") + self.start_commands_started = [True, True] + self.start_commands_finished = [True, True] + self.start_commands_processes = [_FakeProcess(), _FakeProcess()] + self.start_commands_start_time = [10, 20] + self.tunnel_engine_started = True + self.failed = False + return None + + def get_start_commands(self): + return ["uvicorn", "cloudflared"] + + def _maybe_close_start_commands(self): + self._lifecycle_events.append("stop_commands") + for process in self.start_commands_processes: + if process is not None: + process.running = False + # endfor + return + + def _maybe_read_and_stop_all_log_readers(self): + self._lifecycle_events.append("stop_log_readers") + return + + def maybe_stop_tunnel_engine(self): + self._lifecycle_events.append("stop_tunnel") + self.tunnel_engine_started = False + return + + def reset_tunnel_engine(self): + self._lifecycle_events.append("reset_tunnel") + return + + def _start_request_monitor_thread(self): + self._lifecycle_events.append("start_monitor") + self._request_monitor_thread = _FakeThread() + return + + def set_plugin_ready(self, ready=True): + self._is_plugin_ready = ready + return + + def on_log_handler(self, text, key=None): # pylint: disable=unused-argument + self._lifecycle_events.append("log") + return + + +class _FakeDauthMixin: + pass + + +class _FakeNodeTagsMixin: + pass + + +class _FakeRequestTrackingMixin: + pass + + +def _load_dauth_manager_class(): + source_path = ROOT / "extensions" / "business" / "dauth" / "dauth_manager.py" + source = source_path.read_text(encoding="utf-8") + source = source.replace( + "from extensions.business.mixins.node_tags_mixin import _NodeTagsMixin\n", + "", + ) + source = source.replace( + "from naeural_core.business.default.web_app.supervisor_fast_api_web_app import SupervisorFastApiWebApp as BasePlugin\n", + "", + ) + source = source.replace( + "from extensions.business.mixins.request_tracking_mixin import _RequestTrackingMixin\n", + "", + ) + source = source.replace( + "from extensions.business.dauth.dauth_mixin import (\n" + " DAUTH_JOB_SECRETS_CSTORE_HKEY,\n" + " _DauthMixin,\n" + ")\n", + "", + ) + source = source.replace( + "from extensions.business.dauth.dauth_registry import (\n" + " dauth_registry_write_kwargs,\n" + " load_dauth_registry_snapshot,\n" + ")\n", + "", + ) + namespace = { + "BasePlugin": _FakeBasePlugin, + "DAUTH_JOB_SECRETS_CSTORE_HKEY": DAUTH_JOB_SECRETS_CSTORE_HKEY, + "_DauthMixin": _FakeDauthMixin, + "_NodeTagsMixin": _FakeNodeTagsMixin, + "_RequestTrackingMixin": _FakeRequestTrackingMixin, + "dauth_registry_write_kwargs": ( + lambda plugin: { + "extra_peers": list(plugin._dauth_registry_internal_peers), + "include_default_peers": False, + "include_configured_peers": False, + } + ), + "load_dauth_registry_snapshot": load_dauth_registry_snapshot, + "__name__": "loaded_dauth_manager", + } + exec(compile(source, str(source_path), "exec"), namespace) # noqa: S102 + return namespace["DauthManagerPlugin"] + + +DauthManagerPlugin = _load_dauth_manager_class() + + +class _FakeDauthConst: + DAUTH_NONCE = "nonce" + DAUTH_ENV_KEYS_PREFIX = "EE_" + DAUTH_WHITELIST = "DAUTH_WHITELIST" + + +class _FakeBCBaseConst: + SENDER = "EE_SENDER" + ETH_SENDER = "EE_ETH_SENDER" + + +class _FakeBaseConst: + BCctbase = _FakeBCBaseConst + dAuth = _FakeDauthConst + + +class _FakeConst: + BASE_CT = _FakeBaseConst + ADMIN_PIPELINE = { + "DAUTH_MANAGER": { + "AUTH_ENV_KEYS": [], + "AUTH_NODE_ENV_KEYS": [], + "AUTH_PREDEFINED_KEYS": {}, + } + } + + +class _FakeBC: + + def __init__(self, *, dauth_oracle=True, protocol_oracles=None, valid_signature=True): + self.dauth_oracle = dauth_oracle + self.protocol_oracles = protocol_oracles or ["node-oracle"] + self.valid_signature = valid_signature + self.encrypt_calls = [] + self.node_eth = { + "node-oracle": "0xORACLE", + "node-runner": "0xRUNNER", + "node-other": "0xOTHER", + } + + def get_oracles(self, include_eth_addrs=False): + names = ["Oracle"] * len(self.protocol_oracles) + eth_addresses = ["0xDAUTH"] * len(self.protocol_oracles) + if include_eth_addrs: + return self.protocol_oracles, names, eth_addresses + return self.protocol_oracles, names + + def get_whitelist_with_names(self): + return [], [] + + def is_dauth_oracle(self, node_address_eth=None): # pylint: disable=unused-argument + if isinstance(self.dauth_oracle, Exception): + raise self.dauth_oracle + return self.dauth_oracle + + def get_eth_oracles(self): + return [self.node_eth.get(node, "0xORACLE") for node in self.protocol_oracles] + + def node_address_to_eth_address(self, node_address): + return self.node_eth[node_address] + + def verify(self, body, return_full_info=False): # pylint: disable=unused-argument + class _VerifyData: + pass + + data = _VerifyData() + data.valid = self.valid_signature + data.message = "ok" if self.valid_signature else "bad signature" + return data + + def maybe_add_prefix(self, node_address): + if node_address.startswith("0xai_"): + return node_address + return "0xai_" + node_address + + def encrypt_str(self, str_data, str_recipient): + self.encrypt_calls.append((str_data, str_recipient)) + return "encrypted-secret-bundle" + + +class _FakeR1FS: + + def __init__(self, data): + self.data = data + + def get_json(self, cid, show_logs=False): # pylint: disable=unused-argument + return self.data[cid] + + +class _DauthHarness(_DauthMixin): + pass + + +def _make_dauth_harness(*, dauth_oracle=True, protocol_oracles=None, valid_signature=True): + plugin = _DauthHarness() + plugin.const = _FakeConst + plugin.bc = _FakeBC( + dauth_oracle=dauth_oracle, + protocol_oracles=protocol_oracles, + valid_signature=valid_signature, + ) + plugin.deepcopy = deepcopy + plugin.json_dumps = json.dumps + plugin.time = lambda: REQUEST_TIME + plugin._chainstore = {} + plugin._r1fs_data = {} + plugin.r1fs = _FakeR1FS(plugin._r1fs_data) + plugin.evm_network = "devnet" + plugin.cfg_auth_env_keys = [] + plugin.cfg_auth_node_env_keys = [] + plugin.cfg_auth_predefined_keys = {} + plugin.cfg_supervisor_keys = [ + "EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER", + "EE_NGROK_EDGE_LABEL_DAUTH_MANAGER", + "EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER", + ] + plugin.cfg_dauth_oracle_only_supervisor_keys = [ + "EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER", + ] + plugin.cfg_comms_host_key = "EE_MQTT_HOST" + plugin.cfg_comms_host_seed_key = "EE_MQTT_HOST_SEED" + plugin.cfg_dauth_log_response = False + plugin.cfg_dauth_verbose = False + plugin.os_environ = { + "EE_MQTT_HOST_SEED": "mqtt-a", + "EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER": "cloudflare-secret", + "EE_NGROK_EDGE_LABEL_DAUTH_MANAGER": "ngrok-label", + "EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER": "deeploy-secret", + } + plugin.fetch_node_tags = lambda node_address_eth=None: {} + plugin.P = lambda *args, **kwargs: None + plugin.Pd = lambda *args, **kwargs: None + plugin._dauth_registry_internal_peers = ["node-oracle"] + plugin.chainstore_hset = lambda hkey, key, value, **kwargs: plugin._chainstore.__setitem__( + (hkey, str(key)), + deepcopy(value), + ) or True + plugin.chainstore_hget = lambda hkey, key: plugin._chainstore.get((hkey, str(key))) + return plugin + + +class DauthRegistrySecretGatingTests(unittest.TestCase): + + def test_supervisor_keys_are_sent_to_protocol_oracles_registered_for_dauth(self): + plugin = _make_dauth_harness(dauth_oracle=True) + + data = plugin.fill_dauth_data( + dauth_data={}, + requester_node_address="node-oracle", + is_node=True, + sender_eth_address="0xDAUTH", + ) + + self.assertTrue(data["EE_SUPERVISOR"]) + self.assertEqual(data["EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER"], "cloudflare-secret") + self.assertEqual(data["EE_NGROK_EDGE_LABEL_DAUTH_MANAGER"], "ngrok-label") + self.assertEqual(data["EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER"], "deeploy-secret") + + def test_only_dauth_token_is_omitted_for_protocol_oracles_not_registered_for_dauth(self): + plugin = _make_dauth_harness(dauth_oracle=False) + + data = plugin.fill_dauth_data( + dauth_data={}, + requester_node_address="node-oracle", + is_node=True, + sender_eth_address="0xOTHER", + ) + + self.assertTrue(data["EE_SUPERVISOR"]) + self.assertNotIn("EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER", data) + self.assertEqual(data["EE_NGROK_EDGE_LABEL_DAUTH_MANAGER"], "ngrok-label") + self.assertEqual(data["EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER"], "deeploy-secret") + + def test_supervisor_keys_still_require_protocol_oracle_membership(self): + plugin = _make_dauth_harness(dauth_oracle=True, protocol_oracles=["node-other"]) + + data = plugin.fill_dauth_data( + dauth_data={}, + requester_node_address="node-oracle", + is_node=True, + sender_eth_address="0xDAUTH", + ) + + self.assertFalse(data["EE_SUPERVISOR"]) + self.assertNotIn("EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER", data) + self.assertNotIn("EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER", data) + + def test_dauth_token_fails_closed_when_dauth_registry_check_fails(self): + plugin = _make_dauth_harness(dauth_oracle=RuntimeError("registry unavailable")) + + data = plugin.fill_dauth_data( + dauth_data={}, + requester_node_address="node-oracle", + is_node=True, + sender_eth_address="0xDAUTH", + ) + + self.assertTrue(data["EE_SUPERVISOR"]) + self.assertNotIn("EE_CLOUDFLARE_TOKEN_DAUTH_MANAGER", data) + self.assertEqual(data["EE_NGROK_EDGE_LABEL_DAUTH_MANAGER"], "ngrok-label") + self.assertEqual(data["EE_CLOUDFLARE_TOKEN_DEEPLOY_MANAGER"], "deeploy-secret") + + +class DauthJobSecretEndpointTests(unittest.TestCase): + + def test_secret_request_nonce_accepts_only_last_120_seconds(self): + plugin = _make_dauth_harness() + + self.assertEqual( + plugin._validate_dauth_secret_request_nonce({"nonce": REQUEST_NONCE}), + REQUEST_NONCE, + ) + boundary_nonce = hex(int((REQUEST_TIME - 120) * 1000)) + self.assertEqual( + plugin._validate_dauth_secret_request_nonce({"nonce": boundary_nonce}), + boundary_nonce, + ) + + invalid_nonces = ( + ({}, "required"), + ({"nonce": "not-hex"}, "invalid"), + ({"nonce": hex(int((REQUEST_TIME + 1) * 1000))}, "future"), + ({"nonce": hex(int((REQUEST_TIME - 121) * 1000))}, "expired"), + ) + for body, message in invalid_nonces: + with self.subTest(body=body): + with self.assertRaisesRegex(ValueError, message): + plugin._validate_dauth_secret_request_nonce(body) + + def test_add_secrets_allows_protocol_oracle_and_overwrites_bundle(self): + plugin = _make_dauth_harness(protocol_oracles=["node-oracle"]) + plugin._chainstore[(DAUTH_JOB_SECRETS_CSTORE_HKEY, "7")] = { + "job_id": "7", + "old": True, + } + body = { + "EE_SENDER": "node-oracle", + "EE_ETH_SENDER": "0xORACLE", + "nonce": REQUEST_NONCE, + "job_id": 7, + "job_secrets": { + "plugins": { + "CONTAINER_APP_RUNNER": [{ + "instance_conf": { + "ENV": { + "API_KEY": "secret", + }, + }, + }], + }, + }, + } + + response = plugin.process_dauth_add_secrets_request(body) + + self.assertEqual(response["status"], "success") + self.assertEqual(response["job_id"], "7") + self.assertEqual(response["nonce"], REQUEST_NONCE) + self.assertEqual( + plugin._chainstore[(DAUTH_JOB_SECRETS_CSTORE_HKEY, "7")], + { + "job_id": "7", + "job_secrets": body["job_secrets"], + }, + ) + + def test_add_secrets_rejects_non_oracle_writer(self): + plugin = _make_dauth_harness(protocol_oracles=["node-oracle"]) + body = { + "EE_SENDER": "node-runner", + "EE_ETH_SENDER": "0xRUNNER", + "nonce": REQUEST_NONCE, + "job_id": "7", + "job_secrets": {"plugins": {}}, + } + + with self.assertRaisesRegex(ValueError, "not an oracle"): + plugin.process_dauth_add_secrets_request(body) + + self.assertNotIn((DAUTH_JOB_SECRETS_CSTORE_HKEY, "7"), plugin._chainstore) + + def test_add_secrets_rejects_expired_nonce_before_write(self): + plugin = _make_dauth_harness() + body = { + "EE_SENDER": "node-oracle", + "EE_ETH_SENDER": "0xORACLE", + "nonce": hex(int((REQUEST_TIME - 121) * 1000)), + "job_id": "7", + "job_secrets": {"plugins": {}}, + } + + with self.assertRaisesRegex(ValueError, "nonce is expired"): + plugin.process_dauth_add_secrets_request(body) + + self.assertNotIn((DAUTH_JOB_SECRETS_CSTORE_HKEY, "7"), plugin._chainstore) + + def test_add_secrets_rejects_invalid_signature(self): + plugin = _make_dauth_harness(valid_signature=False) + body = { + "EE_SENDER": "node-oracle", + "EE_ETH_SENDER": "0xORACLE", + "nonce": REQUEST_NONCE, + "job_id": "7", + "job_secrets": {"plugins": {}}, + } + + with self.assertRaisesRegex(ValueError, "Invalid request signature"): + plugin.process_dauth_add_secrets_request(body) + + self.assertNotIn((DAUTH_JOB_SECRETS_CSTORE_HKEY, "7"), plugin._chainstore) + + def test_add_secrets_rejects_legacy_plugin_secrets_shape(self): + plugin = _make_dauth_harness() + body = { + "EE_SENDER": "node-oracle", + "EE_ETH_SENDER": "0xORACLE", + "nonce": REQUEST_NONCE, + "job_id": "7", + "plugin_secrets": {"plugins": {}}, + } + + with self.assertRaisesRegex(ValueError, "job_secrets must be a dictionary"): + plugin.process_dauth_add_secrets_request(body) + + self.assertNotIn((DAUTH_JOB_SECRETS_CSTORE_HKEY, "7"), plugin._chainstore) + + def test_get_secrets_returns_bundle_for_node_running_job_from_r1fs_pipeline(self): + plugin = _make_dauth_harness() + bundle = { + "job_id": "7", + "job_secrets": { + "plugins": { + "CONTAINER_APP_RUNNER": [{ + "instance_conf": { + "ENV": { + "API_KEY": "secret", + }, + }, + }], + }, + }, + } + plugin._chainstore[(DAUTH_JOB_SECRETS_CSTORE_HKEY, "7")] = bundle + plugin._chainstore[(DEEPLOY_JOBS_CSTORE_HKEY, "7")] = "cid-7" + plugin._r1fs_data["cid-7"] = { + "deeploy_specs": { + "current_target_nodes": ["node-runner"], + }, + } + body = { + "EE_SENDER": "node-runner", + "EE_ETH_SENDER": "0xRUNNER", + "nonce": REQUEST_NONCE, + "job_id": "7", + } + + response = plugin.process_dauth_get_secret_request(body) + + self.assertEqual(response["status"], "success") + self.assertEqual(response["job_id"], "7") + self.assertEqual(response["nonce"], REQUEST_NONCE) + self.assertEqual( + response["encrypted_secret_bundle"], + "encrypted-secret-bundle", + ) + self.assertNotIn("secret_bundle", response) + self.assertEqual( + plugin.bc.encrypt_calls, + [(json.dumps(bundle), "node-runner")], + ) + + def test_get_secrets_rejects_node_not_running_job(self): + plugin = _make_dauth_harness() + plugin._chainstore[(DAUTH_JOB_SECRETS_CSTORE_HKEY, "7")] = { + "job_id": "7", + "job_secrets": {"plugins": {}}, + } + plugin._chainstore[(DEEPLOY_JOBS_CSTORE_HKEY, "7")] = "cid-7" + plugin._r1fs_data["cid-7"] = { + "DEEPLOY_SPECS": { + "current_target_nodes": ["node-runner"], + }, + } + body = { + "EE_SENDER": "node-other", + "EE_ETH_SENDER": "0xOTHER", + "nonce": REQUEST_NONCE, + "job_id": "7", + } + + with self.assertRaisesRegex(ValueError, "not running job"): + plugin.process_dauth_get_secret_request(body) + + +class DauthServerRegistryGateTests(unittest.TestCase): + + def _make_manager(self, *, dauth_oracle): + class _ManagerBC: + def __init__(self, result): + self.result = result + self.calls = 0 + + def get_eth_dauth_oracles(self): + self.calls += 1 + if isinstance(self.result, Exception): + raise self.result + if isinstance(self.result, list): + return self.result + return ["0xNODE", "0xPEER"] if self.result else ["0xPEER"] + + def eth_addr_to_internal_addr(self, eth_address): + return { + "0xnode": "node-address", + "0xpeer": "peer-address", + "0xnew": "new-peer-address", + }.get(eth_address.lower()) + + plugin = DauthManagerPlugin.__new__(DauthManagerPlugin) + plugin.bc = _ManagerBC(dauth_oracle) + plugin._dauth_server_enabled = None + plugin._dauth_server_enabled_message = None + plugin._dauth_web_app_initialized = False + plugin._dauth_pause_teardown_succeeded = True + plugin._base_init_calls = 0 + plugin._lifecycle_events = [] + plugin._messages = [] + plugin._now = 100.0 + plugin.P = lambda msg, *args, **kwargs: plugin._messages.append(msg) + plugin.time = lambda: plugin._now + plugin.sleep = lambda seconds: setattr(plugin, "_now", plugin._now + seconds) + plugin.cfg_auth_env_keys = [] + plugin.cfg_auth_predefined_keys = {} + plugin._init_request_tracking = lambda: None + plugin.bc.address = "node-address" + plugin.bc.eth_address = "0xNODE" + plugin._dauth_registry_eth_oracles = None + plugin._dauth_registry_internal_peers = None + plugin._last_dauth_registry_refresh = None + plugin._dauth_registry_refresh_failed = False + plugin._last_dauth_job_secrets_hsync = None + plugin.cfg_dauth_job_secrets_hsync_interval = 10 * 60 + plugin.cfg_dauth_registry_refresh_interval = 60 * 60 + plugin.cfg_dauth_registry_refresh_retry_interval = 60 + plugin._is_plugin_ready = None + plugin._hsync_calls = [] + plugin.chainstore_hsync = lambda **kwargs: plugin._hsync_calls.append(kwargs) or { + "hkey": kwargs["hkey"], + } + plugin._DauthManagerPlugin__get_response = lambda data: data + return plugin + + def test_secret_endpoint_errors_echo_request_nonce(self): + plugin = self._make_manager(dauth_oracle=True) + plugin._dauth_server_enabled = True + plugin.process_dauth_add_secrets_request = lambda body: (_ for _ in ()).throw( + ValueError("add failed") + ) + plugin.process_dauth_get_secret_request = lambda body: (_ for _ in ()).throw( + ValueError("get failed") + ) + body = {"nonce": REQUEST_NONCE} + + add_response = plugin.add_secrets(body) + get_response = plugin.get_secrets(body) + + self.assertEqual(add_response["nonce"], REQUEST_NONCE) + self.assertEqual(add_response["error"], "add failed") + self.assertEqual(get_response["nonce"], REQUEST_NONCE) + self.assertEqual(get_response["error"], "get failed") + + def test_registry_lookup_is_cached_between_hourly_lifecycle_refreshes(self): + plugin = self._make_manager(dauth_oracle=True) + + plugin.on_init() + + self.assertFalse(plugin.should_pause()) + self.assertFalse(plugin.should_pause()) + self.assertTrue(plugin.should_resume()) + self.assertTrue(plugin.should_resume()) + self.assertTrue(plugin._check_dauth_server_enabled_on_start()) # pylint: disable=protected-access + self.assertEqual(plugin.bc.calls, 1) + self.assertEqual(plugin._base_init_calls, 1) + self.assertEqual(plugin._lifecycle_events, ["base_init"]) + self.assertEqual( + plugin._dauth_registry_internal_peers, + ["node-address", "peer-address"], + ) + + plugin._now += (60 * 60) - 1 + self.assertFalse(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 1) + + plugin._now += 1 + self.assertFalse(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 2) + + def test_secret_hsync_runs_at_startup_and_every_ten_minutes_on_cached_peers(self): + plugin = self._make_manager(dauth_oracle=True) + + plugin.on_init() + plugin.process() + plugin._now += (10 * 60) - 1 + plugin.process() + plugin._now += 1 + plugin.process() + + self.assertEqual(plugin.bc.calls, 1) + self.assertEqual(len(plugin._hsync_calls), 2) + for call in plugin._hsync_calls: + self.assertEqual(call["hkey"], DAUTH_JOB_SECRETS_CSTORE_HKEY) + self.assertEqual(call["extra_peers"], ["node-address", "peer-address"]) + self.assertFalse(call["include_default_peers"]) + self.assertFalse(call["include_configured_peers"]) + + def test_secret_hsync_failure_waits_until_next_interval(self): + plugin = self._make_manager(dauth_oracle=True) + attempts = [] + + def fail_hsync(**kwargs): + attempts.append(kwargs) + raise ValueError("sync unavailable") + + plugin.chainstore_hsync = fail_hsync + plugin.on_init() + plugin.process() + plugin._now += 10 * 60 + plugin.process() + + self.assertEqual(len(attempts), 2) + self.assertTrue(any("sync unavailable" in message for message in plugin._messages)) + + def test_false_startup_lookup_fails_closed_and_tears_down_fastapi(self): + plugin = self._make_manager(dauth_oracle=False) + + plugin.on_init() + + self.assertTrue(plugin.should_pause()) + self.assertFalse(plugin.should_resume()) + self.assertEqual(plugin.bc.calls, 1) + self.assertEqual(plugin._hsync_calls, []) + self.assertTrue(plugin._stop_request_monitor.is_set()) + self.assertFalse(plugin._request_monitor_thread.is_alive()) + self.assertEqual(plugin.start_commands_started, [False, False]) + self.assertEqual(plugin.start_commands_finished, [False, False]) + self.assertEqual(plugin.start_commands_processes, [None, None]) + self.assertEqual(plugin.start_commands_start_time, [None, None]) + self.assertEqual(list(plugin._incoming_requests), []) + self.assertEqual(list(plugin.postponed_requests), []) + self.assertTrue(plugin._server_queue.empty()) + self.assertEqual( + plugin._lifecycle_events, + [ + "base_init", + "stop_commands", + "stop_log_readers", + "stop_tunnel", + "reset_tunnel", + ], + ) + + def test_startup_lookup_error_fails_closed_and_retries_after_one_minute(self): + plugin = self._make_manager(dauth_oracle=RuntimeError("registry unavailable")) + + plugin.on_init() + + self.assertTrue(plugin.should_pause()) + self.assertFalse(plugin.should_resume()) + self.assertTrue(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 1) + self.assertEqual(plugin._dauth_server_enabled_message, "registry unavailable") + self.assertTrue(plugin._dauth_pause_teardown_succeeded) + + plugin._now += 59 + self.assertTrue(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 1) + + plugin._now += 1 + self.assertTrue(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 2) + + def test_hourly_refresh_revokes_server_and_secret_replication(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + plugin.bc.result = False + + plugin._now += (60 * 60) - 1 + self.assertFalse(plugin.should_pause()) + self.assertEqual(plugin.bc.calls, 1) + + plugin._now += 1 + self.assertTrue(plugin.should_pause()) + plugin.on_pause() + self.assertEqual(plugin.bc.calls, 2) + self.assertIsNone(plugin._dauth_registry_eth_oracles) + self.assertIsNone(plugin._dauth_registry_internal_peers) + self.assertTrue(plugin._stop_request_monitor.is_set()) + self.assertEqual(plugin.start_commands_processes, [None, None]) + + hsync_calls = len(plugin._hsync_calls) + plugin._now += 10 * 60 + plugin.process() + self.assertEqual(len(plugin._hsync_calls), hsync_calls) + + def test_hourly_refresh_replaces_removed_replication_peers(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + plugin.bc.result = ["0xNODE", "0xNEW"] + + plugin._now += 60 * 60 + self.assertFalse(plugin.should_pause()) + + self.assertEqual(plugin.bc.calls, 2) + self.assertEqual(plugin._dauth_registry_eth_oracles, ["0xNODE", "0xNEW"]) + self.assertEqual( + plugin._dauth_registry_internal_peers, + ["node-address", "new-peer-address"], + ) + plugin.process() + self.assertEqual( + plugin._hsync_calls[-1]["extra_peers"], + ["node-address", "new-peer-address"], + ) + + def test_hourly_refresh_allows_newly_registered_server_to_resume(self): + plugin = self._make_manager(dauth_oracle=False) + plugin.on_init() + plugin.bc.result = True + + plugin._now += 60 * 60 + self.assertTrue(plugin.should_resume()) + + self.assertEqual(plugin.bc.calls, 2) + self.assertEqual( + plugin._dauth_registry_internal_peers, + ["node-address", "peer-address"], + ) + + def test_hourly_refresh_error_revokes_server_and_clears_peers(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + plugin.bc.result = RuntimeError("registry unavailable") + + plugin._now += 60 * 60 + self.assertTrue(plugin.should_pause()) + + self.assertEqual(plugin.bc.calls, 2) + self.assertEqual(plugin._dauth_server_enabled_message, "registry unavailable") + self.assertIsNone(plugin._dauth_registry_eth_oracles) + self.assertIsNone(plugin._dauth_registry_internal_peers) + + plugin.bc.result = True + plugin._now += 59 + self.assertFalse(plugin.should_resume()) + self.assertEqual(plugin.bc.calls, 2) + + plugin._now += 1 + self.assertTrue(plugin.should_resume()) + self.assertEqual(plugin.bc.calls, 3) + + def test_pause_tears_down_and_resume_restarts_only_request_monitor(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + plugin._incoming_requests.append("second-incoming") + plugin.postponed_requests.append("second-postponed") + plugin._server_queue.put("second-server") + + plugin.on_pause() + + self.assertTrue(plugin._dauth_pause_teardown_succeeded) + self.assertFalse(plugin._is_plugin_ready) + self.assertEqual(plugin.start_commands_processes, [None, None]) + self.assertEqual(list(plugin._incoming_requests), []) + self.assertEqual(list(plugin.postponed_requests), []) + self.assertTrue(plugin._server_queue.empty()) + + plugin.failed = True + plugin.on_resume() + + self.assertFalse(plugin.failed) + self.assertFalse(plugin._stop_request_monitor.is_set()) + self.assertTrue(plugin._request_monitor_thread.is_alive()) + self.assertFalse(plugin._is_plugin_ready) + self.assertEqual(plugin._lifecycle_events[-1], "start_monitor") + self.assertEqual(plugin.bc.calls, 1) + + plugin.on_log_handler("Uvicorn running on http://0.0.0.0:1234 (Press CTRL+C to quit)") + self.assertTrue(plugin._is_plugin_ready) + + def test_ineligible_server_cannot_resume(self): + plugin = self._make_manager(dauth_oracle=False) + plugin.on_init() + + with self.assertRaisesRegex(RuntimeError, "ineligible dAuth server"): + plugin.on_resume() + + self.assertTrue(plugin._stop_request_monitor.is_set()) + self.assertFalse(plugin._request_monitor_thread.is_alive()) + self.assertNotIn("start_monitor", plugin._lifecycle_events) + + def test_failed_teardown_is_verified_and_blocks_resume(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + running_processes = list(plugin.start_commands_processes) + plugin._maybe_close_start_commands = lambda: None + + with self.assertRaisesRegex(RuntimeError, "Failed to stop start commands"): + plugin.on_pause() + + self.assertFalse(plugin._dauth_pause_teardown_succeeded) + self.assertEqual(plugin.start_commands_processes, running_processes) + with self.assertRaisesRegex(RuntimeError, "incomplete web app teardown"): + plugin.on_resume() + + def test_failed_tunnel_stop_is_verified_before_state_is_reset(self): + plugin = self._make_manager(dauth_oracle=True) + plugin.on_init() + plugin.maybe_stop_tunnel_engine = lambda: plugin._lifecycle_events.append( + "stop_tunnel" + ) + + with self.assertRaisesRegex(RuntimeError, "Failed to stop tunnel engine"): + plugin.on_pause() + + self.assertFalse(plugin._dauth_pause_teardown_succeeded) + self.assertTrue(plugin.tunnel_engine_started) + self.assertNotIn("reset_tunnel", plugin._lifecycle_events) + + def test_pause_callback_before_on_init_is_safe(self): + plugin = self._make_manager(dauth_oracle=True) + + plugin.on_pause() + + self.assertEqual(plugin._lifecycle_events, []) + self.assertEqual(plugin.bc.calls, 0) + + def test_initially_disabled_then_enabled_ineligible_init_stays_torn_down(self): + plugin = self._make_manager(dauth_oracle=False) + + plugin.on_pause() + plugin.on_init() + remains_stopped = not plugin.should_resume() + if not remains_stopped: + plugin.on_resume() + # endif + + self.assertTrue(remains_stopped) + self.assertTrue(plugin.should_pause()) + self.assertTrue(plugin._stop_request_monitor.is_set()) + self.assertFalse(plugin._request_monitor_thread.is_alive()) + self.assertEqual(plugin.start_commands_processes, [None, None]) + self.assertEqual(plugin.bc.calls, 1) + + def test_endpoint_authorization_uses_cached_startup_result(self): + plugin = self._make_manager(dauth_oracle=False) + plugin.on_init() + plugin._DauthManagerPlugin__get_response = lambda data: data + plugin.process_dauth_request = lambda body: self.fail( + f"process_dauth_request unexpectedly called with {body}" + ) + + response = plugin.get_auth_data({"nonce": "value"}) + + self.assertEqual( + response, + {"error": "dAuth server is not registered as a dAuth oracle"}, + ) + self.assertEqual(plugin.bc.calls, 1) + + +if __name__ == "__main__": + unittest.main() diff --git a/extensions/business/dauth/test_dauth_secret_routing.py b/extensions/business/dauth/test_dauth_secret_routing.py new file mode 100644 index 000000000..6b4fe8432 --- /dev/null +++ b/extensions/business/dauth/test_dauth_secret_routing.py @@ -0,0 +1,105 @@ +import unittest + +from extensions.business.dauth.dauth_mixin import _DauthMixin +from extensions.business.dauth.dauth_registry import ( + dauth_registry_write_kwargs, + load_dauth_registry_snapshot, + pipeline_registry_write_kwargs, +) + + +class _RegistryBCStub: + address = "node-local" + eth_address = "0xLOCAL" + + def __init__(self, remote_available=True): + self.calls = 0 + self.remote_available = remote_available + + def get_eth_dauth_oracles(self): + self.calls += 1 + return ["0xLOCAL", "0xREMOTE", "0xUNKNOWN"] + + def eth_addr_to_internal_addr(self, eth_address): + if self.remote_available and eth_address == "0xREMOTE": + return "node-remote" + return None + + +class _DauthStub(_DauthMixin): + def __init__(self): + self._dauth_registry_internal_peers = ["dauth-a", "dauth-b"] + self.writes = [] + + def chainstore_hset(self, **kwargs): + self.writes.append(kwargs) + return True + + +class DauthSecretRoutingTests(unittest.TestCase): + def test_registry_snapshot_uses_one_rpc_and_maps_every_known_peer(self): + class _Plugin: + bc = _RegistryBCStub() + + peers, eth_oracles = load_dauth_registry_snapshot(_Plugin()) + + self.assertEqual(_Plugin.bc.calls, 1) + self.assertEqual(peers, ["node-local", "node-remote"]) + self.assertEqual(eth_oracles, ["0xLOCAL", "0xREMOTE", "0xUNKNOWN"]) + + def test_routing_refreshes_internal_mappings_without_another_rpc(self): + class _Plugin: + bc = _RegistryBCStub(remote_available=False) + + plugin = _Plugin() + peers, eth_oracles = load_dauth_registry_snapshot(plugin) + plugin._dauth_registry_eth_oracles = eth_oracles + plugin._dauth_registry_internal_peers = peers + plugin.bc.remote_available = True + + routing = dauth_registry_write_kwargs(plugin) + + self.assertEqual(plugin.bc.calls, 1) + self.assertEqual(routing["extra_peers"], ["node-local", "node-remote"]) + + def test_secret_storage_targets_only_cached_dauth_registry_peers(self): + plugin = _DauthStub() + + plugin._save_dauth_job_secret_bundle( + "7", + {"job_id": "7", "job_secrets": {}}, + ) + + write = plugin.writes[0] + self.assertEqual(write["extra_peers"], ["dauth-a", "dauth-b"]) + self.assertFalse(write["include_default_peers"]) + self.assertFalse(write["include_configured_peers"]) + + def test_explicit_peer_builders_support_non_manager_callers(self): + plugin = object() + + secret_routing = dauth_registry_write_kwargs(plugin, peers=["dauth-a"]) + pipeline_routing = pipeline_registry_write_kwargs(plugin, peers=["dauth-a"]) + + self.assertEqual(secret_routing["extra_peers"], ["dauth-a"]) + self.assertFalse(secret_routing["include_default_peers"]) + self.assertFalse(secret_routing["include_configured_peers"]) + self.assertEqual(pipeline_routing["extra_peers"], ["dauth-a"]) + self.assertTrue(pipeline_routing["include_default_peers"]) + self.assertTrue(pipeline_routing["include_configured_peers"]) + + def test_secret_storage_fails_without_startup_cached_peers(self): + plugin = _DauthStub() + plugin._dauth_registry_internal_peers = [] + + with self.assertRaisesRegex(ValueError, "not cached"): + plugin._save_dauth_job_secret_bundle( + "7", + {"job_id": "7", "job_secrets": {}}, + ) + + self.assertEqual(plugin.writes, []) + + +if __name__ == "__main__": + unittest.main() diff --git a/ver.py b/ver.py index c239b235f..90d83e448 100644 --- a/ver.py +++ b/ver.py @@ -1 +1 @@ -__VER__ = '2.10.400' +__VER__ = '2.10.410'