diff --git a/pyproject.toml b/pyproject.toml index dc5b8b78..63ab355d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -63,6 +63,7 @@ Issues = "https://github.com/zenitraM/malla/issues" malla-web = "malla.web_ui:main" malla-web-gunicorn = "malla.wsgi:main" malla-capture = "malla.mqtt_capture:main" +malla-backfill-traceroutes = "malla.backfill_traceroutes:main" [project.optional-dependencies] dev = [ diff --git a/src/malla/backfill_traceroutes.py b/src/malla/backfill_traceroutes.py new file mode 100644 index 00000000..975611da --- /dev/null +++ b/src/malla/backfill_traceroutes.py @@ -0,0 +1,146 @@ +"""Explicit, resumable preparation of stored traceroute packets.""" + +import argparse +import json +import sqlite3 +import sys +from collections.abc import Callable +from contextlib import closing +from pathlib import Path + +from .database.traceroute_schema import TRACEROUTE_PREDICATE, ensure_traceroute_schema +from .database.traceroutes import ( + PARSER_VERSION, + inspect_traceroutes, + write_traceroute, +) + + +def prepare_traceroutes( + conn: sqlite3.Connection, + batch_size: int = 1000, + progress: Callable[[int], None] | None = None, +) -> dict: + """Fill missing/pending/older-version records, then validate under a write lock. + + Each batch commits independently. An ID boundary keeps continuous capture + from extending the run indefinitely; captures beyond it use the same writer. + Final validation includes any raw-only imports that arrived during the run. + """ + if batch_size < 1: + raise ValueError("batch_size must be positive") + conn.row_factory = sqlite3.Row + with conn: + conn.execute("BEGIN IMMEDIATE") + cursor = conn.cursor() + ensure_traceroute_schema(cursor) + # Building this raw-history index is explicit maintenance work, never + # part of web/capture startup. It also speeds subsequent resumptions. + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_packet_history_traceroute_id " + f"ON packet_history(id) WHERE {TRACEROUTE_PREDICATE}" + ) + if cursor.execute( + "SELECT 1 FROM traceroute_routes WHERE parser_version > ? LIMIT 1", + (PARSER_VERSION,), + ).fetchone(): + raise ValueError("Database uses a newer traceroute decoder; upgrade Malla") + upper_id = cursor.execute( + "SELECT COALESCE(MAX(id), 0) FROM packet_history" + ).fetchone()[0] + + last_id = None + processed = 0 + while True: + with conn: + conn.execute("BEGIN IMMEDIATE") + cursor = conn.cursor() + lower_bound = "" if last_id is None else "AND id > ?" + params = [upper_id] + if last_id is not None: + params.append(last_id) + rows = cursor.execute( + f""" + SELECT id, timestamp, mesh_packet_id, from_node_id, to_node_id, + hop_start, hop_limit, raw_payload + FROM packet_history + WHERE {TRACEROUTE_PREDICATE} AND id <= ? {lower_bound} + AND NOT EXISTS ( + SELECT 1 FROM traceroute_routes + WHERE packet_id = packet_history.id AND parser_version = ? + ) + ORDER BY id LIMIT ? + """, + [*params, PARSER_VERSION, batch_size], + ).fetchall() + if not rows: + break + for row in rows: + write_traceroute(cursor, dict(row)) + last_id = rows[-1]["id"] + processed += len(rows) + if progress is not None: + progress(processed) + + with conn: + conn.execute("BEGIN IMMEDIATE") + cursor = conn.cursor() + result = inspect_traceroutes(cursor) + return {**result, "processed_this_run": processed} + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--database", required=True, type=Path, help="Existing database path" + ) + parser.add_argument( + "--batch-size", type=int, default=1000, help="Packets per transaction" + ) + parser.add_argument( + "--check", + "--dry-run", + action="store_true", + help="Report counts without modifying the database", + ) + args = parser.parse_args(argv) + if args.batch_size < 1: + parser.error("--batch-size must be positive") + path = args.database.expanduser().resolve() + print(f"Database: {path}", flush=True) + try: + # mode=rw refuses to create a database for a misspelled path. Neither + # configuration defaults nor capture startup participate in this command. + mode = "ro" if args.check else "rw" + with closing( + sqlite3.connect(f"{path.as_uri()}?mode={mode}", uri=True, timeout=30.0) + ) as conn: + conn.row_factory = sqlite3.Row + conn.execute("PRAGMA foreign_keys=ON") + if args.check: + with conn: + conn.execute("BEGIN") + result = inspect_traceroutes(conn.cursor()) + else: + result = prepare_traceroutes( + conn, + batch_size=args.batch_size, + progress=lambda count: print( + f"Prepared {count} traceroutes", flush=True + ), + ) + print(json.dumps(result, indent=2), flush=True) + return 0 if result["complete"] else 1 + except KeyboardInterrupt: + print( + "Interrupted. Committed batches are safe; run the same command to resume.", + file=sys.stderr, + ) + return 130 + except (OSError, sqlite3.Error, ValueError) as exc: + print(f"Traceroute preparation failed: {exc}", file=sys.stderr) + return 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/src/malla/database/repositories.py b/src/malla/database/repositories.py index 36a72126..ef4ae5b0 100644 --- a/src/malla/database/repositories.py +++ b/src/malla/database/repositories.py @@ -4,7 +4,6 @@ This module provides data access layer with business logic for different entities. """ -import json import logging import time from datetime import UTC, datetime @@ -16,6 +15,7 @@ from ..config import get_config from ..utils.decryption import try_decrypt_mesh_packet from ..utils.formatting import format_time_ago +from ..utils.geo_utils import is_valid_position from ..utils.node_utils import convert_node_id, get_bulk_node_short_names from ..utils.signal_quality import ( is_plausible_rssi, @@ -3336,62 +3336,6 @@ def get_node_telemetry_history( class TracerouteRepository: """Repository for traceroute operations.""" - @staticmethod - def get_traceroute_packets_for_graph( - limit: int = 5000, - filters: dict[str, Any] | None = None, - ) -> list[dict[str, Any]]: - """Get minimal traceroute packet fields for network graph extraction.""" - if filters is None: - filters = {} - - try: - conn = get_db_connection() - cursor = conn.cursor() - - where_conditions = ["portnum_name = 'TRACEROUTE_APP'"] - params: list[Any] = [] - - if filters.get("start_time"): - where_conditions.append("timestamp >= ?") - params.append(filters["start_time"]) - - if filters.get("end_time"): - where_conditions.append("timestamp <= ?") - params.append(filters["end_time"]) - - if filters.get("gateway_id"): - where_conditions.append("gateway_id = ?") - params.append(filters["gateway_id"]) - - if filters.get("processed_successfully_only"): - where_conditions.append("processed_successfully = 1") - - where_clause = "WHERE " + " AND ".join(where_conditions) - query = f""" - SELECT - id, - timestamp, - from_node_id, - to_node_id, - gateway_id, - hop_start, - hop_limit, - raw_payload - FROM packet_history - {where_clause} - ORDER BY timestamp DESC - LIMIT ? - """ - - cursor.execute(query, [*params, limit]) - rows = [dict(row) for row in cursor.fetchall()] - conn.close() - return rows - except Exception as e: - logger.error(f"Error getting traceroute packets for graph: {e}") - raise - @staticmethod def get_traceroute_packets( limit: int = 100, @@ -3403,548 +3347,37 @@ def get_traceroute_packets( group_packets: bool = False, ) -> dict[str, Any]: """Get traceroute packets with filtering and optional grouping.""" - if filters is None: - filters = {} - - try: - conn = get_db_connection() - cursor = conn.cursor() - - # Build WHERE clause - where_conditions = ["portnum_name = 'TRACEROUTE_APP'"] - params = [] - - if filters.get("start_time"): - where_conditions.append("timestamp >= ?") - params.append(filters["start_time"]) - - if filters.get("end_time"): - where_conditions.append("timestamp <= ?") - params.append(filters["end_time"]) - - if filters.get("from_node"): - where_conditions.append("from_node_id = ?") - params.append(filters["from_node"]) - - if filters.get("to_node"): - where_conditions.append("to_node_id = ?") - params.append(filters["to_node"]) - - if filters.get("gateway_id"): - where_conditions.append("gateway_id = ?") - params.append(filters["gateway_id"]) - - # New: Optional filtering by primary_channel (matches packet.channel_id field) - if filters.get("primary_channel"): - where_conditions.append("channel_id = ?") - params.append(filters["primary_channel"]) - - if filters.get("processed_successfully_only"): - where_conditions.append("processed_successfully = 1") - - # Check if route_node filtering is needed - route_node_filter = filters.get("route_node") - needs_route_filtering = route_node_filter is not None - - # Add search functionality - if search: - search_conditions = [ - "gateway_id LIKE ?", - "CAST(from_node_id AS TEXT) LIKE ?", - "CAST(to_node_id AS TEXT) LIKE ?", - ] - search_param = f"%{search}%" - where_conditions.append(f"({' OR '.join(search_conditions)})") - params.extend([search_param] * len(search_conditions)) - - where_clause = "WHERE " + " AND ".join(where_conditions) - - if group_packets: - # Determine time window (default: 7 days for traceroutes) - time_window_days = 7 - - # If no specific time filters, use default window - if not filters.get("start_time") and not filters.get("end_time"): - import time - - current_time = time.time() - window_start = current_time - (time_window_days * 24 * 3600) - where_conditions.append("timestamp >= ?") - params.append(window_start) - - # Add mesh_packet_id filter and exclude special cases - where_conditions.append("mesh_packet_id IS NOT NULL") - where_conditions.append("mesh_packet_id != 0") # Exclude problematic ID - - where_clause = "WHERE " + " AND ".join(where_conditions) - - # PERFORMANCE FIX: Skip expensive total count for grouped traceroute queries - # The COUNT(DISTINCT ...) query was taking too long on large datasets - # Instead, estimate total count based on results (much faster) - total_count = None # Will be estimated after getting results - - # ULTRA-OPTIMIZED: Use much smaller fetch limits for better performance - # The original approach was fetching 1k-100k records which is too expensive - # Instead, use a more reasonable approach with smaller multipliers - if offset == 0: - # For first page, use a smaller multiplier for traceroutes - fetch_limit = min( - limit * 15, 3000 - ) # Smaller: 375-3000 instead of 1k-40k - else: - # For subsequent pages, use a reasonable multiplier - grouping_ratio = 2.0 # More realistic estimate - estimated_individual_needed = (offset + limit) * grouping_ratio - - # Cap at much smaller limits for performance - fetch_limit = min( - max(estimated_individual_needed, limit * 8), 8000 - ) # Max 8k instead of 100k - - # Fetch individual packets using efficient ORDER BY timestamp DESC LIMIT - query = f""" - SELECT - id, timestamp, from_node_id, to_node_id, gateway_id, - channel_id, hop_start, hop_limit, rssi, snr, payload_length, raw_payload, - processed_successfully, mesh_packet_id, - datetime(timestamp, 'unixepoch') as timestamp_str - FROM packet_history - {where_clause} - ORDER BY timestamp DESC - LIMIT ? - """ - - cursor.execute(query, params + [fetch_limit]) - rows = cursor.fetchall() - individual_packets: list[dict[str, Any]] = [dict(row) for row in rows] - - # Group packets in memory by (mesh_packet_id, from_node_id, to_node_id) - groups: dict[tuple[Any, Any, Any], list[dict[str, Any]]] = {} - for packet in individual_packets: - # Skip if missing required fields - if not packet.get("mesh_packet_id") or not packet.get( - "from_node_id" - ): - continue - - # Create grouping key - group_key = ( - packet["mesh_packet_id"], - packet["from_node_id"], - packet["to_node_id"], - ) - - if group_key not in groups: - groups[group_key] = [] - groups[group_key].append(packet) - - # Convert groups to aggregated packets - aggregated_packets = [] - for _group_key, packets_in_group in groups.items(): - # Sort by timestamp (newest first) within group - packets_in_group.sort(key=lambda x: x["timestamp"], reverse=True) - - # Use the first (newest) packet as the base - base_packet = packets_in_group[0] - - # Calculate aggregations - gateway_ids = [ - p["gateway_id"] for p in packets_in_group if p["gateway_id"] - ] - unique_gateways = list(set(gateway_ids)) - - rssi_values = [ - p["rssi"] - for p in packets_in_group - if is_plausible_rssi(p["rssi"]) - ] - snr_values = [ - p["snr"] for p in packets_in_group if is_plausible_snr(p["snr"]) - ] - hop_values = [] - for p in packets_in_group: - if ( - p.get("hop_start") is not None - and p.get("hop_limit") is not None - ): - hop_values.append(p["hop_start"] - p["hop_limit"]) - - payload_lengths = [ - p["payload_length"] - for p in packets_in_group - if p["payload_length"] - ] - - # Find the packet with the longest payload (most complete route data) - best_payload_packet = max( - packets_in_group, key=lambda x: len(x.get("raw_payload", b"")) - ) - - # Create aggregated packet - aggregated = { - "id": base_packet["id"], - "timestamp": base_packet["timestamp"], - "timestamp_str": base_packet["timestamp_str"], - "from_node_id": base_packet["from_node_id"], - "to_node_id": base_packet["to_node_id"], - "channel_id": base_packet.get("channel_id"), - "mesh_packet_id": base_packet["mesh_packet_id"], - "gateway_count": len(unique_gateways), - "gateway_list": ",".join(unique_gateways), - "reception_count": len(packets_in_group), - "processed_successfully": any( - p["processed_successfully"] for p in packets_in_group - ), - "raw_payload": best_payload_packet.get("raw_payload"), - "is_grouped": True, - } - - # RSSI aggregation - if rssi_values: - aggregated["min_rssi"] = min(rssi_values) - aggregated["max_rssi"] = max(rssi_values) - if aggregated["min_rssi"] == aggregated["max_rssi"]: - aggregated["rssi_range"] = ( - f"{aggregated['min_rssi']:.1f} dBm" - ) - else: - aggregated["rssi_range"] = ( - f"{aggregated['min_rssi']:.1f} to {aggregated['max_rssi']:.1f} dBm" - ) - aggregated["rssi"] = aggregated["rssi_range"] - else: - aggregated["min_rssi"] = None - aggregated["max_rssi"] = None - aggregated["rssi_range"] = None - aggregated["rssi"] = None - - # SNR aggregation - if snr_values: - aggregated["min_snr"] = min(snr_values) - aggregated["max_snr"] = max(snr_values) - if aggregated["min_snr"] == aggregated["max_snr"]: - aggregated["snr_range"] = f"{aggregated['min_snr']:.2f} dB" - else: - aggregated["snr_range"] = ( - f"{aggregated['min_snr']:.2f} to {aggregated['max_snr']:.2f} dB" - ) - aggregated["snr"] = aggregated["snr_range"] - else: - aggregated["min_snr"] = None - aggregated["max_snr"] = None - aggregated["snr_range"] = None - aggregated["snr"] = None - - # Hop count aggregation - if hop_values: - aggregated["min_hops"] = min(hop_values) - aggregated["max_hops"] = max(hop_values) - if aggregated["min_hops"] == aggregated["max_hops"]: - aggregated["hop_range"] = str(aggregated["min_hops"]) - else: - aggregated["hop_range"] = ( - f"{aggregated['min_hops']}-{aggregated['max_hops']}" - ) - aggregated["hop_count"] = aggregated["min_hops"] - else: - aggregated["min_hops"] = None - aggregated["max_hops"] = None - aggregated["hop_range"] = None - aggregated["hop_count"] = None - - # Payload length aggregation - if payload_lengths: - aggregated["avg_payload_length"] = sum(payload_lengths) / len( - payload_lengths - ) - else: - aggregated["avg_payload_length"] = None - - # Success indicator - aggregated["success"] = aggregated["processed_successfully"] - - # Enhanced route display using TraceroutePacket - aggregated["route"] = None - aggregated["route_display"] = "No route data" - if aggregated.get("raw_payload"): - try: - from ..models.traceroute import TraceroutePacket - - tr_packet = TraceroutePacket(aggregated, resolve_names=True) - if tr_packet.route_data["route_nodes"]: - aggregated["route"] = json.dumps( - tr_packet.route_data["route_nodes"] - ) - # Get enhanced route display with node names - aggregated["route_display"] = ( - tr_packet.format_path_display("display") - ) - except Exception as e: - logger.debug( - f"Failed to parse route for grouped packet {aggregated['id']}: {e}" - ) - - aggregated_packets.append(aggregated) - - if needs_route_filtering: - filtered_packets: list[dict[str, Any]] = [] - for packet in aggregated_packets: - # Direct match on source/destination - if ( - packet.get("from_node_id") == route_node_filter - or packet.get("to_node_id") == route_node_filter - ): - filtered_packets.append(packet) - continue - - # Attempt to match within the hop route - # Prefer already extracted route information if available - route_nodes: list[int] | None = None - if packet.get("route"): - try: - import json as _json - - route_nodes = _json.loads(packet["route"]) - except Exception: - route_nodes = None - - # If not available, fall back to parsing the raw payload - if route_nodes is None and packet.get("raw_payload"): - try: - from ..models.traceroute import TraceroutePacket as _TRP - - tr_packet = _TRP(packet, resolve_names=False) - route_nodes = tr_packet.route_data.get( - "route_nodes", [] - ) - except Exception as e: - logger.debug( - f"Failed to parse route for grouped route_node filtering: {e}" - ) - if route_nodes and route_node_filter in route_nodes: - filtered_packets.append(packet) - - aggregated_packets = filtered_packets - # For grouped queries we can now set an accurate total_count - total_count = len(aggregated_packets) - - # Apply sorting to aggregated packets - reverse_sort = order_dir.lower() == "desc" - - if order_by == "gateway_id" or order_by == "gateway_count": - # Sort by gateway count - aggregated_packets.sort( - key=lambda x: x["gateway_count"], reverse=reverse_sort - ) - elif order_by == "timestamp": - aggregated_packets.sort( - key=lambda x: x["timestamp"], reverse=reverse_sort - ) - elif order_by == "from_node_id": - aggregated_packets.sort( - key=lambda x: x.get("from_node_id", 0), reverse=reverse_sort - ) - elif order_by == "to_node_id": - aggregated_packets.sort( - key=lambda x: x.get("to_node_id", 0), reverse=reverse_sort - ) - elif order_by == "rssi": - aggregated_packets.sort( - key=lambda x: x.get("min_rssi", -999), reverse=reverse_sort - ) - elif order_by == "snr": - aggregated_packets.sort( - key=lambda x: x.get("min_snr", -999), reverse=reverse_sort - ) - elif order_by == "hop_count": - aggregated_packets.sort( - key=lambda x: x.get("min_hops", 999), reverse=reverse_sort - ) - elif order_by == "payload_length": - aggregated_packets.sort( - key=lambda x: x.get("avg_payload_length", 0), - reverse=reverse_sort, - ) - else: - # Default to timestamp - aggregated_packets.sort( - key=lambda x: x["timestamp"], reverse=reverse_sort - ) - - # Apply pagination - packets = aggregated_packets[offset : offset + limit] - - # Handle None total_count for grouped queries - if total_count is None: - # Estimate total_count based on results for grouped queries - if len(packets) == limit: - total_count = ( - offset + limit + 1 - ) # Estimate at least one more page - else: - total_count = offset + len( - packets - ) # Exact count for partial page - - else: - # Original ungrouped behavior - - # If route_node filtering is needed, we need to fetch more data - # to account for filtering before pagination - if needs_route_filtering: - # Fetch a larger dataset to ensure we have enough results after filtering - # Use a multiplier based on how selective route_node filtering typically is - fetch_multiplier = 20 # Empirically determined - adjust as needed - fetch_limit = max((offset + limit) * fetch_multiplier, 1000) - fetch_offset = 0 # Start from beginning when route filtering - else: - fetch_limit = limit - fetch_offset = offset - - # Get total count (before route filtering) - cursor.execute( - f"SELECT COUNT(*) as total FROM packet_history {where_clause}", - params, - ) - total_count_before_filter = cursor.fetchone()["total"] - - # Main query - valid_order_columns = [ - "timestamp", - "from_node_id", - "to_node_id", - "gateway_id", - "rssi", - "snr", - "payload_length", - "hop_count", # Allow ordering by computed hops - ] - if order_by not in valid_order_columns: - order_by = "timestamp" - - order_dir_sql = "DESC" if order_dir.lower() == "desc" else "ASC" - - query = f""" - SELECT - id, timestamp, from_node_id, to_node_id, gateway_id, - channel_id, hop_start, hop_limit, rssi, snr, payload_length, raw_payload, - processed_successfully, mesh_packet_id, - datetime(timestamp, 'unixepoch') as timestamp_str, - (hop_start - hop_limit) AS hop_count - FROM packet_history - {where_clause} - ORDER BY {order_by} {order_dir_sql} - LIMIT ? OFFSET ? - """ - - cursor.execute(query, params + [fetch_limit, fetch_offset]) - all_packets = [] - for row in cursor.fetchall(): - packet = dict(row) - - # Format timestamp if not already formatted - if packet["timestamp_str"] is None: - packet["timestamp_str"] = datetime.fromtimestamp( - packet["timestamp"], UTC - ).strftime("%Y-%m-%d %H:%M:%S UTC") - - # Add success indicator - packet["success"] = packet["processed_successfully"] - packet["is_grouped"] = False - - # Extract route data from raw_payload if available - packet["route"] = None - if packet.get("raw_payload"): - try: - from ..models.traceroute import TraceroutePacket - - tr_packet = TraceroutePacket(packet, resolve_names=False) - if tr_packet.route_data["route_nodes"]: - packet["route"] = json.dumps( - tr_packet.route_data["route_nodes"] - ) - except Exception as e: - logger.debug( - f"Failed to parse route for packet {packet['id']}: {e}" - ) - - # Calculate hop count from hop_start and hop_limit - if ( - packet.get("hop_start") is not None - and packet.get("hop_limit") is not None - ): - packet["hop_count"] = packet["hop_start"] - packet["hop_limit"] - else: - packet["hop_count"] = None - - all_packets.append(packet) - - # Apply route_node filtering if specified - if needs_route_filtering: - filtered_packets = [] - for packet in all_packets: - # Check if the route_node appears in from_node_id, to_node_id, or route_nodes - if ( - packet.get("from_node_id") == route_node_filter - or packet.get("to_node_id") == route_node_filter - ): - filtered_packets.append(packet) - continue - - # Check if the node appears in the route_nodes array - if packet.get("raw_payload"): - try: - from ..models.traceroute import TraceroutePacket - - tr_packet = TraceroutePacket( - packet, resolve_names=False - ) - if route_node_filter in tr_packet.route_data.get( - "route_nodes", [] - ): - filtered_packets.append(packet) - except Exception as e: - logger.debug( - f"Failed to parse route for route_node filtering: {e}" - ) - - # Now apply pagination to filtered results - total_count = len(filtered_packets) - packets = filtered_packets[offset : offset + limit] - else: - # No route filtering needed, use all packets - packets = all_packets - total_count = total_count_before_filter - - conn.close() - - return { - "packets": packets, - "total_count": total_count, - "limit": limit, - "offset": offset, - "is_grouped": group_packets, - } - - except Exception as e: - logger.error(f"Error getting traceroute packets: {e}") - raise + from .traceroute_read_repository import get_traceroute_packets + + return get_traceroute_packets( + limit=limit, + offset=offset, + filters=filters, + order_by=order_by, + order_dir=order_dir, + search=search, + group_packets=group_packets, + ) @staticmethod def get_traceroute_details(packet_id: int) -> dict[str, Any] | None: - """Get details for a specific traceroute packet.""" + """Get details for a specific traceroute packet using materialized routes.""" try: conn = get_db_connection() cursor = conn.cursor() query = """ SELECT - id, timestamp, from_node_id, to_node_id, gateway_id, - hop_start, hop_limit, rssi, snr, payload_length, raw_payload, - processed_successfully, - datetime(timestamp, 'unixepoch') as timestamp_str - FROM packet_history - WHERE id = ? AND portnum_name = 'TRACEROUTE_APP' + p.id, p.timestamp, p.from_node_id, p.to_node_id, p.gateway_id, + p.channel_id, p.hop_start, p.hop_limit, p.rssi, p.snr, + p.payload_length, p.raw_payload, p.processed_successfully, + r.route_nodes_json, r.snr_towards_json, r.route_back_json, r.snr_back_json, + r.forward_complete, r.return_complete, r.parse_status, r.parse_error, + r.parser_version, + datetime(p.timestamp, 'unixepoch') as timestamp_str + FROM packet_history p + LEFT JOIN traceroute_routes r ON r.packet_id = p.id + WHERE p.id = ? AND (p.portnum = 70 OR p.portnum_name = 'TRACEROUTE_APP') """ cursor.execute(query, (packet_id,)) @@ -3962,6 +3395,13 @@ def get_traceroute_details(packet_id: int) -> dict[str, Any] | None: class LocationRepository: """Repository for location operations.""" + # How many of each node's most recent position packets to consider when + # resolving its latest valid location. Firmware bugs intermittently emit + # garbage (~0/~0) coordinates; the newest valid packet wins so nodes fall + # back to their previous good fix instead of vanishing or jumping to null + # island. + POSITION_LOOKUP_DEPTH = 5 + @staticmethod def get_node_locations( filters: dict[str, Any] | None = None, @@ -4023,41 +3463,47 @@ def get_node_locations( if extra_conditions: extra_where = "AND " + " AND ".join(extra_conditions) - # Optimized query using window function instead of correlated subquery + # Optimized query using window function instead of correlated subquery. + # Ranks each node's recent position packets so the decode loop can + # skip invalid ones (null-island garbage) and fall back to the + # previous valid fix. query = f""" - WITH max_timestamps AS ( + WITH ranked AS ( SELECT - from_node_id, - MAX(timestamp) as max_timestamp - FROM packet_history - WHERE portnum = 3 -- POSITION_APP - AND raw_payload IS NOT NULL - AND from_node_id IS NOT NULL + ph.from_node_id, + ph.timestamp, + ph.raw_payload, + ROW_NUMBER() OVER ( + PARTITION BY ph.from_node_id + ORDER BY ph.timestamp DESC + ) AS rank + FROM packet_history ph + WHERE ph.portnum = 3 -- POSITION_APP + AND ph.raw_payload IS NOT NULL + AND ph.from_node_id IS NOT NULL {node_ids_clause} {extra_where} - GROUP BY from_node_id ) SELECT - ph.from_node_id as node_id, - ph.timestamp, - ph.raw_payload, + r.from_node_id as node_id, + r.timestamp, + r.raw_payload, ni.long_name, ni.short_name, ni.hw_model, ni.role, ni.primary_channel, - printf('!%08x', ph.from_node_id) as hex_id - FROM packet_history ph - INNER JOIN max_timestamps mt ON ph.from_node_id = mt.from_node_id - AND ph.timestamp = mt.max_timestamp - LEFT JOIN node_info ni ON ph.from_node_id = ni.node_id - WHERE ph.portnum = 3 - AND ph.raw_payload IS NOT NULL - ORDER BY ph.timestamp DESC + printf('!%08x', r.from_node_id) as hex_id + FROM ranked r + LEFT JOIN node_info ni ON r.from_node_id = ni.node_id + WHERE r.rank <= ? + ORDER BY r.timestamp DESC """ query_start = time.time() - cursor.execute(query, [*node_ids_params, *extra_params]) + cursor.execute( + query, [*node_ids_params, *extra_params, LocationRepository.POSITION_LOOKUP_DEPTH] + ) raw_rows = cursor.fetchall() timing_breakdown["sql_query"] = time.time() - query_start @@ -4066,9 +3512,14 @@ def get_node_locations( locations = [] decode_count = 0 skip_count = 0 + seen_nodes: set[int] = set() for row in raw_rows: try: + # A newer row for this node was already accepted + if row["node_id"] in seen_nodes: + continue + if not row["raw_payload"]: skip_count += 1 continue @@ -4146,12 +3597,7 @@ def get_node_locations( precision_meters = math.exp(log_result) - if ( - latitude is None - or longitude is None - or latitude == 0 - or longitude == 0 - ): + if not is_valid_position(latitude, longitude): skip_count += 1 continue @@ -4161,6 +3607,8 @@ def get_node_locations( or f"Node {row['node_id']:08x}" ) + seen_nodes.add(row["node_id"]) + locations.append( { "node_id": row["node_id"], @@ -4272,8 +3720,8 @@ def get_node_location_history( ) altitude = position.altitude if position.altitude else None - # Skip invalid coordinates - if not latitude or not longitude: + # Skip invalid coordinates (unset, null-island garbage, out of range) + if not is_valid_position(latitude, longitude): continue locations.append( @@ -4298,6 +3746,86 @@ def get_node_location_history( logger.error(f"Error getting node location history: {e}") raise + @staticmethod + def get_nodes_location_history( + node_ids: list[int] | set[int], limit_per_node: int = 50 + ) -> dict[int, list[dict[str, Any]]]: + """Batch-fetch location history for multiple nodes from position packets.""" + if not node_ids: + return {} + clean_node_ids = [int(nid) for nid in node_ids if nid is not None and nid != 4294967295] + if not clean_node_ids: + return {} + + results: dict[int, list[dict[str, Any]]] = {nid: [] for nid in clean_node_ids} + try: + conn = get_db_connection() + cursor = conn.cursor() + + chunk_size = 500 + for i in range(0, len(clean_node_ids), chunk_size): + chunk = clean_node_ids[i : i + chunk_size] + placeholders = ",".join("?" * len(chunk)) + query = f""" + WITH ranked AS ( + SELECT + from_node_id, + timestamp, + raw_payload, + datetime(timestamp, 'unixepoch') AS timestamp_str, + ROW_NUMBER() OVER ( + PARTITION BY from_node_id ORDER BY timestamp DESC + ) AS rank + FROM packet_history + WHERE from_node_id IN ({placeholders}) + AND portnum = 3 + AND raw_payload IS NOT NULL + ) + SELECT from_node_id, timestamp, raw_payload, timestamp_str + FROM ranked + WHERE rank <= ? + ORDER BY from_node_id, timestamp DESC + """ + cursor.execute(query, [*chunk, limit_per_node]) + for row in cursor.fetchall(): + try: + raw_payload = row["raw_payload"] + if not raw_payload: + continue + position = mesh_pb2.Position() + position.ParseFromString(raw_payload) + latitude = ( + position.latitude_i / 1e7 if position.latitude_i else None + ) + longitude = ( + position.longitude_i / 1e7 if position.longitude_i else None + ) + altitude = position.altitude if position.altitude else None + if not is_valid_position(latitude, longitude): + continue + + results[row["from_node_id"]].append( + { + "latitude": latitude, + "longitude": longitude, + "altitude": altitude, + "timestamp": row["timestamp"], + "timestamp_str": row["timestamp_str"], + } + ) + except Exception as e: + logger.warning( + f"Failed to parse location for node {row['from_node_id']} timestamp {row['timestamp']}: {e}" + ) + continue + + conn.close() + return results + + except Exception as e: + logger.error(f"Error getting batch node location history: {e}") + raise + @staticmethod def get_latest_node_location(node_id: int) -> dict[str, Any] | None: """Return the most recent decoded location packet for a single node. @@ -4321,7 +3849,9 @@ def get_latest_node_location(node_id: int) -> dict[str, Any] | None: else: node_id = int(node_id) - # Fetch the most recent POSITION_APP packet for this node + # Fetch the node's most recent POSITION_APP packets so a garbage + # newest fix (null-island coordinates) can fall back to the last + # valid one cursor.execute( """ SELECT timestamp, raw_payload @@ -4330,43 +3860,45 @@ def get_latest_node_location(node_id: int) -> dict[str, Any] | None: AND portnum = 3 -- POSITION_APP AND raw_payload IS NOT NULL ORDER BY timestamp DESC - LIMIT 1 + LIMIT ? """, - (node_id,), + (node_id, LocationRepository.POSITION_LOOKUP_DEPTH), ) - row = cursor.fetchone() + rows = cursor.fetchall() - if not row: + if not rows: conn.close() return None - # Decode protobuf – this is the same logic used elsewhere but for a single row - try: - position = mesh_pb2.Position() - position.ParseFromString(row["raw_payload"]) + # Decode protobuf – this is the same logic used elsewhere but for + # a handful of rows, returning the newest valid position + for row in rows: + try: + position = mesh_pb2.Position() + position.ParseFromString(row["raw_payload"]) - latitude = position.latitude_i / 1e7 if position.latitude_i else None - longitude = position.longitude_i / 1e7 if position.longitude_i else None - altitude = position.altitude if position.altitude else None + latitude = position.latitude_i / 1e7 if position.latitude_i else None + longitude = position.longitude_i / 1e7 if position.longitude_i else None + altitude = position.altitude if position.altitude else None - if not latitude or not longitude or latitude == 0 or longitude == 0: - conn.close() - return None + if not is_valid_position(latitude, longitude): + continue - result = { - "latitude": latitude, - "longitude": longitude, - "altitude": altitude, - "timestamp": row["timestamp"], - } - conn.close() - return result - except Exception as e: - logger.warning( - f"Failed to decode position payload for node {node_id}: {e}" - ) - conn.close() - return None + result = { + "latitude": latitude, + "longitude": longitude, + "altitude": altitude, + "timestamp": row["timestamp"], + } + conn.close() + return result + except Exception as e: + logger.warning( + f"Failed to decode position payload for node {node_id}: {e}" + ) + continue + conn.close() + return None except Exception as e: logger.error(f"Error getting latest location for node {node_id}: {e}") raise @@ -4380,58 +3912,75 @@ def get_node_location_at_timestamp( conn = get_db_connection() cursor = conn.cursor() - # First try to get the most recent location before or at the target timestamp - query_before = """ - SELECT timestamp, raw_payload - FROM packet_history - WHERE from_node_id = ? - AND portnum = 3 -- POSITION_APP - AND timestamp <= ? - AND raw_payload IS NOT NULL - ORDER BY timestamp DESC - LIMIT 1 - """ - - cursor.execute(query_before, (node_id, target_timestamp)) - location_before = cursor.fetchone() + def _first_valid_location( + rows: list[Any], *, later: bool + ) -> dict[str, Any] | None: + for location_row in rows: + try: + position = mesh_pb2.Position() + position.ParseFromString(location_row["raw_payload"]) - if location_before: - try: - # Decode position from raw protobuf payload - position = mesh_pb2.Position() - position.ParseFromString(location_before["raw_payload"]) + latitude = ( + position.latitude_i / 1e7 if position.latitude_i else None + ) + longitude = ( + position.longitude_i / 1e7 + if position.longitude_i + else None + ) + altitude = position.altitude if position.altitude else None - # Extract coordinates (stored as integers, need to divide by 1e7) - latitude = ( - position.latitude_i / 1e7 if position.latitude_i else None - ) - longitude = ( - position.longitude_i / 1e7 if position.longitude_i else None - ) - altitude = position.altitude if position.altitude else None + if not is_valid_position(latitude, longitude): + continue - if latitude and longitude: - age_seconds = target_timestamp - location_before["timestamp"] + if later: + age_seconds = location_row["timestamp"] - target_timestamp + else: + age_seconds = target_timestamp - location_row["timestamp"] age_hours = age_seconds / 3600 + unit = "later" if later else "ago" if age_hours <= 24: - age_warning = f"from {age_hours:.1f}h ago" + age_warning = f"from {age_hours:.1f}h {unit}" elif age_hours <= 168: # 1 week - age_warning = f"from {age_hours / 24:.1f}d ago" + age_warning = f"from {age_hours / 24:.1f}d {unit}" else: - age_warning = f"from {age_hours / 168:.1f}w ago" + age_warning = f"from {age_hours / 168:.1f}w {unit}" return { "latitude": latitude, "longitude": longitude, "altitude": altitude, - "timestamp": location_before["timestamp"], + "timestamp": location_row["timestamp"], "age_warning": age_warning, } - except Exception as e: - logger.warning(f"Failed to decode position from raw payload: {e}") + except Exception as e: + logger.warning(f"Failed to decode position from raw payload: {e}") + return None + + # Try to get the most recent valid location before or at the target + # timestamp, falling back past invalid (null-island) fixes + query_before = """ + SELECT timestamp, raw_payload + FROM packet_history + WHERE from_node_id = ? + AND portnum = 3 -- POSITION_APP + AND timestamp <= ? + AND raw_payload IS NOT NULL + ORDER BY timestamp DESC + LIMIT ? + """ + + cursor.execute( + query_before, + (node_id, target_timestamp, LocationRepository.POSITION_LOOKUP_DEPTH), + ) + location_before = _first_valid_location(cursor.fetchall(), later=False) + if location_before: + conn.close() + return location_before - # If no location before target, try to get the earliest location after + # If no valid location before target, try the earliest location after query_after = """ SELECT timestamp, raw_payload FROM packet_history @@ -4440,47 +3989,17 @@ def get_node_location_at_timestamp( AND timestamp > ? AND raw_payload IS NOT NULL ORDER BY timestamp ASC - LIMIT 1 + LIMIT ? """ - cursor.execute(query_after, (node_id, target_timestamp)) - location_after = cursor.fetchone() - + cursor.execute( + query_after, + (node_id, target_timestamp, LocationRepository.POSITION_LOOKUP_DEPTH), + ) + location_after = _first_valid_location(cursor.fetchall(), later=True) if location_after: - try: - # Decode position from raw protobuf payload - position = mesh_pb2.Position() - position.ParseFromString(location_after["raw_payload"]) - - # Extract coordinates (stored as integers, need to divide by 1e7) - latitude = ( - position.latitude_i / 1e7 if position.latitude_i else None - ) - longitude = ( - position.longitude_i / 1e7 if position.longitude_i else None - ) - altitude = position.altitude if position.altitude else None - - if latitude and longitude: - age_seconds = location_after["timestamp"] - target_timestamp - age_hours = age_seconds / 3600 - - if age_hours <= 24: - age_warning = f"from {age_hours:.1f}h later" - elif age_hours <= 168: # 1 week - age_warning = f"from {age_hours / 24:.1f}d later" - else: - age_warning = f"from {age_hours / 168:.1f}w later" - - return { - "latitude": latitude, - "longitude": longitude, - "altitude": altitude, - "timestamp": location_after["timestamp"], - "age_warning": age_warning, - } - except Exception as e: - logger.warning(f"Failed to decode position from raw payload: {e}") + conn.close() + return location_after conn.close() return None diff --git a/src/malla/database/schema.py b/src/malla/database/schema.py index c84fbfcd..618276c5 100644 --- a/src/malla/database/schema.py +++ b/src/malla/database/schema.py @@ -3,6 +3,8 @@ import logging import sqlite3 +from .traceroute_schema import ensure_traceroute_schema + logger = logging.getLogger(__name__) @@ -208,6 +210,9 @@ def ensure_startup_schema( cursor.execute(ACTIVITY_ROLLUP_TABLE_SQL) + if "packet_history" in existing_tables: + ensure_traceroute_schema(cursor) + if "node_info" in existing_tables: cursor.execute("PRAGMA table_info(node_info)") node_info_columns = {row[1] for row in cursor.fetchall()} diff --git a/src/malla/database/traceroute_read_repository.py b/src/malla/database/traceroute_read_repository.py new file mode 100644 index 00000000..5e36b2a0 --- /dev/null +++ b/src/malla/database/traceroute_read_repository.py @@ -0,0 +1,614 @@ +"""Queries for materialized traceroutes used by web readers.""" + +from __future__ import annotations + +import json +import time +from typing import Any + +from ..utils.signal_quality import ( + SNR_PLAUSIBLE_MAX, + SNR_PLAUSIBLE_MIN, + TRACEROUTE_UNKNOWN_SNR, + rssi_valid_sql, + snr_valid_sql, +) +from .connection import get_db_connection +from .traceroutes import PARSER_VERSION + + +def _where_clause( + filters: dict[str, Any], + search: str | None, + *, + include_success_filter: bool = True, +) -> tuple[str, list[Any]]: + conditions = ["r.parser_version = ?", "r.parse_status != 'pending'"] + params: list[Any] = [PARSER_VERSION] + field_filters = ( + ("start_time", "r.timestamp >= ?"), + ("end_time", "r.timestamp <= ?"), + ("from_node", "r.from_node_id = ?"), + ("to_node", "r.to_node_id = ?"), + ("gateway_id", "p.gateway_id = ?"), + ("primary_channel", "p.channel_id = ?"), + ) + for key, sql in field_filters: + value = filters.get(key) + if value is not None and value != "": + conditions.append(sql) + params.append(value) + if include_success_filter and ( + filters.get("processed_successfully_only") or filters.get("success_only") + ): + conditions.append("r.parse_status IN ('parsed', 'valid_empty')") + if filters.get("return_path_only"): + conditions.append("r.route_back_json != '[]'") + route_node = filters.get("route_node") + if route_node is not None and route_node != "": + try: + route_node = int(route_node) + except (ValueError, TypeError): + try: + if isinstance(route_node, str) and route_node.startswith(("0x", "0X")): + route_node = int(route_node, 16) + except (ValueError, TypeError): + pass + conditions.append( + "(r.from_node_id = ? OR r.to_node_id = ? " + "OR EXISTS (SELECT 1 FROM traceroute_hops h WHERE h.packet_id = r.packet_id " + "AND (h.from_node_id = ? OR h.to_node_id = ?)) " + "OR EXISTS (SELECT 1 FROM json_each(r.route_nodes_json) WHERE value = ?) " + "OR EXISTS (SELECT 1 FROM json_each(r.route_back_json) WHERE value = ?))" + ) + params.extend([route_node] * 6) + if search: + conditions.append( + "(p.gateway_id LIKE ? OR CAST(r.from_node_id AS TEXT) LIKE ? " + "OR CAST(r.to_node_id AS TEXT) LIKE ?)" + ) + params.extend([f"%{search}%"] * 3) + return " AND ".join(conditions), params + + +_SORT_EXPRESSIONS = { + "timestamp": "timestamp", + "from_node_id": "from_node_id", + "to_node_id": "to_node_id", + "gateway_id": "gateway_count", + "gateway_count": "gateway_count", + "rssi": "min_rssi", + "snr": "min_snr", + "hop_count": "min_hops", + "payload_length": "avg_payload_length", +} + + +def get_traceroute_packets( + *, + limit: int = 100, + offset: int = 0, + filters: dict[str, Any] | None = None, + order_by: str = "timestamp", + order_dir: str = "desc", + search: str | None = None, + group_packets: bool = False, +) -> dict[str, Any]: + """Return exact, post-filter pagination from saved route rows.""" + filters = dict(filters or {}) + if group_packets and "start_time" not in filters and "end_time" not in filters: + filters["start_time"] = time.time() - 7 * 24 * 3600 + conn = get_db_connection() + try: + cursor = conn.cursor() + where, params = _where_clause(filters, search) + direction = "ASC" if order_dir.lower() == "asc" else "DESC" + sort_expression = _SORT_EXPRESSIONS.get(order_by, "timestamp") + + if group_packets: + group_conditions = ["r.mesh_packet_id IS NOT NULL", "r.mesh_packet_id != 0"] + grouped_where = f"{where} AND {' AND '.join(group_conditions)}" + group_key = "mesh_packet_id, from_node_id, to_node_id" + total_count = cursor.execute( + f""" + SELECT COUNT(*) FROM ( + SELECT 1 + FROM traceroute_routes r + JOIN packet_history p ON p.id = r.packet_id + WHERE {grouped_where} + GROUP BY r.mesh_packet_id, r.from_node_id, r.to_node_id + ) + """, + params, + ).fetchone()[0] + rows = cursor.execute( + f""" + WITH filtered AS ( + SELECT + p.id, r.timestamp, r.from_node_id, r.to_node_id, + p.gateway_id, p.channel_id, p.hop_start, p.hop_limit, + p.rssi, p.snr, p.payload_length, p.processed_successfully, + r.mesh_packet_id, r.route_nodes_json, r.snr_towards_json, + r.route_back_json, r.snr_back_json, r.parse_status, + ROW_NUMBER() OVER ( + PARTITION BY r.mesh_packet_id, r.from_node_id, r.to_node_id + ORDER BY p.payload_length DESC, r.timestamp DESC, p.id DESC + ) AS representative_rank + FROM traceroute_routes r + JOIN packet_history p ON p.id = r.packet_id + WHERE {grouped_where} + ), grouped AS ( + SELECT + MAX(CASE WHEN representative_rank = 1 THEN id END) AS id, + MAX(timestamp) AS timestamp, + from_node_id, to_node_id, mesh_packet_id, + COUNT(DISTINCT gateway_id) AS gateway_count, + GROUP_CONCAT(DISTINCT gateway_id) AS gateway_list, + COUNT(*) AS reception_count, + MAX(processed_successfully) AS processed_successfully, + MAX(CASE WHEN representative_rank = 1 THEN channel_id END) AS channel_id, + MIN(CASE WHEN {rssi_valid_sql()} THEN rssi END) AS min_rssi, + MAX(CASE WHEN {rssi_valid_sql()} THEN rssi END) AS max_rssi, + MIN(CASE WHEN {snr_valid_sql()} THEN snr END) AS min_snr, + MAX(CASE WHEN {snr_valid_sql()} THEN snr END) AS max_snr, + MIN(hop_start - hop_limit) AS min_hops, + MAX(hop_start - hop_limit) AS max_hops, + AVG(payload_length) AS avg_payload_length, + MAX(CASE WHEN representative_rank = 1 THEN route_nodes_json END) AS route_nodes_json, + MAX(CASE WHEN representative_rank = 1 THEN snr_towards_json END) AS snr_towards_json, + MAX(CASE WHEN representative_rank = 1 THEN route_back_json END) AS route_back_json, + MAX(CASE WHEN representative_rank = 1 THEN snr_back_json END) AS snr_back_json, + MAX(CASE WHEN representative_rank = 1 THEN parse_status END) AS parse_status + FROM filtered + GROUP BY {group_key} + ) + SELECT *, datetime(timestamp, 'unixepoch') AS timestamp_str + FROM grouped + ORDER BY {sort_expression} {direction}, id {direction} + LIMIT ? OFFSET ? + """, + [*params, limit, offset], + ).fetchall() + else: + total_count = cursor.execute( + f""" + SELECT COUNT(*) + FROM traceroute_routes r + JOIN packet_history p ON p.id = r.packet_id + WHERE {where} + """, + params, + ).fetchone()[0] + ungrouped_sort = { + "gateway_count": "p.gateway_id", + "min_rssi": "p.rssi", + "min_snr": "p.snr", + "min_hops": "(p.hop_start - p.hop_limit)", + "avg_payload_length": "p.payload_length", + }.get(sort_expression, f"r.{sort_expression}") + rows = cursor.execute( + f""" + SELECT + p.id, r.timestamp, r.from_node_id, r.to_node_id, + p.gateway_id, p.channel_id, p.hop_start, p.hop_limit, + p.rssi, p.snr, p.payload_length, p.processed_successfully, + r.mesh_packet_id, r.route_nodes_json, r.snr_towards_json, + r.route_back_json, r.snr_back_json, r.parse_status, + r.route_nodes_json AS route, + (p.hop_start - p.hop_limit) AS hop_count, + datetime(r.timestamp, 'unixepoch') AS timestamp_str + FROM traceroute_routes r + JOIN packet_history p ON p.id = r.packet_id + WHERE {where} + ORDER BY {ungrouped_sort} {direction}, p.id {direction} + LIMIT ? OFFSET ? + """, + [*params, limit, offset], + ).fetchall() + + packets = [dict(row) for row in rows] + for packet in packets: + packet["route"] = packet["route_nodes_json"] + if group_packets: + packet["is_grouped"] = True + packet["rssi_range"] = _range(packet["min_rssi"], packet["max_rssi"], "dBm", 1) + packet["snr_range"] = _range(packet["min_snr"], packet["max_snr"], "dB", 2) + packet["hop_range"] = _range(packet["min_hops"], packet["max_hops"], "", 0) + packet["rssi"] = packet["rssi_range"] + packet["snr"] = packet["snr_range"] + packet["hop_count"] = packet["min_hops"] + return { + "packets": packets, + "total_count": total_count, + } + finally: + conn.close() + + +def get_traceroute_link( + node1_id: int, + node2_id: int, + *, + start_time: float, + end_time: float, + limit: int, + offset: int, +) -> dict[str, Any]: + """Return aggregate link statistics and one page of matching traceroutes. + + The hop table identifies matches without decoding or scanning traceroute + payloads. Statistics cover every matching hop in the requested window, while + the detail query returns each matching packet once. + """ + conn = get_db_connection() + try: + cursor = conn.cursor() + match_sql = """ + h.timestamp >= ? AND h.timestamp <= ? + AND ((h.from_node_id = ? AND h.to_node_id = ?) + OR (h.from_node_id = ? AND h.to_node_id = ?)) + """ + match_params = [ + start_time, + end_time, + node1_id, + node2_id, + node2_id, + node1_id, + ] + stats = cursor.execute( + f""" + SELECT + COUNT(DISTINCT h.packet_id) AS total_attempts, + SUM(CASE WHEN h.from_node_id = ? AND h.to_node_id = ? + THEN 1 ELSE 0 END) AS forward_count, + SUM(CASE WHEN h.from_node_id = ? AND h.to_node_id = ? + THEN 1 ELSE 0 END) AS reverse_count, + AVG(CASE WHEN h.snr = {TRACEROUTE_UNKNOWN_SNR} + OR h.snr BETWEEN {SNR_PLAUSIBLE_MIN} AND {SNR_PLAUSIBLE_MAX} + THEN h.snr END) AS avg_snr + FROM traceroute_hops h + JOIN traceroute_routes r ON r.packet_id = h.packet_id + WHERE r.parser_version = ? AND r.parse_status = 'parsed' + AND {match_sql} + """, + [ + node1_id, + node2_id, + node2_id, + node1_id, + PARSER_VERSION, + *match_params, + ], + ).fetchone() + + total_count = int(stats["total_attempts"] or 0) + rows = cursor.execute( + f""" + WITH matching AS ( + SELECT + h.packet_id, + h.from_node_id AS target_from_node_id, + h.to_node_id AS target_to_node_id, + h.snr AS target_hop_snr, + ROW_NUMBER() OVER ( + PARTITION BY h.packet_id + ORDER BY CASE h.direction WHEN 'forward' THEN 0 ELSE 1 END, + h.hop_index + ) AS target_rank + FROM traceroute_hops h + JOIN traceroute_routes r ON r.packet_id = h.packet_id + WHERE r.parser_version = ? AND r.parse_status = 'parsed' + AND {match_sql} + ) + SELECT + p.id, r.timestamp, r.from_node_id, r.to_node_id, + p.gateway_id, p.channel_id, p.hop_start, p.hop_limit, + p.rssi, p.snr, p.payload_length, p.processed_successfully, + r.mesh_packet_id, r.route_nodes_json, r.snr_towards_json, + r.route_back_json, r.snr_back_json, r.parse_status, + m.target_from_node_id, m.target_to_node_id, m.target_hop_snr, + datetime(r.timestamp, 'unixepoch') AS timestamp_str + FROM matching m + JOIN traceroute_routes r ON r.packet_id = m.packet_id + JOIN packet_history p ON p.id = m.packet_id + WHERE m.target_rank = 1 + ORDER BY r.timestamp DESC, p.id DESC + LIMIT ? OFFSET ? + """, + [PARSER_VERSION, *match_params, limit, offset], + ).fetchall() + + return { + "packets": [dict(row) for row in rows], + "total_count": total_count, + "total_attempts": total_count, + "forward_count": int(stats["forward_count"] or 0), + "reverse_count": int(stats["reverse_count"] or 0), + "avg_snr": stats["avg_snr"], + } + finally: + conn.close() + + +def _range(minimum: Any, maximum: Any, unit: str, decimals: int) -> str | None: + if minimum is None or maximum is None: + return None + suffix = f" {unit}" if unit else "" + if decimals: + low = f"{minimum:.{decimals}f}" + high = f"{maximum:.{decimals}f}" + else: + low, high = str(minimum), str(maximum) + return f"{low}{suffix}" if minimum == maximum else f"{low}-{high}{suffix}" + + +def route_data_from_row(packet: dict[str, Any]) -> dict[str, list[Any]] | None: + """Decode the four saved JSON arrays without touching the packet payload.""" + keys = ("route_nodes", "snr_towards", "route_back", "snr_back") + try: + return {key: json.loads(packet[f"{key}_json"]) for key in keys} + except (KeyError, TypeError, json.JSONDecodeError): + return None + + +def get_traceroute_hops_for_graph( + filters: dict[str, Any] | None = None, +) -> list[dict[str, Any]]: + """Return the complete hop sequences of parsed traceroutes for graph building. + + Hops are returned in path order with their SNR untouched so callers can + preserve path structure: dropping individual rows here (invalid SNR, weak + links) would splice the remaining hops of one route into a shorter, + connected-looking path. Quality filtering and continuity validation are the + caller's responsibility. + """ + filters = dict(filters or {}) + conn = get_db_connection() + try: + cursor = conn.cursor() + conditions = [ + "r.parser_version = ?", + "r.parse_status = 'parsed'", + "h.from_node_id != 4294967295", + "h.to_node_id != 4294967295", + ] + params: list[Any] = [PARSER_VERSION] + + if filters.get("start_time") is not None: + conditions.append("h.timestamp >= ?") + params.append(filters["start_time"]) + if filters.get("end_time") is not None: + conditions.append("h.timestamp <= ?") + params.append(filters["end_time"]) + + join_packet = False + if filters.get("gateway_id"): + join_packet = True + conditions.append("p.gateway_id = ?") + params.append(filters["gateway_id"]) + if filters.get("primary_channel"): + join_packet = True + conditions.append("p.channel_id = ?") + params.append(filters["primary_channel"]) + + join_clause = "JOIN packet_history p ON p.id = h.packet_id" if join_packet else "" + where_clause = " AND ".join(conditions) + + query = f""" + SELECT + h.packet_id, + h.direction, + h.hop_index, + h.timestamp, + h.from_node_id, + h.to_node_id, + h.snr + FROM traceroute_hops h + JOIN traceroute_routes r ON r.packet_id = h.packet_id + {join_clause} + WHERE {where_clause} + ORDER BY h.timestamp DESC, h.packet_id, h.direction, h.hop_index + """ + rows = cursor.execute(query, params).fetchall() + return [dict(row) for row in rows] + finally: + conn.close() + + +def get_traceroute_hops_for_longest_links( + start_time: float, + end_time: float, +) -> list[dict[str, Any]]: + """Return all RF hops from traceroute_hops in the specified time window.""" + return get_traceroute_hops_for_graph( + filters={"start_time": start_time, "end_time": end_time} + ) + + +def get_route_patterns_data( + start_time: float, + end_time: float, + limit: int = 50, +) -> dict[str, Any]: + """Aggregate route patterns directly in SQL from traceroute_routes.""" + conn = get_db_connection() + try: + cursor = conn.cursor() + total_analyzed = cursor.execute( + """ + SELECT COUNT(*) + FROM traceroute_routes + WHERE parser_version = ? AND parse_status = 'parsed' + AND timestamp >= ? AND timestamp <= ? + """, + (PARSER_VERSION, start_time, end_time), + ).fetchone()[0] + + query = """ + WITH ranked_patterns AS ( + SELECT + p.id AS packet_id, + r.timestamp, + r.from_node_id, + r.to_node_id, + r.route_nodes_json, + ROW_NUMBER() OVER ( + PARTITION BY r.from_node_id, r.to_node_id, r.route_nodes_json + ORDER BY r.timestamp DESC, p.id DESC + ) AS rank, + COUNT(*) OVER ( + PARTITION BY r.from_node_id, r.to_node_id, r.route_nodes_json + ) AS pattern_count + FROM traceroute_routes r + JOIN packet_history p ON p.id = r.packet_id + WHERE r.parser_version = ? AND r.parse_status = 'parsed' + AND r.route_nodes_json != '[]' + AND r.timestamp >= ? AND r.timestamp <= ? + ) + SELECT + packet_id, + timestamp, + from_node_id, + to_node_id, + route_nodes_json, + pattern_count, + rank + FROM ranked_patterns + WHERE rank <= 3 + ORDER BY pattern_count DESC, from_node_id, to_node_id + """ + rows = cursor.execute(query, (PARSER_VERSION, start_time, end_time)).fetchall() + + route_patterns: dict[tuple[tuple[int, int], tuple[int, ...]], dict[str, Any]] = {} + + for row in rows: + try: + route_nodes = tuple(json.loads(row["route_nodes_json"])) + except (TypeError, ValueError, json.JSONDecodeError): + continue + if not route_nodes: + continue + + endpoints = tuple(sorted([row["from_node_id"], row["to_node_id"]])) + pattern_key = (endpoints, route_nodes) + + if pattern_key not in route_patterns: + route_patterns[pattern_key] = { + "count": 0, + "endpoints": endpoints, + # Legacy placeholder / dead code: always 0, never computed or updated; + # retained for API response schema backwards-compatibility. + "avg_success_rate": 0, + "examples": [], + } + + if row["rank"] == 1: + route_patterns[pattern_key]["count"] += row["pattern_count"] + + if len(route_patterns[pattern_key]["examples"]) < 3: + route_patterns[pattern_key]["examples"].append( + { + "packet_id": row["packet_id"], + "timestamp": row["timestamp"], + "from_node": row["from_node_id"], + "to_node": row["to_node_id"], + } + ) + + sorted_patterns = sorted( + route_patterns.items(), key=lambda x: x[1]["count"], reverse=True + )[:limit] + + return { + "sorted_patterns": sorted_patterns, + "total_patterns": len(route_patterns), + "analyzed_traceroutes": total_analyzed, + } + finally: + conn.close() + + +def get_node_traceroute_statistics( + node_id: int, + start_time: float | None = None, + end_time: float | None = None, +) -> dict[str, Any]: + """Compute traceroute source, destination, and intermediate participation in SQL.""" + conn = get_db_connection() + try: + cursor = conn.cursor() + time_conditions = [] + time_params: list[Any] = [] + if start_time is not None: + time_conditions.append("timestamp >= ?") + time_params.append(start_time) + if end_time is not None: + time_conditions.append("timestamp <= ?") + time_params.append(end_time) + + time_sql = f" AND {' AND '.join(time_conditions)}" if time_conditions else "" + + # Query 1: As source and as destination + endpoint_query = f""" + SELECT + COUNT(CASE WHEN from_node_id = ? THEN 1 END) AS source_total, + COUNT(CASE WHEN from_node_id = ? AND parse_status IN ('parsed', 'valid_empty') THEN 1 END) AS source_successful, + COUNT(CASE WHEN to_node_id = ? THEN 1 END) AS dest_total, + COUNT(CASE WHEN to_node_id = ? AND parse_status IN ('parsed', 'valid_empty') THEN 1 END) AS dest_successful + FROM traceroute_routes + WHERE parser_version = ? + AND (from_node_id = ? OR to_node_id = ?) + {time_sql} + """ + endpoint_params = [ + node_id, + node_id, + node_id, + node_id, + PARSER_VERSION, + node_id, + node_id, + *time_params, + ] + row = cursor.execute(endpoint_query, endpoint_params).fetchone() + source_total = int(row["source_total"] or 0) + source_successful = int(row["source_successful"] or 0) + dest_total = int(row["dest_total"] or 0) + dest_successful = int(row["dest_successful"] or 0) + + # Query 2: Intermediate participation + intermediate_query = f""" + SELECT COUNT(*) + FROM traceroute_routes + WHERE parser_version = ? + AND parse_status IN ('parsed', 'valid_empty') + {time_sql} + AND EXISTS ( + SELECT 1 FROM json_each(route_nodes_json) WHERE value = ? + ) + """ + intermediate_params = [ + PARSER_VERSION, + *time_params, + node_id, + ] + participation_count = cursor.execute(intermediate_query, intermediate_params).fetchone()[0] + + return { + "node_id": node_id, + "as_source": { + "total": source_total, + "successful": source_successful, + "success_rate": (source_successful / source_total * 100) if source_total > 0 else 0, + }, + "as_destination": { + "total": dest_total, + "successful": dest_successful, + "success_rate": (dest_successful / dest_total * 100) if dest_total > 0 else 0, + }, + "as_intermediate_hop": {"participation_count": participation_count}, + "total_involvement": source_total + dest_total + participation_count, + } + finally: + conn.close() diff --git a/src/malla/database/traceroute_schema.py b/src/malla/database/traceroute_schema.py new file mode 100644 index 00000000..c26dea50 --- /dev/null +++ b/src/malla/database/traceroute_schema.py @@ -0,0 +1,109 @@ +"""Rebuildable traceroute tables and pending records for raw-only writers.""" + +import sqlite3 + +TRACEROUTE_PORT = 70 +TRACEROUTE_PREDICATE = "(portnum = 70 OR portnum_name = 'TRACEROUTE_APP')" + + +def ensure_traceroute_schema(cursor: sqlite3.Cursor) -> None: + """Create empty derived tables; historical decoding is always explicit. + + Triggers leave raw-only inserts/updates pending, including imports and older + capture versions. The shared writer replaces the pending record in the same + transaction. These triggers never decode payloads or examine packet history. + """ + cursor.execute(""" + CREATE TABLE IF NOT EXISTS traceroute_routes ( + packet_id INTEGER PRIMARY KEY REFERENCES packet_history(id) ON DELETE CASCADE, + timestamp REAL NOT NULL, + mesh_packet_id INTEGER, + from_node_id INTEGER, + to_node_id INTEGER, + route_nodes_json TEXT NOT NULL DEFAULT '[]', + snr_towards_json TEXT NOT NULL DEFAULT '[]', + route_back_json TEXT NOT NULL DEFAULT '[]', + snr_back_json TEXT NOT NULL DEFAULT '[]', + forward_complete INTEGER NOT NULL DEFAULT 0 CHECK (forward_complete IN (0, 1)), + return_complete INTEGER NOT NULL DEFAULT 0 CHECK (return_complete IN (0, 1)), + parse_status TEXT NOT NULL DEFAULT 'pending' + CHECK (parse_status IN ('pending', 'parsed', 'valid_empty', 'invalid_payload')), + parse_error TEXT, + parser_version INTEGER NOT NULL DEFAULT 0, + materialized_at REAL, + CHECK ((parse_status = 'pending' AND parser_version = 0) + OR (parse_status != 'pending' AND parser_version > 0)), + CHECK ((parse_status = 'invalid_payload' AND parse_error IS NOT NULL) + OR (parse_status != 'invalid_payload' AND parse_error IS NULL)) + ) + """) + cursor.execute(""" + CREATE TABLE IF NOT EXISTS traceroute_hops ( + packet_id INTEGER NOT NULL REFERENCES traceroute_routes(packet_id) ON DELETE CASCADE, + direction TEXT NOT NULL CHECK (direction IN ('forward', 'return')), + hop_index INTEGER NOT NULL CHECK (hop_index >= 0), + timestamp REAL NOT NULL, + from_node_id INTEGER NOT NULL, + to_node_id INTEGER NOT NULL, + snr REAL, + PRIMARY KEY (packet_id, direction, hop_index) + ) + """) + + for name, table, columns in ( + ("routes_time", "traceroute_routes", "timestamp"), + ("routes_source_time", "traceroute_routes", "from_node_id, timestamp"), + ("routes_target_time", "traceroute_routes", "to_node_id, timestamp"), + ( + "routes_group_time", + "traceroute_routes", + "mesh_packet_id, from_node_id, to_node_id, timestamp", + ), + ("routes_version", "traceroute_routes", "parser_version"), + ("hops_link_time", "traceroute_hops", "from_node_id, to_node_id, timestamp"), + ("hops_time_link", "traceroute_hops", "timestamp, from_node_id, to_node_id"), + ("hops_source_time", "traceroute_hops", "from_node_id, timestamp"), + ("hops_target_time", "traceroute_hops", "to_node_id, timestamp"), + ): + cursor.execute( + f"CREATE INDEX IF NOT EXISTS idx_traceroute_{name} ON {table} ({columns})" + ) + + # Explicit child deletion also handles raw import connections with foreign + # keys disabled, and INSERT OR REPLACE reusing a packet row ID. + for event, operation in ( + ("insert", "INSERT"), + ( + "update", + "UPDATE OF id, timestamp, mesh_packet_id, from_node_id, to_node_id, " + "portnum, portnum_name, raw_payload, hop_start, hop_limit", + ), + ): + old_cleanup = ( + "DELETE FROM traceroute_hops WHERE packet_id = OLD.id; " + "DELETE FROM traceroute_routes WHERE packet_id = OLD.id;" + if event == "update" + else "" + ) + cursor.execute(f""" + CREATE TRIGGER IF NOT EXISTS traceroute_packet_{event} + AFTER {operation} ON packet_history + BEGIN + {old_cleanup} + DELETE FROM traceroute_hops WHERE packet_id = NEW.id; + DELETE FROM traceroute_routes WHERE packet_id = NEW.id; + INSERT INTO traceroute_routes + (packet_id, timestamp, mesh_packet_id, from_node_id, to_node_id) + SELECT NEW.id, NEW.timestamp, NEW.mesh_packet_id, + NEW.from_node_id, NEW.to_node_id + WHERE NEW.portnum = {TRACEROUTE_PORT} OR NEW.portnum_name = 'TRACEROUTE_APP'; + END + """) + cursor.execute(""" + CREATE TRIGGER IF NOT EXISTS traceroute_packet_delete + AFTER DELETE ON packet_history + BEGIN + DELETE FROM traceroute_hops WHERE packet_id = OLD.id; + DELETE FROM traceroute_routes WHERE packet_id = OLD.id; + END + """) diff --git a/src/malla/database/traceroutes.py b/src/malla/database/traceroutes.py new file mode 100644 index 00000000..696114f8 --- /dev/null +++ b/src/malla/database/traceroutes.py @@ -0,0 +1,198 @@ +"""Shared capture/backfill decoder, writer, and traceroute readiness checks.""" + +import json +import sqlite3 +import time +from dataclasses import dataclass +from typing import Any + +from google.protobuf.message import DecodeError + +from ..models.traceroute import TracerouteHop, TraceroutePacket +from ..utils.traceroute_utils import RouteData, decode_traceroute_payload +from .traceroute_schema import TRACEROUTE_PREDICATE + +PARSER_VERSION = 1 + + +@dataclass(frozen=True) +class DecodedTraceroute: + route: RouteData + hops: tuple[TracerouteHop, ...] = () + forward_complete: bool = False + return_complete: bool = False + parse_error: str | None = None + + @property + def parse_status(self) -> str: + if self.parse_error is not None: + return "invalid_payload" + return "parsed" if any(self.route.values()) else "valid_empty" + + +def decode_traceroute(packet: dict[str, Any]) -> DecodedTraceroute: + """Decode once, without database/name/location lookups or discarded hops.""" + try: + if packet.get("raw_payload") is None: + raise TypeError("Traceroute payload is missing") + route = decode_traceroute_payload(packet["raw_payload"]) + except (DecodeError, TypeError) as exc: + return DecodedTraceroute( + route=RouteData(route_nodes=[], snr_towards=[], route_back=[], snr_back=[]), + parse_error=str(exc), + ) + + traceroute = TraceroutePacket( + packet, resolve_names=False, pre_parsed_route_data=route + ) + return DecodedTraceroute( + route=route, + hops=tuple(traceroute.get_rf_hops()), + forward_complete=traceroute.is_complete(), + return_complete=traceroute.is_return_complete(), + ) + + +def write_traceroute(cursor: sqlite3.Cursor, packet: dict[str, Any]) -> None: + """Replace one reception's decoded data inside the caller's transaction. + + The caller must already have inserted the raw packet. This function neither + commits nor swallows storage errors, so raw and derived data stay atomic. + """ + if not cursor.connection.in_transaction: + raise ValueError("Traceroute writes require an active transaction") + decoded = decode_traceroute(packet) + cursor.execute( + """ + INSERT INTO traceroute_routes ( + packet_id, timestamp, mesh_packet_id, from_node_id, to_node_id, + route_nodes_json, snr_towards_json, route_back_json, snr_back_json, + forward_complete, return_complete, parse_status, parse_error, + parser_version, materialized_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(packet_id) DO UPDATE SET + timestamp = excluded.timestamp, + mesh_packet_id = excluded.mesh_packet_id, + from_node_id = excluded.from_node_id, + to_node_id = excluded.to_node_id, + route_nodes_json = excluded.route_nodes_json, + snr_towards_json = excluded.snr_towards_json, + route_back_json = excluded.route_back_json, + snr_back_json = excluded.snr_back_json, + forward_complete = excluded.forward_complete, + return_complete = excluded.return_complete, + parse_status = excluded.parse_status, + parse_error = excluded.parse_error, + parser_version = excluded.parser_version, + materialized_at = excluded.materialized_at + """, + ( + packet["id"], + packet["timestamp"], + packet.get("mesh_packet_id"), + packet.get("from_node_id"), + packet.get("to_node_id"), + json.dumps(decoded.route["route_nodes"]), + json.dumps(decoded.route["snr_towards"]), + json.dumps(decoded.route["route_back"]), + json.dumps(decoded.route["snr_back"]), + decoded.forward_complete, + decoded.return_complete, + decoded.parse_status, + decoded.parse_error, + PARSER_VERSION, + time.time(), + ), + ) + cursor.execute("DELETE FROM traceroute_hops WHERE packet_id = ?", (packet["id"],)) + indices = {"forward_rf": 0, "return_rf": 0} + for hop in decoded.hops: + cursor.execute( + """ + INSERT INTO traceroute_hops + (packet_id, direction, hop_index, timestamp, from_node_id, to_node_id, snr) + VALUES (?, ?, ?, ?, ?, ?, ?) + """, + ( + packet["id"], + "forward" if hop.direction == "forward_rf" else "return", + indices[hop.direction], + packet["timestamp"], + hop.from_node_id, + hop.to_node_id, + hop.snr, + ), + ) + indices[hop.direction] += 1 + + +def inspect_traceroutes(cursor: sqlite3.Cursor) -> dict[str, Any]: + """Audit stored history explicitly, including databases not yet prepared. + + This is an offline validation query, not a web request helper. The caller + holds a transaction for a consistent snapshot. + """ + tables = { + row[0] + for row in cursor.execute("SELECT name FROM sqlite_master WHERE type = 'table'") + } + raw = cursor.execute( + f"SELECT COUNT(*) FROM packet_history WHERE {TRACEROUTE_PREDICATE}" + ).fetchone()[0] + counts: dict[str, Any] = { + "raw_traceroutes": raw, + "routes": 0, + "parsed": 0, + "valid_empty": 0, + "invalid_payload": 0, + "pending": 0, + "outdated": 0, + "missing": raw, + "hops": 0, + "orphan_routes": 0, + "orphan_hops": 0, + "parser_version": PARSER_VERSION, + "complete": False, + } + if not {"traceroute_routes", "traceroute_hops"} <= tables: + return counts + + for status, version, count in cursor.execute( + "SELECT parse_status, parser_version, COUNT(*) FROM traceroute_routes " + "GROUP BY parse_status, parser_version" + ): + counts["routes"] += count + counts[status] += count + if version not in (0, PARSER_VERSION): + counts["outdated"] += count + counts["missing"] = cursor.execute( + f""" + SELECT COUNT(*) FROM packet_history + WHERE {TRACEROUTE_PREDICATE} AND NOT EXISTS ( + SELECT 1 FROM traceroute_routes WHERE packet_id = packet_history.id + ) + """ + ).fetchone()[0] + counts["hops"] = cursor.execute("SELECT COUNT(*) FROM traceroute_hops").fetchone()[ + 0 + ] + counts["orphan_routes"] = cursor.execute( + f""" + SELECT COUNT(*) FROM traceroute_routes WHERE NOT EXISTS ( + SELECT 1 FROM packet_history + WHERE id = traceroute_routes.packet_id AND {TRACEROUTE_PREDICATE} + ) + """ + ).fetchone()[0] + counts["orphan_hops"] = cursor.execute( + """ + SELECT COUNT(*) FROM traceroute_hops WHERE NOT EXISTS ( + SELECT 1 FROM traceroute_routes WHERE packet_id = traceroute_hops.packet_id + ) + """ + ).fetchone()[0] + counts["complete"] = not any( + counts[key] + for key in ("missing", "pending", "outdated", "orphan_routes", "orphan_hops") + ) + return counts diff --git a/src/malla/models/traceroute.py b/src/malla/models/traceroute.py index 55e24fba..a19043eb 100644 --- a/src/malla/models/traceroute.py +++ b/src/malla/models/traceroute.py @@ -6,6 +6,7 @@ interface for analyzing traceroute packets and extracting path information. """ +import json import logging import math from dataclasses import dataclass @@ -105,6 +106,17 @@ def __init__( self.to_node_name: str | None = None # Parse (or accept pre-parsed) traceroute payload + if pre_parsed_route_data is None: + try: + pre_parsed_route_data = RouteData( + route_nodes=json.loads(packet_data["route_nodes_json"]), + snr_towards=json.loads(packet_data["snr_towards_json"]), + route_back=json.loads(packet_data["route_back_json"]), + snr_back=json.loads(packet_data["snr_back_json"]), + ) + except (KeyError, TypeError, json.JSONDecodeError): + pass + if pre_parsed_route_data is not None: self.route_data = pre_parsed_route_data else: diff --git a/src/malla/mqtt_capture.py b/src/malla/mqtt_capture.py index a9831877..1c426a32 100644 --- a/src/malla/mqtt_capture.py +++ b/src/malla/mqtt_capture.py @@ -38,6 +38,7 @@ import sqlite3 import threading import time +from contextlib import closing from typing import Any import paho.mqtt.client as mqtt @@ -59,6 +60,8 @@ from .database.connection import seed_query_planner_stats_async from .database.schema import ensure_startup_schema +from .database.traceroutes import write_traceroute +from .utils.geo_utils import is_valid_position # Load the singleton configuration once at module import time. This ensures the # capture tool honours the same YAML + optional environment override mechanism @@ -846,10 +849,10 @@ def log_packet_to_database( relay_node = getattr(mesh_packet, "relay_node", None) if mesh_packet else None tx_after = getattr(mesh_packet, "tx_after", None) if mesh_packet else None - with db_lock: - conn = sqlite3.connect(DATABASE_FILE, timeout=30.0) + with db_lock, closing(sqlite3.connect(DATABASE_FILE, timeout=30.0)) as conn, conn: conn.row_factory = sqlite3.Row cursor = conn.cursor() + cursor.execute("PRAGMA foreign_keys=ON") cursor.execute( """ @@ -894,8 +897,20 @@ def log_packet_to_database( ), ) - conn.commit() - conn.close() + if portnum == portnums_pb2.PortNum.TRACEROUTE_APP: + write_traceroute( + cursor, + { + "id": cursor.lastrowid, + "timestamp": current_time, + "mesh_packet_id": mesh_packet_id, + "from_node_id": from_node_id, + "to_node_id": to_node_id, + "hop_start": hop_start, + "hop_limit": hop_limit, + "raw_payload": raw_payload, + }, + ) def get_packet_history( @@ -945,6 +960,7 @@ def cleanup_old_data() -> None: conn = sqlite3.connect(DATABASE_FILE, timeout=30.0) conn.row_factory = sqlite3.Row cursor = conn.cursor() + cursor.execute("PRAGMA foreign_keys=ON") try: # Delete old packet history records @@ -1250,9 +1266,19 @@ def on_message(client: mqtt.Client, userdata: Any, msg: mqtt.MQTTMessage) -> Non via_mqtt_str = ( " (via MQTT)" if getattr(mesh_packet, "via_mqtt", False) else "" ) - logging.info( - f"📍 Position from {from_node_display}{via_mqtt_str}: {lat:.5f}, {lon:.5f} (alt: {alt}m)" - ) + if not is_valid_position( + lat if position_data.latitude_i else None, + lon if position_data.longitude_i else None, + ): + logging.warning( + f"⚠️ Invalid position from {from_node_display}{via_mqtt_str}: " + f"{lat:.5f}, {lon:.5f} (near null island or out of range) - " + f"will be ignored by UI queries" + ) + else: + logging.info( + f"📍 Position from {from_node_display}{via_mqtt_str}: {lat:.5f}, {lon:.5f} (alt: {alt}m)" + ) processed_successfully = True elif mesh_packet.decoded.portnum == portnums_pb2.PortNum.NODEINFO_APP: diff --git a/src/malla/routes/api_routes.py b/src/malla/routes/api_routes.py index 2e097e4c..d11304a7 100644 --- a/src/malla/routes/api_routes.py +++ b/src/malla/routes/api_routes.py @@ -17,12 +17,16 @@ TracerouteRepository, get_db_connection, ) +from ..database.traceroute_read_repository import ( + get_traceroute_link, + route_data_from_row, +) from ..models.traceroute import TraceroutePacket from ..services.analytics_service import AnalyticsService from ..services.location_service import LocationService from ..services.meshtastic_service import MeshtasticService from ..services.node_service import NodeService -from ..services.traceroute_service import DEFAULT_GRAPH_PACKET_LIMIT, TracerouteService +from ..services.traceroute_service import TracerouteService from ..utils.node_utils import ( convert_node_id, get_bulk_node_names, @@ -30,7 +34,6 @@ ) from ..utils.serialization_utils import convert_bytes_to_base64, sanitize_floats from ..utils.signal_quality import is_plausible_traceroute_snr -from ..utils.traceroute_utils import parse_traceroute_payload logger = logging.getLogger(__name__) api_bp = Blueprint("api", __name__, url_prefix="/api") @@ -717,54 +720,97 @@ def api_traceroute_details(packet_id): def api_locations(): """ API endpoint for node location data with network topology. - Returns up to 14 days of data for client-side filtering. + + Link aggregates (traceroute RF hops and direct packet receptions) are + computed server-side over the exact time window requested via the + ``start_time``/``end_time`` (epoch seconds) or ``hours``/``max_age_hours`` + request parameters, so observation counts and signal statistics always + match the period selected on the map. + + Node positions deliberately keep a wide lookup window (14 days) + independent of the link window: an actively routing node whose last + position report is older than the selected window must stay visible at + its last known good position instead of disappearing. """ import time start_time_perf = time.time() logger.info("API locations endpoint accessed") try: - # Build filters from request parameters - filters = {} - - # Always limit to last 14 days for performance from datetime import datetime, timedelta - end_time = datetime.now() - start_time = end_time - timedelta(days=14) - filters["start_time"] = start_time.timestamp() - filters["end_time"] = end_time.timestamp() + now = datetime.now() + max_window_seconds = 14 * 24 * 3600 + + # Wide position-lookup window (performance cap only). This window is + # intentionally NOT narrowed by the client's time selection so nodes + # with stale GPS fixes remain visible while active. + position_filters: dict[str, Any] = { + "start_time": (now - timedelta(seconds=max_window_seconds)).timestamp(), + "end_time": now.timestamp(), + } + + # ------------------------------------------------------------------ + # Resolve the link aggregation window from request parameters. + # Explicit start/end win; otherwise hours/max_age_hours is applied + # relative to now; with no time parameter at all the endpoint keeps + # its historical 14-day default. + # ------------------------------------------------------------------ + start_arg = request.args.get("start_time", type=float) + end_arg = request.args.get("end_time", type=float) + hours_arg = request.args.get("hours", type=float) + if hours_arg is None: + hours_arg = request.args.get("max_age_hours", type=float) + + link_filters: dict[str, Any] = {} + if start_arg is not None or end_arg is not None or hours_arg: + if start_arg is None: + lookback = hours_arg * 3600 if hours_arg else max_window_seconds + start_arg = (end_arg or now.timestamp()) - lookback + if end_arg is None: + end_arg = now.timestamp() + if start_arg >= end_arg: + return jsonify({"error": "start_time must be before end_time"}), 400 + # Cap the aggregation window for performance + if end_arg - start_arg > max_window_seconds: + start_arg = end_arg - max_window_seconds + link_filters["start_time"] = start_arg + link_filters["end_time"] = end_arg + else: + link_filters.update(position_filters) # Gateway filter (keep this server-side for performance) gateway_id_arg = request.args.get("gateway_id") if gateway_id_arg is not None: try: - filters["gateway_id"] = int(gateway_id_arg) + gateway_id = int(gateway_id_arg) except ValueError: return jsonify({"error": "Invalid gateway_id format"}), 400 + link_filters["gateway_id"] = gateway_id + position_filters["gateway_id"] = gateway_id # Search filter (keep this server-side for performance) if request.args.get("search"): - filters["search"] = request.args.get("search") + link_filters["search"] = request.args.get("search") + position_filters["search"] = request.args.get("search") # ------------------------------------------------------------------ # OPTIMIZATION: Call expensive operations ONCE and pass results down # ------------------------------------------------------------------ # 1. Get network topology data (used by both get_node_locations and get_traceroute_links) - from ..services.traceroute_service import TracerouteService - - hours = 24 # Default to 24 hours for network analysis - time_diff = filters["end_time"] - filters["start_time"] - hours = max(1, min(168, int(time_diff / 3600))) # Between 1 and 168 hours - network_filters = {} - if filters.get("start_time"): - network_filters["start_time"] = filters["start_time"] - if filters.get("end_time"): - network_filters["end_time"] = filters["end_time"] - if filters.get("gateway_id"): - network_filters["gateway_id"] = filters["gateway_id"] + if link_filters.get("start_time"): + network_filters["start_time"] = link_filters["start_time"] + if link_filters.get("end_time"): + network_filters["end_time"] = link_filters["end_time"] + if link_filters.get("gateway_id"): + network_filters["gateway_id"] = link_filters["gateway_id"] + + # The explicit start/end filters take precedence inside the service; + # hours is kept consistent with the resolved window for cache keys. + time_diff = link_filters["end_time"] - link_filters["start_time"] + hours = max(1, min(168, int(time_diff / 3600))) # Between 1 and 168 hours network_data = TracerouteService.get_network_graph_data( hours=hours, @@ -773,16 +819,18 @@ def api_locations(): ) # 2. Get packet links (used by get_node_locations and returned in response) - packet_links = LocationService.get_packet_links(filters) + packet_links = LocationService.get_packet_links(link_filters) - # 3. Get enhanced location data, passing pre-computed data + # 3. Get enhanced location data, passing pre-computed data. Position + # lookups use the wide window; node activity timestamps come from + # the window-scoped network/packet data computed above. locations = LocationService.get_node_locations( - filters, network_data=network_data, packet_links=packet_links + position_filters, network_data=network_data, packet_links=packet_links ) # 4. Get traceroute links, passing pre-computed network data traceroute_links = LocationService.get_traceroute_links( - filters, network_data=network_data + link_filters, network_data=network_data ) duration = time.time() - start_time_perf @@ -794,7 +842,7 @@ def api_locations(): "traceroute_links": traceroute_links, "packet_links": packet_links, "total_count": len(locations) if isinstance(locations, list) else 0, - "filters_applied": filters, + "filters_applied": link_filters, "data_period_days": 14, } ) @@ -1041,13 +1089,10 @@ def api_traceroute_hops_nodes(): conn = get_db_connection() cursor = conn.cursor() - # Collect all nodes involved in recent traceroutes: initiators, targets - # and intermediate route nodes. Route arrays live inside raw_payload - # blobs so they cannot be extracted in SQL; parse them in Python using - # the same 7-day window as the hop/link analysis endpoints so every - # listed node can actually be analyzed (and so map "View History" - # links resolve for hop-only nodes). The parse is expensive, so the - # resulting id set is cached briefly. + # Collect initiators, targets, and intermediate RF-hop endpoints from + # the saved tables over the same seven-day window as link analysis. + # Cache the resulting ID set briefly; readiness is checked before the + # cache so a new raw-only import cannot look complete. global _hop_nodes_cache now = time.time() if ( @@ -1060,38 +1105,40 @@ def api_traceroute_hops_nodes(): window_end = datetime.now() window_start = window_end - timedelta(days=7) - cursor.execute( - f""" - SELECT - id, from_node_id, to_node_id, raw_payload - FROM packet_history - WHERE portnum_name = 'TRACEROUTE_APP' - AND processed_successfully = 1 - AND raw_payload IS NOT NULL - AND timestamp >= ? AND timestamp <= ? - ORDER BY timestamp DESC - LIMIT {DEFAULT_GRAPH_PACKET_LIMIT} - """, - (window_start.timestamp(), window_end.timestamp()), - ) - BROADCAST_NODE_ID = 4294967295 # 0xFFFFFFFF - hop_node_ids = set() - for packet in cursor.fetchall(): - # RF hop endpoints are consecutive nodes of the sequence - # from -> route... -> to (and the same for the return path), - # so the union of all those node ids is exactly the set of - # nodes that can appear as a hop endpoint. - route_data = parse_traceroute_payload(packet["raw_payload"]) - packet_node_ids = [ - packet["from_node_id"], - packet["to_node_id"], - *route_data["route_nodes"], - *route_data["route_back"], - ] - for hop_node_id in packet_node_ids: - if hop_node_id is not None and hop_node_id != BROADCAST_NODE_ID: - hop_node_ids.add(hop_node_id) + hop_node_ids = { + row[0] + for row in cursor.execute( + """ + SELECT from_node_id FROM traceroute_hops + WHERE timestamp >= ? AND timestamp <= ? + UNION + SELECT to_node_id FROM traceroute_hops + WHERE timestamp >= ? AND timestamp <= ? + UNION + SELECT from_node_id FROM traceroute_routes + WHERE timestamp >= ? AND timestamp <= ? AND parse_status != 'pending' + UNION + SELECT to_node_id FROM traceroute_routes + WHERE timestamp >= ? AND timestamp <= ? AND parse_status != 'pending' + UNION + SELECT CAST(value AS INTEGER) FROM traceroute_routes, json_each(route_nodes_json) + WHERE timestamp >= ? AND timestamp <= ? AND parse_status != 'pending' + UNION + SELECT CAST(value AS INTEGER) FROM traceroute_routes, json_each(route_back_json) + WHERE timestamp >= ? AND timestamp <= ? AND parse_status != 'pending' + """, + ( + window_start.timestamp(), window_end.timestamp(), + window_start.timestamp(), window_end.timestamp(), + window_start.timestamp(), window_end.timestamp(), + window_start.timestamp(), window_end.timestamp(), + window_start.timestamp(), window_end.timestamp(), + window_start.timestamp(), window_end.timestamp(), + ), + ) + if row[0] is not None and row[0] != BROADCAST_NODE_ID + } _hop_nodes_cache = (now, frozenset(hop_node_ids)) @@ -1211,118 +1258,99 @@ def api_traceroute_link(node1_id, node2_id): node1_id_int = convert_node_id(node1_id) node2_id_int = convert_node_id(node2_id) - # Get ALL recent traceroute packets to search for RF hops between these nodes from datetime import datetime, timedelta end_time = datetime.now() - start_time = end_time - timedelta(days=7) # Look at last 7 days - - filters = { - "start_time": start_time.timestamp(), - "end_time": end_time.timestamp(), - "processed_successfully_only": True, - } - - all_packets = TracerouteRepository.get_traceroute_packets( - limit=15000, filters=filters + start_time = end_time - timedelta(days=7) + limit = _clamp_limit(default=100, maximum=1000) + page = _clamp_page() + offset = (page - 1) * limit + link_result = get_traceroute_link( + node1_id_int, + node2_id_int, + start_time=start_time.timestamp(), + end_time=end_time.timestamp(), + limit=limit, + offset=offset, ) - # Don't convert bytes to base64 yet - TraceroutePacket needs raw bytes - # We'll convert only at the end for JSON serialization + route_data_by_id = {} + node_ids = {node1_id_int, node2_id_int} + gateway_ids: dict[int, int] = {} + for packet in link_result["packets"]: + route_data = route_data_from_row(packet) + if route_data is None: + continue + route_data_by_id[packet["id"]] = route_data + node_ids.update(route_data["route_nodes"]) + node_ids.update(route_data["route_back"]) + node_ids.update( + node_id + for node_id in (packet["from_node_id"], packet["to_node_id"]) + if node_id is not None + ) + gateway_id = packet.get("gateway_id") + try: + parsed_gateway_id = ( + int(gateway_id[1:], 16) + if isinstance(gateway_id, str) and gateway_id.startswith("!") + else int(gateway_id) + ) + except (TypeError, ValueError): + continue + gateway_ids[packet["id"]] = parsed_gateway_id + node_ids.add(parsed_gateway_id) + + node_names = NodeRepository.get_bulk_node_names(list(node_ids)) - # Get node names - node_names = NodeRepository.get_bulk_node_names([node1_id_int, node2_id_int]) + def display_name(node_id: int) -> str: + return node_names.get(node_id, f"!{node_id:08x}") - # Process each packet to find RF hops between our target nodes processed_traceroutes: list[dict[str, Any]] = [] - direction_counts: dict[str, int] = {} - snr_values: list[float] = [] - - for packet in all_packets["packets"]: + for packet in link_result["packets"]: try: - # Create TraceroutePacket for analysis - tr_packet = TraceroutePacket(packet, resolve_names=True) - - # Get RF hops (no need to calculate distances for this analysis) - rf_hops = tr_packet.get_rf_hops() - - # Find any RF hop between our two target nodes - target_hop = None - for hop in rf_hops: - if ( - hop.from_node_id == node1_id_int - and hop.to_node_id == node2_id_int - ) or ( - hop.from_node_id == node2_id_int - and hop.to_node_id == node1_id_int - ): - target_hop = hop - break - - if target_hop: - # Determine direction - if target_hop.from_node_id == node1_id_int: - direction = f"{node_names.get(node1_id_int, f'!{node1_id_int:08x}')} → {node_names.get(node2_id_int, f'!{node2_id_int:08x}')}" - else: - direction = f"{node_names.get(node2_id_int, f'!{node2_id_int:08x}')} → {node_names.get(node1_id_int, f'!{node1_id_int:08x}')}" - - direction_counts[direction] = direction_counts.get(direction, 0) + 1 - - if is_plausible_traceroute_snr(target_hop.snr): - snr_values.append(target_hop.snr) - - # Create route_hops structure for UI - include ALL RF hops (forward and return) - route_hops = [] - all_rf_hops = tr_packet.get_rf_hops() - - for i, hop in enumerate(all_rf_hops): - route_hops.append( - { - "hop_number": i + 1, - "from_node_id": hop.from_node_id, - "to_node_id": hop.to_node_id, - "from_node_name": hop.from_node_name, - "to_node_name": hop.to_node_name, - # Null garbage payload SNR so the per-hop rows - # can't display corrupt values as real signal. - "snr": hop.snr - if is_plausible_traceroute_snr(hop.snr) - else None, - "direction": hop.direction, # Include direction info (forward_rf, return_rf) - "is_target_hop": ( - ( - hop.from_node_id == node1_id_int - and hop.to_node_id == node2_id_int - ) - or ( - hop.from_node_id == node2_id_int - and hop.to_node_id == node1_id_int - ) - ), - } - ) - - # Get gateway node name if available - gateway_node_name = None - if tr_packet.gateway_id: - try: - # Convert gateway_id to int if it's a hex string - if isinstance( - tr_packet.gateway_id, str - ) and tr_packet.gateway_id.startswith("!"): - gateway_id_int = int(tr_packet.gateway_id[1:], 16) - else: - gateway_id_int = int(tr_packet.gateway_id) - - gateway_names = NodeRepository.get_bulk_node_names( - [gateway_id_int] - ) - gateway_node_name = gateway_names.get(gateway_id_int) - except (ValueError, TypeError): - pass - - # Create traceroute entry for UI - traceroute_entry = { + route_data = route_data_by_id.get(packet["id"]) + if route_data is None: + continue + tr_packet = TraceroutePacket( + packet, resolve_names=False, pre_parsed_route_data=route_data + ) + tr_packet.from_node_name = display_name(packet["from_node_id"]) + tr_packet.to_node_name = display_name(packet["to_node_id"]) + for path in ( + tr_packet.forward_path, + tr_packet.return_path, + tr_packet.actual_rf_path, + ): + if path is None: + continue + path.node_names = [display_name(node_id) for node_id in path.node_ids] + for hop in path.hops: + hop.from_node_name = display_name(hop.from_node_id) + hop.to_node_name = display_name(hop.to_node_id) + + route_hops = [ + { + "hop_number": i + 1, + "from_node_id": hop.from_node_id, + "to_node_id": hop.to_node_id, + "from_node_name": hop.from_node_name, + "to_node_name": hop.to_node_name, + "snr": hop.snr + if is_plausible_traceroute_snr(hop.snr) + else None, + "direction": hop.direction, + "is_target_hop": ( + {hop.from_node_id, hop.to_node_id} + == {node1_id_int, node2_id_int} + ), + } + for i, hop in enumerate(tr_packet.get_rf_hops()) + ] + gateway_node_id = gateway_ids.get(packet["id"]) + target_snr = packet["target_hop_snr"] + processed_traceroutes.append( + { "id": packet["id"], "timestamp": packet["timestamp"], "timestamp_str": packet["timestamp_str"], @@ -1330,51 +1358,50 @@ def api_traceroute_link(node1_id, node2_id): "to_node_id": packet["to_node_id"], "from_node_name": tr_packet.from_node_name, "to_node_name": tr_packet.to_node_name, - "gateway_id": tr_packet.gateway_id, - "gateway_node_name": gateway_node_name, - # Null out garbage SNR so the SNR-over-time chart - # (which autoscales over hop_snr) can't be flattened. - "hop_snr": target_hop.snr - if is_plausible_traceroute_snr(target_hop.snr) + "gateway_id": packet.get("gateway_id"), + "gateway_node_name": ( + node_names.get(gateway_node_id) + if gateway_node_id is not None + else None + ), + "hop_snr": target_snr + if is_plausible_traceroute_snr(target_snr) else None, "route_hops": route_hops, - "complete_path_display": tr_packet.format_path_display( - "display" - ), + "complete_path_display": tr_packet.format_path_display("display"), } - - processed_traceroutes.append(traceroute_entry) - + ) except Exception as e: logger.warning( f"Error processing traceroute packet {packet['id']}: {e}" ) continue - # Sort by timestamp (most recent first) - processed_traceroutes.sort(key=lambda x: x["timestamp"], reverse=True) - - # Calculate summary statistics - total_attempts = len(processed_traceroutes) - avg_snr = sum(snr_values) / len(snr_values) if snr_values else None - - # Ensure direction_counts has the expected format even when empty - if not direction_counts: - # Create default direction labels for the two nodes - node_names.get(node1_id_int, f"!{node1_id_int:08x}") - node_names.get(node2_id_int, f"!{node2_id_int:08x}") + if link_result["total_attempts"]: + direction_counts = { + f"{display_name(node1_id_int)} → {display_name(node2_id_int)}": link_result[ + "forward_count" + ], + f"{display_name(node2_id_int)} → {display_name(node1_id_int)}": link_result[ + "reverse_count" + ], + } + else: direction_counts = {"forward": 0, "reverse": 0} - # Create response in the format expected by the UI response_data = { "from_node_id": node1_id_int, "to_node_id": node2_id_int, - "from_node_name": node_names.get(node1_id_int, f"!{node1_id_int:08x}"), - "to_node_name": node_names.get(node2_id_int, f"!{node2_id_int:08x}"), - "total_attempts": total_attempts, - "avg_snr": avg_snr, + "from_node_name": display_name(node1_id_int), + "to_node_name": display_name(node2_id_int), + "total_attempts": link_result["total_attempts"], + "avg_snr": link_result["avg_snr"], "direction_counts": direction_counts, "traceroutes": processed_traceroutes, + "page": page, + "limit": limit, + "total_count": link_result["total_count"], + "total_pages": (link_result["total_count"] + limit - 1) // limit, } # Convert any remaining bytes to base64 for JSON serialization @@ -1917,15 +1944,10 @@ def api_traceroute_data(): gateway_node_ids.add(gateway_node_id) except ValueError: pass - if tr.get("raw_payload"): - try: - route_data = parse_traceroute_payload(tr["raw_payload"]) - if route_data.get("route_nodes"): - for route_node_id in route_data["route_nodes"]: - node_ids.add(route_node_id) - except Exception: - # If parsing fails, we'll handle it in the individual processing below - pass + route_data = route_data_from_row(tr) + if route_data and route_data.get("route_nodes"): + for route_node_id in route_data["route_nodes"]: + node_ids.add(route_node_id) node_names = get_bulk_node_names(list(node_ids | gateway_node_ids)) @@ -1976,20 +1998,13 @@ def api_traceroute_data(): route_nodes = [] route_names = [] - # If no route data from repository, try parsing raw_payload - if not route_nodes and tr.get("raw_payload"): - try: - route_data = parse_traceroute_payload(tr["raw_payload"]) - if route_data.get("route_nodes"): - route_nodes = route_data["route_nodes"] - # Get names for each node in the route - for node_id in route_nodes: - node_name = node_short_names.get( - node_id, f"!{node_id:08x}"[-4:] - ) - route_names.append(node_name) - except Exception: - pass + if not route_nodes: + route_data = route_data_from_row(tr) + if route_data and route_data.get("route_nodes"): + route_nodes = route_data["route_nodes"] + for node_id in route_nodes: + node_name = node_short_names.get(node_id, f"!{node_id:08x}"[-4:]) + route_names.append(node_name) # Final fallback: use from -> to if not route_nodes: diff --git a/src/malla/routes/packet_routes.py b/src/malla/routes/packet_routes.py index 56f00dc3..88dc3f82 100644 --- a/src/malla/routes/packet_routes.py +++ b/src/malla/routes/packet_routes.py @@ -79,17 +79,20 @@ def get_packet_details(packet_id: int) -> dict[str, Any] | None: has_envelope_col = PacketRepository.has_raw_service_envelope_column(cursor) # Get the main packet information - env_col = ", raw_service_envelope" if has_envelope_col else "" + env_col = ", p.raw_service_envelope" if has_envelope_col else "" cursor.execute( f""" SELECT - id, timestamp, from_node_id, to_node_id, portnum, portnum_name, - gateway_id, channel_id, mesh_packet_id, rssi, snr, hop_limit, hop_start, - payload_length, processed_successfully, raw_payload, - via_mqtt, want_ack, priority, delayed, channel_index, rx_time, - pki_encrypted, next_hop, relay_node, tx_after{env_col} - FROM packet_history - WHERE id = ? + p.id, p.timestamp, p.from_node_id, p.to_node_id, p.portnum, p.portnum_name, + p.gateway_id, p.channel_id, p.mesh_packet_id, p.rssi, p.snr, p.hop_limit, p.hop_start, + p.payload_length, p.processed_successfully, p.raw_payload, + p.via_mqtt, p.want_ack, p.priority, p.delayed, p.channel_index, p.rx_time, + p.pki_encrypted, p.next_hop, p.relay_node, p.tx_after{env_col}, + r.route_nodes_json, r.snr_towards_json, r.route_back_json, r.snr_back_json, + r.parse_status, r.parse_error + FROM packet_history p + LEFT JOIN traceroute_routes r ON r.packet_id = p.id + WHERE p.id = ? """, (packet_id,), ) @@ -167,14 +170,17 @@ def get_packet_details(packet_id: int) -> dict[str, Any] | None: cursor.execute( """ SELECT - id, timestamp, gateway_id, channel_id, rssi, snr, hop_limit, hop_start, - payload_length, processed_successfully, - raw_payload, from_node_id, to_node_id, portnum, portnum_name, relay_node, - mesh_packet_id - FROM packet_history - WHERE mesh_packet_id = ? - AND id != ? - ORDER BY timestamp ASC + p.id, p.timestamp, p.gateway_id, p.channel_id, p.rssi, p.snr, p.hop_limit, p.hop_start, + p.payload_length, p.processed_successfully, + p.raw_payload, p.from_node_id, p.to_node_id, p.portnum, p.portnum_name, p.relay_node, + p.mesh_packet_id, + r.route_nodes_json, r.snr_towards_json, r.route_back_json, r.snr_back_json, + r.parse_status, r.parse_error + FROM packet_history p + LEFT JOIN traceroute_routes r ON r.packet_id = p.id + WHERE p.mesh_packet_id = ? + AND p.id != ? + ORDER BY p.timestamp ASC """, (packet["mesh_packet_id"], packet_id), ) @@ -188,16 +194,19 @@ def get_packet_details(packet_id: int) -> dict[str, Any] | None: cursor.execute( """ SELECT - id, timestamp, gateway_id, channel_id, rssi, snr, hop_limit, hop_start, - payload_length, processed_successfully, - raw_payload, from_node_id, to_node_id, portnum, portnum_name, relay_node, - mesh_packet_id - FROM packet_history - WHERE from_node_id = ? - AND timestamp BETWEEN ? AND ? - AND portnum = ? - AND id != ? - ORDER BY timestamp ASC + p.id, p.timestamp, p.gateway_id, p.channel_id, p.rssi, p.snr, p.hop_limit, p.hop_start, + p.payload_length, p.processed_successfully, + p.raw_payload, p.from_node_id, p.to_node_id, p.portnum, p.portnum_name, p.relay_node, + p.mesh_packet_id, + r.route_nodes_json, r.snr_towards_json, r.route_back_json, r.snr_back_json, + r.parse_status, r.parse_error + FROM packet_history p + LEFT JOIN traceroute_routes r ON r.packet_id = p.id + WHERE p.from_node_id = ? + AND p.timestamp BETWEEN ? AND ? + AND p.portnum = ? + AND p.id != ? + ORDER BY p.timestamp ASC """, ( packet["from_node_id"], @@ -521,6 +530,80 @@ def decode_packet_payload(packet: dict[str, Any]) -> dict[str, Any] | None: "error": None, } + # For TRACEROUTE_APP: use pre-parsed route data from traceroute_routes if available + if packet.get("portnum_name") == "TRACEROUTE_APP": + parse_status = packet.get("parse_status") + if parse_status == "invalid_payload": + payload_info["decoded"] = False + payload_info["error"] = packet.get("parse_error") or "Invalid traceroute payload" + return payload_info + + if packet.get("route_nodes_json") is not None: + try: + tr_packet = TraceroutePacket(packet, resolve_names=True) + tr_packet.calculate_hop_distances(calculate_for_all_paths=True) + forward_hops_with_distances = ( + tr_packet.get_display_hops_with_distances() + ) + return_hops_with_distances = ( + tr_packet.get_return_hops_with_distances() + ) + + payload_info["decoded"] = True + payload_info["data"] = { + "route_nodes": tr_packet.route_data["route_nodes"], + "snr_towards": tr_packet.route_data["snr_towards"], + "route_back": tr_packet.route_data["route_back"], + "snr_back": tr_packet.route_data["snr_back"], + "route_node_names": {}, + "traceroute_packet": tr_packet, + "has_return_path": tr_packet.has_return_path(), + "is_complete": tr_packet.is_complete(), + "forward_path_display": tr_packet.format_path_display("display"), + "return_path_display": tr_packet.format_path_display("return") + if tr_packet.has_return_path() + else None, + "actual_rf_path_display": tr_packet.format_path_display("actual_rf"), + "forward_hops": forward_hops_with_distances, + "return_hops": return_hops_with_distances, + "total_forward_distance": sum( + hop.distance_meters + for hop in forward_hops_with_distances + if hop.distance_meters is not None + ) + if forward_hops_with_distances + else None, + "total_return_distance": sum( + hop.distance_meters + for hop in return_hops_with_distances + if hop.distance_meters is not None + ) + if return_hops_with_distances + else None, + "parse_status": parse_status, + "parse_error": packet.get("parse_error"), + } + + all_route_nodes = set() + if tr_packet.route_data["route_nodes"]: + all_route_nodes.update(tr_packet.route_data["route_nodes"]) + if tr_packet.route_data["route_back"]: + all_route_nodes.update(tr_packet.route_data["route_back"]) + if packet.get("from_node_id"): + all_route_nodes.add(packet["from_node_id"]) + if packet.get("to_node_id"): + all_route_nodes.add(packet["to_node_id"]) + + if all_route_nodes: + route_node_names = get_bulk_node_names(list(all_route_nodes)) + payload_info["data"]["route_node_names"] = route_node_names + + return payload_info + except Exception as e: + logger.warning( + f"Failed to decode traceroute using materialized data for packet {packet.get('id')}: {e}" + ) + # Use the new generic protobuf decoding system decoded_payload = decode_protobuf_payload(packet) diff --git a/src/malla/services/location_service.py b/src/malla/services/location_service.py index 82f025aa..408e7531 100644 --- a/src/malla/services/location_service.py +++ b/src/malla/services/location_service.py @@ -186,7 +186,10 @@ def get_node_locations( } ) - # Get direct packet links to include in neighbor data + # Track latest direct packet reception timestamp per node + packet_last_seen: dict[int, float] = {} + + # Process packet links to add to neighbor details try: # Use pre-computed packet links if provided, otherwise fetch them if packet_links is None: @@ -206,6 +209,12 @@ def get_node_locations( from_node_id = link["from_node_id"] to_node_id = link["to_node_id"] packet_count = link.get("total_hops_seen", 0) + link_ts = link.get("last_seen") + if link_ts: + if from_node_id not in packet_last_seen or link_ts > packet_last_seen[from_node_id]: + packet_last_seen[from_node_id] = link_ts + if to_node_id not in packet_last_seen or link_ts > packet_last_seen[to_node_id]: + packet_last_seen[to_node_id] = link_ts # Initialize neighbor tracking if not already present if from_node_id not in neighbor_counts: @@ -280,18 +289,33 @@ def get_node_locations( for location in locations: node_id = location["node_id"] - # Calculate age in hours - age_hours = (current_time - location["timestamp"]) / 3600 - - # Format timestamp string - timestamp_dt = datetime.fromtimestamp(location["timestamp"], UTC) - timestamp_str = timestamp_dt.strftime("%Y-%m-%d %H:%M:%S UTC") - # Get network data for this node network_node = network_nodes.get(node_id, {}) direct_neighbors = neighbor_counts.get(node_id, 0) neighbors = neighbor_details.get(node_id, []) + network_last_seen = network_node.get("last_seen") + pkt_last_seen = packet_last_seen.get(node_id) + + # Unified active timestamp: node is active if it broadcast position, + # participated in a traceroute RF hop, or exchanged direct packets. + active_candidates = [location["timestamp"]] + if network_last_seen: + active_candidates.append(network_last_seen) + if pkt_last_seen: + active_candidates.append(pkt_last_seen) + active_timestamp = max(active_candidates) + + # Calculate age in hours relative to the latest active timestamp + age_hours = (current_time - active_timestamp) / 3600 + + # Format timestamp strings + timestamp_dt = datetime.fromtimestamp(active_timestamp, UTC) + timestamp_str = timestamp_dt.strftime("%Y-%m-%d %H:%M:%S UTC") + + pos_dt = datetime.fromtimestamp(location["timestamp"], UTC) + pos_str = pos_dt.strftime("%Y-%m-%d %H:%M:%S UTC") + enhanced_location = { # Original location data "node_id": location["node_id"], @@ -305,7 +329,9 @@ def get_node_locations( "latitude": location["latitude"], "longitude": location["longitude"], "altitude": location["altitude"], - "timestamp": location["timestamp"], + "timestamp": active_timestamp, + "position_timestamp": location["timestamp"], + "position_timestamp_str": pos_str, # Enhanced fields for map display "age_hours": round(age_hours, 2), "timestamp_str": timestamp_str, @@ -317,7 +343,8 @@ def get_node_locations( # Network analysis data "packet_count": network_node.get("packet_count", 0), "avg_snr": network_node.get("avg_snr"), - "last_seen_network": network_node.get("last_seen"), + "last_seen_network": network_last_seen, + "last_seen_packet": pkt_last_seen, } enhanced_locations.append(enhanced_location) diff --git a/src/malla/services/node_service.py b/src/malla/services/node_service.py index ee9385e7..9b3b317c 100644 --- a/src/malla/services/node_service.py +++ b/src/malla/services/node_service.py @@ -167,7 +167,6 @@ def get_traceroute_related_nodes(node_id) -> dict[str, Any]: Dictionary containing related nodes and their RF hop counts """ from ..database import get_db_connection - from ..models.traceroute import TraceroutePacket node_id_int = convert_node_id(node_id) @@ -178,78 +177,26 @@ def get_traceroute_related_nodes(node_id) -> dict[str, Any]: end_time = datetime.now() start_time = end_time - timedelta(days=7) # Look at last 7 days - query = """ + rows = cursor.execute( + """ SELECT - id, - timestamp, - from_node_id, - to_node_id, - gateway_id, - hop_start, - hop_limit, - raw_payload - FROM packet_history - WHERE portnum_name = 'TRACEROUTE_APP' - AND processed_successfully = 1 - AND raw_payload IS NOT NULL - AND timestamp >= ? - AND timestamp <= ? - """ - - cursor.execute(query, (start_time.timestamp(), end_time.timestamp())) - packets = cursor.fetchall() - - # Track nodes with direct RF hops and their connection counts - related_nodes = {} - - for packet in packets: + CASE WHEN from_node_id = ? THEN to_node_id ELSE from_node_id END AS other_node, + COUNT(*) AS observation_count + FROM traceroute_hops + WHERE timestamp >= ? AND timestamp <= ? + AND (from_node_id = ? OR to_node_id = ?) + GROUP BY other_node + ORDER BY observation_count DESC, other_node + """, ( - packet_id, - timestamp, - from_node_id, - to_node_id, - gateway_id, - hop_start, - hop_limit, - raw_payload, - ) = packet - - try: - # Create TraceroutePacket to analyze RF hops - packet_data = { - "id": packet_id, - "timestamp": timestamp, - "from_node_id": from_node_id, - "to_node_id": to_node_id, - "gateway_id": gateway_id, - "hop_start": hop_start, - "hop_limit": hop_limit, - "raw_payload": raw_payload, - } - - tr_packet = TraceroutePacket(packet_data, resolve_names=False) - - # Get all RF hops (both forward and return) - rf_hops = tr_packet.get_rf_hops() - - # Check if any RF hop involves our target node - for hop in rf_hops: - if hop.from_node_id == node_id_int: - # Target node is the sender in this RF hop - other_node = hop.to_node_id - if other_node not in related_nodes: - related_nodes[other_node] = 0 - related_nodes[other_node] += 1 - elif hop.to_node_id == node_id_int: - # Target node is the receiver in this RF hop - other_node = hop.from_node_id - if other_node not in related_nodes: - related_nodes[other_node] = 0 - related_nodes[other_node] += 1 - - except Exception as e: - logger.warning(f"Failed to analyze RF hops for packet {packet_id}: {e}") - continue + node_id_int, + start_time.timestamp(), + end_time.timestamp(), + node_id_int, + node_id_int, + ), + ).fetchall() + related_nodes = {row["other_node"]: row["observation_count"] for row in rows} # Get node info for all related nodes if related_nodes: diff --git a/src/malla/services/traceroute_service.py b/src/malla/services/traceroute_service.py index 4756d959..7b0d7af0 100644 --- a/src/malla/services/traceroute_service.py +++ b/src/malla/services/traceroute_service.py @@ -12,19 +12,25 @@ import math import time from datetime import datetime, timedelta -from typing import Any, cast +from typing import Any from ..database.repositories import ( LocationRepository, TracerouteRepository, ) +from ..database.traceroute_read_repository import ( + get_node_traceroute_statistics, + get_route_patterns_data, + get_traceroute_hops_for_graph, + get_traceroute_hops_for_longest_links, + route_data_from_row, +) from ..models.traceroute import ( - RouteData, TraceroutePacket, # Use the correct TraceroutePacket class ) +from ..utils.geo_utils import calculate_distance from ..utils.node_utils import get_bulk_node_names from ..utils.signal_quality import is_plausible_traceroute_snr -from ..utils.traceroute_utils import parse_traceroute_payload logger = logging.getLogger(__name__) @@ -32,13 +38,6 @@ _NETWORK_GRAPH_CACHE_TTL_SECONDS = 60 _NETWORK_GRAPH_CACHE_MAX_ENTRIES = 32 -# Default cap on traceroute packets analyzed per network graph build. The graph -# advertises a multi-day window (e.g. 168h for the map), so this must be large -# enough that the newest N packets still span that window; on busy brokers a -# low cap silently shrinks the analysis to the last few hours and older RF -# links drop off the map. ~20k packets parse in a couple of seconds. -DEFAULT_GRAPH_PACKET_LIMIT = 20000 - def _network_graph_cache_key( hours: int, @@ -89,6 +88,47 @@ def _prune_network_graph_cache(now: float) -> None: _NETWORK_GRAPH_CACHE.pop(key, None) +def _rf_hop_qualifies(hop: dict[str, Any], min_snr: float | None = None) -> bool: + """True when a hop is an evidenced RF link (usable SNR, real endpoints).""" + snr = hop.get("snr") + if not is_plausible_traceroute_snr(snr) or snr == 0: + return False + if min_snr is not None and min_snr != -200 and snr < min_snr: + return False + return 4294967295 not in (hop["from_node_id"], hop["to_node_id"]) + + +def _contiguous_path_segments( + hops: list[dict[str, Any]], + qualifies: Any = _rf_hop_qualifies, +) -> list[list[dict[str, Any]]]: + """Split hops in path order into maximal contiguous runs of qualifying hops. + + A run continues only while each hop starts where the previous one ended. + Removing a hop from the middle of a route (zero/invalid SNR, weak link) + must not splice the survivors into a shorter path: for A->B->C->D with + B->C filtered, A->B and C->D stay two disconnected segments instead of + aggregating into a bogus A->D path. + """ + segments: list[list[dict[str, Any]]] = [] + current: list[dict[str, Any]] = [] + for hop in hops: + if not qualifies(hop): + if len(current) > 1: + segments.append(current) + current = [] + continue + if current and hop["from_node_id"] != current[-1]["to_node_id"]: + if len(current) > 1: + segments.append(current) + current = [hop] + else: + current.append(hop) + if len(current) > 1: + segments.append(current) + return segments + + class TracerouteService: """Service for traceroute analysis and management.""" @@ -192,9 +232,10 @@ def get_traceroute_analysis(hours: int = 24) -> dict[str, Any]: "end_time": end_time.timestamp(), } - # Get raw traceroute data + # Saved JSON arrays make full-window analysis practical without + # decoding packet payloads during the request. result = TracerouteRepository.get_traceroute_packets( - limit=1000, # Large limit for analysis + limit=-1, filters=filters, ) @@ -210,9 +251,8 @@ def get_traceroute_analysis(hours: int = 24) -> dict[str, Any]: if tr["processed_successfully"]: successful_traceroutes += 1 - # Parse route data - if tr["raw_payload"]: - route_data = parse_traceroute_payload(tr["raw_payload"]) + route_data = route_data_from_row(tr) + if route_data is not None: if route_data["route_back"]: traceroutes_with_return += 1 @@ -287,75 +327,41 @@ def get_traceroute_analysis(hours: int = 24) -> dict[str, Any]: raise @staticmethod - def get_route_patterns(limit: int = 50) -> dict[str, Any]: + def get_route_patterns( + limit: int = 50, + hours: int = 168, + filters: dict | None = None, + ) -> dict[str, Any]: """ - Analyze common route patterns in the mesh network. + Analyze common route patterns in the mesh network using materialized routes. Args: limit: Maximum number of patterns to return + hours: Hours window for analysis (default 168h / 7 days) + filters: Optional filters with explicit start_time and end_time Returns: Dictionary with route pattern analysis """ - logger.info(f"Getting route patterns (limit={limit})") + logger.info(f"Getting route patterns (limit={limit}, hours={hours})") try: - # Get recent successful traceroutes - filters = {"processed_successfully_only": True} - result = TracerouteRepository.get_traceroute_packets( - limit=1000, # Analyze more data - filters=filters, + now = datetime.now() + start_time = (now - timedelta(hours=hours)).timestamp() + end_time = now.timestamp() + if filters: + if filters.get("start_time"): + start_time = float(filters["start_time"]) + if filters.get("end_time"): + end_time = float(filters["end_time"]) + + raw_result = get_route_patterns_data( + start_time=start_time, + end_time=end_time, + limit=limit, ) - # Analyze patterns - route_patterns: dict[ - tuple[tuple[int, int], tuple[int, ...]], dict[str, Any] - ] = {} - directional_patterns: dict[tuple[int, int, tuple[int, ...]], int] = {} - - for tr in result["packets"]: - if tr["raw_payload"] and tr["processed_successfully"]: - route_data = parse_traceroute_payload(tr["raw_payload"]) - - # Create pattern key (normalized) - route_nodes = tuple(route_data["route_nodes"]) - if route_nodes: - # Bidirectional pattern (normalized by sorting endpoints) - endpoints = tuple( - sorted([tr["from_node_id"], tr["to_node_id"]]) - ) - pattern_key = (endpoints, route_nodes) - - if pattern_key not in route_patterns: - route_patterns[pattern_key] = { - "count": 0, - "endpoints": endpoints, - "route_nodes": route_nodes, - "avg_success_rate": 0, - "examples": [], - } - - route_patterns[pattern_key]["count"] += 1 - if len(route_patterns[pattern_key]["examples"]) < 3: - route_patterns[pattern_key]["examples"].append( - { - "packet_id": tr["id"], - "timestamp": tr["timestamp"], - "from_node": tr["from_node_id"], - "to_node": tr["to_node_id"], - } - ) - - # Directional pattern - dir_key = (tr["from_node_id"], tr["to_node_id"], route_nodes) - directional_patterns[dir_key] = ( - directional_patterns.get(dir_key, 0) + 1 - ) - - # Sort patterns by frequency - sorted_patterns = sorted( - route_patterns.items(), key=lambda x: x[1]["count"], reverse=True - )[:limit] + sorted_patterns = raw_result["sorted_patterns"] # Enhance with node names all_node_ids: set[int] = set() @@ -380,8 +386,9 @@ def get_route_patterns(limit: int = 50) -> dict[str, Any]: return { "patterns": enhanced_patterns, - "total_patterns": len(route_patterns), - "analyzed_traceroutes": len(result["packets"]), + "total_patterns": raw_result["total_patterns"], + "analyzed_traceroutes": raw_result["analyzed_traceroutes"], + "time_period_hours": hours, } except Exception as e: @@ -389,12 +396,18 @@ def get_route_patterns(limit: int = 50) -> dict[str, Any]: raise @staticmethod - def get_node_traceroute_stats(node_id: int) -> dict[str, Any]: + def get_node_traceroute_stats( + node_id: int, + start_time: float | None = None, + end_time: float | None = None, + ) -> dict[str, Any]: """ Get traceroute statistics for a specific node. Args: node_id: Node ID to analyze + start_time: Optional start timestamp filter + end_time: Optional end timestamp filter Returns: Dictionary with node's traceroute statistics @@ -402,67 +415,15 @@ def get_node_traceroute_stats(node_id: int) -> dict[str, Any]: logger.info(f"Getting traceroute stats for node {node_id}") try: - # Get traceroutes involving this node as source or destination - source_filters = {"from_node": node_id} - dest_filters = {"to_node": node_id} - - source_result = TracerouteRepository.get_traceroute_packets( - limit=1000, filters=source_filters - ) - dest_result = TracerouteRepository.get_traceroute_packets( - limit=1000, filters=dest_filters - ) - - # Analyze as source - source_total = len(source_result["packets"]) - source_successful = sum( - 1 for tr in source_result["packets"] if tr["processed_successfully"] - ) - - # Analyze as destination - dest_total = len(dest_result["packets"]) - dest_successful = sum( - 1 for tr in dest_result["packets"] if tr["processed_successfully"] + stats = get_node_traceroute_statistics( + node_id=node_id, + start_time=start_time, + end_time=end_time, ) - - # Get node name node_names = get_bulk_node_names([node_id]) node_name = node_names.get(node_id, f"!{node_id:08x}") - - # Analyze route participation (as intermediate hop) - # This requires checking all traceroutes for this node in route_nodes - participation_count = 0 - all_traceroutes = TracerouteRepository.get_traceroute_packets( - limit=1000, filters={"processed_successfully_only": True} - ) - - for tr in all_traceroutes["packets"]: - if tr["raw_payload"]: - route_data = parse_traceroute_payload(tr["raw_payload"]) - if node_id in route_data.get("route_nodes", []): - participation_count += 1 - - return { - "node_id": node_id, - "node_name": node_name, - "as_source": { - "total": source_total, - "successful": source_successful, - "success_rate": (source_successful / source_total * 100) - if source_total > 0 - else 0, - }, - "as_destination": { - "total": dest_total, - "successful": dest_successful, - "success_rate": (dest_successful / dest_total * 100) - if dest_total > 0 - else 0, - }, - "as_intermediate_hop": {"participation_count": participation_count}, - "total_involvement": source_total + dest_total + participation_count, - } - + stats["node_name"] = node_name + return stats except Exception as e: logger.error(f"Error getting node traceroute stats: {e}") raise @@ -490,96 +451,49 @@ def get_longest_links_analysis( try: # ------------------------------------------------------------------ - # Fetch raw data (only the last 7 days & successfully processed) + # Fetch RF hops for the last 7 days from materialized traceroute_hops # ------------------------------------------------------------------ fetch_start = time.time() - from datetime import datetime, timedelta - end_time = datetime.now() start_time_filter = end_time - timedelta(days=7) - filters = { - "start_time": start_time_filter.timestamp(), - "end_time": end_time.timestamp(), - "processed_successfully_only": True, - } - - result = TracerouteRepository.get_traceroute_packets( - # Fetch a larger sample of packets to cover busy networks - # 25k packets ≈ several hours of traffic on busy meshes but still manageable - limit=25000, - filters=filters, + hops = get_traceroute_hops_for_longest_links( + start_time=start_time_filter.timestamp(), + end_time=end_time.timestamp(), ) fetch_duration = time.time() - fetch_start logger.info( - f"TIMING: Data fetch took {fetch_duration:.3f}s for {len(result['packets'])} packets" + f"TIMING: Data fetch took {fetch_duration:.3f}s for {len(hops)} hops" ) # ------------------------------------------------------------------ - # Pre-fetch node location history using a single query per node - # (replaces the previous expensive nested node->packet cache fill). + # Batch fetch node location history and names # ------------------------------------------------------------------ - # First collect all unique node ids that appear in the packets unique_node_ids: set[int] = set() - parsed_route_cache: dict[int, RouteData] = {} - for packet in result["packets"]: - if not packet.get("raw_payload"): - continue - try: - route_data = parse_traceroute_payload(packet["raw_payload"]) - parsed_route_cache[packet["id"]] = route_data - nodes_for_packet = {packet["from_node_id"], packet["to_node_id"]} - nodes_for_packet.update(route_data.get("route_nodes", [])) - # Remove invalid placeholders - nodes_for_packet.discard(None) - nodes_for_packet.discard(4294967295) - unique_node_ids.update(nodes_for_packet) - except Exception as e: - logger.warning( - f"Error parsing packet {packet.get('id', 'unknown')} for node collection: {e}" - ) - continue + for hop in hops: + for nid in (hop["from_node_id"], hop["to_node_id"]): + if nid and nid != 4294967295: + unique_node_ids.add(nid) prefetch_start = time.time() - - from ..utils import traceroute_utils as _tru # Local import to avoid cycles - - # Build a dict: node_id -> list[location_dict] (DESC by timestamp) - location_history_cache: dict[int, list[dict[str, Any]]] = {} - for node_id in unique_node_ids: - try: - locations = LocationRepository.get_node_location_history( - node_id, limit=50 - ) - if locations: - location_history_cache[node_id] = ( - locations # already DESC order - ) - except Exception as e: - logger.warning( - f"Error fetching location history for node {node_id}: {e}" - ) - continue - + node_ids_list = list(unique_node_ids) + location_history_cache = ( + LocationRepository.get_nodes_location_history( + node_ids_list, limit_per_node=50 + ) + if node_ids_list + else {} + ) + node_names = get_bulk_node_names(node_ids_list) if node_ids_list else {} prefetch_duration = time.time() - prefetch_start logger.info( - f"TIMING: Location history pre-fetch took {prefetch_duration:.3f}s for {len(location_history_cache)} nodes" + f"TIMING: Batch pre-fetch took {prefetch_duration:.3f}s for {len(node_ids_list)} nodes" ) - # ------------------------------------------------------------------ - # Inject a fast in-memory location lookup to avoid per-hop DB hits. - # ------------------------------------------------------------------ - _orig_get_location = _tru.get_node_location_at_timestamp # Backup original - - def _fast_location_lookup( - node_id: int, target_ts: float - ) -> dict[str, Any] | None: - """Return the best location for a node at target_ts from the pre-loaded history. + # Fast in-memory location lookup + location_cache: dict[tuple[int, int], dict[str, Any] | None] = {} - Falls back to the original DB implementation if the node isn't in cache - (keeps behaviour identical for very old/unknown nodes). - """ - # Use an hour bucket to maximise cache hits without losing much accuracy + def get_node_loc(node_id: int, target_ts: float) -> dict[str, Any] | None: bucket = int(target_ts // 3600) memo_key = (node_id, bucket) if memo_key in location_cache: @@ -587,430 +501,248 @@ def _fast_location_lookup( history = location_history_cache.get(node_id) if not history: - loc = _orig_get_location(node_id, target_ts) - if loc: - location_cache[memo_key] = loc - return loc + location_cache[memo_key] = None + return None - # histories are DESC (newest first). Find first <= ts. + # Find first location with timestamp <= target_ts (history is newest first) best = None for loc in history: if loc["timestamp"] <= target_ts: best = loc break if best is None: - # No past location found – use oldest future record. best = history[-1] - age_sec = target_ts - best["timestamp"] - age_hours = abs(age_sec) / 3600 - if age_sec >= 0: # Past record - if age_hours <= 24: - age_warning = f"from {age_hours:.1f}h ago" - elif age_hours <= 168: - age_warning = f"from {age_hours / 24:.1f}d ago" - else: - age_warning = f"from {age_hours / 168:.1f}w ago" - else: # Future record - if age_hours <= 24: - age_warning = f"from {age_hours:.1f}h later" - elif age_hours <= 168: - age_warning = f"from {age_hours / 24:.1f}d later" - else: - age_warning = f"from {age_hours / 168:.1f}w later" - - loc_dict = { - "latitude": best["latitude"], - "longitude": best["longitude"], - "altitude": best.get("altitude"), - "timestamp": best["timestamp"], - "age_warning": age_warning, - } - location_cache[memo_key] = loc_dict - return loc_dict - - # Override location lookup with fast in-memory version for this analysis - _tru.get_node_location_at_timestamp = cast(Any, _fast_location_lookup) - - # Location cache used by TraceroutePacket.calculate_hop_distances - location_cache: dict[tuple, Any] = {} + location_cache[memo_key] = best + return best - try: - # ------------------------------------------------------------------ - # Stream-process each hop to avoid holding large intermediate lists. - # ------------------------------------------------------------------ - process_start = time.time() - link_stats: dict[tuple, dict[str, Any]] = {} - # Track multi-hop path statistics (indirect links) - path_stats: dict[tuple, dict[str, Any]] = {} - - logger.info( - f"Processing {len(result['packets'])} packets with pre-populated location cache" - ) - - packets_processed = 0 - hops_processed = 0 - distance_calculations = 0 - cache_hits = 0 - cache_misses = 0 - early_filtered = 0 - - for packet in result["packets"]: - packet_start = time.time() - try: - # Early filtering: skip packets that won't contribute any valid hops + # ------------------------------------------------------------------ + # Group hops by (packet_id, direction) to evaluate direct and indirect paths + # ------------------------------------------------------------------ + process_start = time.time() + link_stats: dict[tuple[int, int], dict[str, Any]] = {} + path_stats: dict[tuple[int, int], dict[str, Any]] = {} + + # Group hops + hops_by_path: dict[tuple[int, str], list[dict[str, Any]]] = {} + for hop in hops: + path_key = (hop["packet_id"], hop.get("direction", "forward")) + if path_key not in hops_by_path: + hops_by_path[path_key] = [] + hops_by_path[path_key].append(hop) + + for (packet_id, _direction), path_hops in hops_by_path.items(): + for hop in path_hops: + from_id = hop["from_node_id"] + to_id = hop["to_node_id"] + ts = hop["timestamp"] + + dist: float | None = None + if from_id != 4294967295 and to_id != 4294967295: + loc_from = get_node_loc(from_id, ts) + loc_to = get_node_loc(to_id, ts) if ( - not packet["raw_payload"] - or not packet["processed_successfully"] + loc_from + and loc_to + and loc_from.get("latitude") is not None + and loc_from.get("longitude") is not None + and loc_to.get("latitude") is not None + and loc_to.get("longitude") is not None ): - early_filtered += 1 - continue - - tr_packet = TraceroutePacket( - packet_data=packet, - resolve_names=True, - ) - - # Track cache performance before distance calculation - cache_size_before = len(location_cache) - - # Populate distance information – uses pre-populated location_cache - distance_calc_start = time.time() - tr_packet.calculate_hop_distances(location_cache=location_cache) - # Track timing (result not used but calculation is important) - _ = time.time() - distance_calc_start - distance_calculations += 1 - - # Track cache performance after distance calculation - cache_size_after = len(location_cache) - if cache_size_after > cache_size_before: - cache_misses += cache_size_after - cache_size_before - else: - cache_hits += 1 - - rf_hops = tr_packet.get_rf_hops() - hops_processed += len(rf_hops) - - for hop in rf_hops: - # Early filtering: skip hops that won't meet criteria - if ( - not is_plausible_traceroute_snr(hop.snr) - or hop.snr == 0 - or hop.snr < min_snr - or not hop.distance_km - or hop.distance_km < min_distance_km - or 4294967295 in [hop.from_node_id, hop.to_node_id] - ): - continue - - # Use a bidirectional key so A<->B == B<->A - key = tuple(sorted([hop.from_node_id, hop.to_node_id])) - - if key not in link_stats: - # Determine the correct orientation for names - node1_id, node2_id = key - if hop.from_node_id == node1_id: - from_name = hop.from_node_name - to_name = hop.to_node_name - else: - from_name = hop.to_node_name - to_name = hop.from_node_name - - link_stats[key] = { - "from_node_name": from_name, - "to_node_name": to_name, - "total_distance": 0.0, - "total_snr": 0.0, - "traceroute_count": 0, - "max_distance": 0.0, - "best_snr": None, - "recent_packets": [], # keep last 5 ids - } - - stats_dict = link_stats[key] - - # Update aggregates - stats_dict["traceroute_count"] += 1 - stats_dict["total_distance"] += hop.distance_km - stats_dict["total_snr"] += hop.snr - stats_dict["max_distance"] = max( - stats_dict["max_distance"], hop.distance_km + dist = calculate_distance( + loc_from["latitude"], + loc_from["longitude"], + loc_to["latitude"], + loc_to["longitude"], ) - - if ( - stats_dict["best_snr"] is None - or hop.snr > stats_dict["best_snr"] - ): - stats_dict["best_snr"] = hop.snr - - # Maintain only last 5 packet ids (newest first) - stats_dict["recent_packets"].append(packet["id"]) + hop["_distance_km"] = dist + + # Direct link processing + snr = hop["snr"] + if ( + dist is not None + and dist >= min_distance_km + and is_plausible_traceroute_snr(snr) + and snr != 0 + and snr >= min_snr + ): + node1_id, node2_id = sorted((from_id, to_id)) + key: tuple[int, int] = (node1_id, node2_id) + from_name = node_names.get(node1_id, f"!{node1_id:08x}") + to_name = node_names.get(node2_id, f"!{node2_id:08x}") + + if key not in link_stats: + link_stats[key] = { + "from_node_name": from_name, + "to_node_name": to_name, + "total_distance": 0.0, + "total_snr": 0.0, + "traceroute_count": 0, + "max_distance": 0.0, + "best_snr": None, + "recent_packets": [], + "last_seen": ts, + } + stats_dict = link_stats[key] + stats_dict["traceroute_count"] += 1 + stats_dict["total_distance"] += dist + stats_dict["total_snr"] += snr + stats_dict["max_distance"] = max(stats_dict["max_distance"], dist) + if stats_dict["best_snr"] is None or snr > stats_dict["best_snr"]: + stats_dict["best_snr"] = snr + if ts > stats_dict["last_seen"]: + stats_dict["last_seen"] = ts + if packet_id not in stats_dict["recent_packets"]: + stats_dict["recent_packets"].append(packet_id) if len(stats_dict["recent_packets"]) > 5: stats_dict["recent_packets"].pop(0) - packets_processed += 1 - - # Log progress every 100 packets - if packets_processed % 100 == 0: - packet_duration = time.time() - packet_start - logger.info( - f"TIMING: Processed {packets_processed} packets, last packet took {packet_duration:.3f}s" - ) - - # -------------------------------------------------- - # Indirect path processing (entire traceroute path) - # -------------------------------------------------- - if len(rf_hops) > 1: - # Calculate total distance of the full path - path_distance_km = sum( - h.distance_km or 0.0 for h in rf_hops - ) - - # Skip if it doesn't meet distance threshold - if path_distance_km < min_distance_km: - pass # too short – ignore - else: - # Average SNR across hops (ignore missing/garbage values) - valid_snrs = [ - h.snr - for h in rf_hops - if is_plausible_traceroute_snr(h.snr) - ] - avg_path_snr = ( - (sum(valid_snrs) / len(valid_snrs)) - if valid_snrs - else None - ) - - # Apply SNR filter (only if we actually have a value) - if avg_path_snr is None or avg_path_snr < min_snr: - pass # SNR below threshold – ignore - else: - from_node_id_path = rf_hops[0].from_node_id - to_node_id_path = rf_hops[-1].to_node_id - - path_key = (from_node_id_path, to_node_id_path) - - if path_key not in path_stats: - path_stats[path_key] = { - "from_node_name": rf_hops[0].from_node_name, - "to_node_name": rf_hops[-1].to_node_name, - "total_distance": 0.0, - "total_snr": 0.0, - "traceroute_count": 0, - "hop_count_total": 0, - "recent_packets": [], - "route_preview": [ - h.from_node_name for h in rf_hops - ] - + [rf_hops[-1].to_node_name], - "max_distance": 0.0, - } - - pstats = path_stats[path_key] - - # Update aggregates - pstats["traceroute_count"] += 1 - pstats["total_distance"] += path_distance_km - pstats["hop_count_total"] += len(rf_hops) - pstats["total_snr"] += avg_path_snr - pstats["max_distance"] = max( - pstats["max_distance"], path_distance_km - ) - - pstats["recent_packets"].append(packet["id"]) - if len(pstats["recent_packets"]) > 5: - pstats["recent_packets"].pop(0) - - except Exception as e: - logger.warning( - f"Error processing packet {packet.get('id', 'unknown')} for longest links: {e}" - ) + # Indirect path processing: only contiguous runs of qualifying + # hops are real multi-hop paths. When a middle hop fails the + # filters, the remaining hops are disconnected segments whose + # endpoints and distance sums must not be joined. + for segment in _contiguous_path_segments(path_hops): + segment_distances = [h["_distance_km"] for h in segment] + if any(d is None for d in segment_distances): continue - - process_duration = time.time() - process_start - logger.info(f"TIMING: Packet processing took {process_duration:.3f}s") - logger.info( - f"TIMING: Processed {packets_processed} packets, {hops_processed} hops, {distance_calculations} distance calculations" - ) - logger.info( - f"TIMING: Cache performance - {cache_hits} hits, {cache_misses} misses, final size: {len(location_cache)}" - ) - - logger.info( - f"Location cache efficiency: {len(location_cache)} unique location lookups cached" - ) - - # ------------------------------------------------------------------ - # Build the final list from aggregated statistics. - # ------------------------------------------------------------------ - build_start = time.time() - analyzed_links: list[dict[str, Any]] = [] - analyzed_paths: list[dict[str, Any]] = [] - - for (node1_id, node2_id), stats in link_stats.items(): - if stats["traceroute_count"] == 0: + path_distance_km = sum(segment_distances) + if path_distance_km < min_distance_km: continue - - avg_distance = stats["total_distance"] / stats["traceroute_count"] - avg_snr = stats["total_snr"] / stats["traceroute_count"] - - # Get last_seen timestamp from the most recent packet - last_seen = None - if stats["recent_packets"] and len(stats["recent_packets"]) > 0: - packet_id = stats["recent_packets"][0] - pkt = next( - (p for p in result["packets"] if p["id"] == packet_id), None - ) - if pkt and "timestamp" in pkt: - last_seen = pkt["timestamp"] - - packet_id = ( - stats["recent_packets"][0] if stats["recent_packets"] else None - ) - packet_url = ( - f"/packet/{packet_id}" if packet_id is not None else None - ) - - analyzed_links.append( - { - "from_node_id": node1_id, - "to_node_id": node2_id, - "from_node_name": stats["from_node_name"], - "to_node_name": stats["to_node_name"], - "distance_km": round(avg_distance, 2), - "avg_snr": round(avg_snr, 1), - "traceroute_count": stats["traceroute_count"], - "recent_packets": sorted( - stats["recent_packets"], reverse=True - ), - "packet_id": packet_id, - "packet_url": packet_url, - "last_seen": last_seen, - } - ) - - # Build indirect paths results - for (from_id, to_id), stats in path_stats.items(): - if stats["traceroute_count"] == 0: + avg_path_snr = sum(h["snr"] for h in segment) / len(segment) + if avg_path_snr < min_snr: continue + from_id_path = segment[0]["from_node_id"] + to_id_path = segment[-1]["to_node_id"] + p_key = (from_id_path, to_id_path) + if p_key not in path_stats: + from_name = node_names.get(from_id_path, f"!{from_id_path:08x}") + to_name = node_names.get(to_id_path, f"!{to_id_path:08x}") + route_preview = [ + node_names.get(h["from_node_id"], f"!{h['from_node_id']:08x}") + for h in segment + ] + [node_names.get(to_id_path, f"!{to_id_path:08x}")] + path_stats[p_key] = { + "from_node_name": from_name, + "to_node_name": to_name, + "total_distance": 0.0, + "total_snr": 0.0, + "traceroute_count": 0, + "hop_count_total": 0, + "recent_packets": [], + "route_preview": route_preview, + "max_distance": 0.0, + "last_seen": segment[0]["timestamp"], + } + pstats = path_stats[p_key] + pstats["traceroute_count"] += 1 + pstats["total_distance"] += path_distance_km + pstats["hop_count_total"] += len(segment) + pstats["total_snr"] += avg_path_snr + pstats["max_distance"] = max(pstats["max_distance"], path_distance_km) + ts = segment[0]["timestamp"] + if ts > pstats["last_seen"]: + pstats["last_seen"] = ts + if packet_id not in pstats["recent_packets"]: + pstats["recent_packets"].append(packet_id) + if len(pstats["recent_packets"]) > 5: + pstats["recent_packets"].pop(0) + + process_duration = time.time() - process_start + logger.info(f"TIMING: Hop processing took {process_duration:.3f}s") - avg_distance = stats["total_distance"] / stats["traceroute_count"] - avg_snr = ( - (stats["total_snr"] / stats["traceroute_count"]) - if stats["total_snr"] - else None - ) + # ------------------------------------------------------------------ + # Build the final list from aggregated statistics. + # ------------------------------------------------------------------ + build_start = time.time() + analyzed_links: list[dict[str, Any]] = [] + analyzed_paths: list[dict[str, Any]] = [] - # Determine last_seen timestamp - last_seen = None - if stats["recent_packets"]: - pkt_id = stats["recent_packets"][0] - pkt_obj = next( - (p for p in result["packets"] if p["id"] == pkt_id), None - ) - if pkt_obj and "timestamp" in pkt_obj: - last_seen = pkt_obj["timestamp"] + for (node1_id, node2_id), stats in link_stats.items(): + if stats["traceroute_count"] == 0: + continue - pkt_id = ( - stats["recent_packets"][0] if stats["recent_packets"] else None - ) - pkt_url = f"/packet/{pkt_id}" if pkt_id is not None else None + avg_distance = stats["total_distance"] / stats["traceroute_count"] + avg_snr = stats["total_snr"] / stats["traceroute_count"] + packet_id = stats["recent_packets"][0] if stats["recent_packets"] else None + packet_url = f"/packet/{packet_id}" if packet_id is not None else None - analyzed_paths.append( - { - "from_node_id": from_id, - "to_node_id": to_id, - "from_node_name": stats["from_node_name"], - "to_node_name": stats["to_node_name"], - "total_distance_km": round(avg_distance, 2), - "hop_count": int( - round( - stats["hop_count_total"] / stats["traceroute_count"] - ) - ), - "avg_snr": round(avg_snr, 1) - if avg_snr is not None - else None, - "traceroute_count": stats["traceroute_count"], - "route_preview": stats["route_preview"], - "recent_packets": sorted( - stats["recent_packets"], reverse=True - ), - "packet_id": pkt_id, - "packet_url": pkt_url, - "last_seen": last_seen, - } - ) + analyzed_links.append( + { + "from_node_id": node1_id, + "to_node_id": node2_id, + "from_node_name": stats["from_node_name"], + "to_node_name": stats["to_node_name"], + "distance_km": round(avg_distance, 2), + "avg_snr": round(avg_snr, 1), + "traceroute_count": stats["traceroute_count"], + "recent_packets": sorted(stats["recent_packets"], reverse=True), + "packet_id": packet_id, + "packet_url": packet_url, + "last_seen": stats["last_seen"], + } + ) - # Sort and trim results to the requested maximum - sort_start = time.time() - analyzed_links.sort(key=lambda x: x["distance_km"], reverse=True) - analyzed_links = analyzed_links[:max_results] - sort_duration = time.time() - sort_start + for (from_id, to_id), stats in path_stats.items(): + if stats["traceroute_count"] == 0: + continue - # Sort and trim indirect paths - analyzed_paths.sort(key=lambda x: x["total_distance_km"], reverse=True) - analyzed_paths = analyzed_paths[:max_results] + avg_distance = stats["total_distance"] / stats["traceroute_count"] + avg_snr = (stats["total_snr"] / stats["traceroute_count"]) if stats["total_snr"] else None + pkt_id = stats["recent_packets"][0] if stats["recent_packets"] else None + pkt_url = f"/packet/{pkt_id}" if pkt_id is not None else None - build_duration = time.time() - build_start - logger.info( - f"TIMING: Result building took {build_duration:.3f}s (sort: {sort_duration:.3f}s)" + analyzed_paths.append( + { + "from_node_id": from_id, + "to_node_id": to_id, + "from_node_name": stats["from_node_name"], + "to_node_name": stats["to_node_name"], + "total_distance_km": round(avg_distance, 2), + "hop_count": int(round(stats["hop_count_total"] / stats["traceroute_count"])), + "avg_snr": round(avg_snr, 1) if avg_snr is not None else None, + "traceroute_count": stats["traceroute_count"], + "route_preview": stats["route_preview"], + "recent_packets": sorted(stats["recent_packets"], reverse=True), + "packet_id": pkt_id, + "packet_url": pkt_url, + "last_seen": stats["last_seen"], + } ) - # ------------------------------------------------------------------ - # Compose summary. - # ------------------------------------------------------------------ - summary_start = time.time() - total_links = len(analyzed_links) + len(analyzed_paths) - - # Format longest distances as strings for summary - longest_direct = None - if analyzed_links: - longest_direct = f"{analyzed_links[0]['distance_km']:.2f} km" - - longest_path = None - if analyzed_paths: - longest_path = f"{analyzed_paths[0]['total_distance_km']:.2f} km" - - result_dict = { - "summary": { - "total_links": total_links, - "direct_links": len(analyzed_links), - "longest_direct": longest_direct, - "longest_path": longest_path, - }, - "direct_links": analyzed_links, - "indirect_links": analyzed_paths, - "criteria": { - "min_distance_km": min_distance_km, - "min_snr": min_snr, - "max_results": max_results, - "analysis_period_days": 7, - }, - "cache_stats": { - "location_lookups_cached": len(location_cache), - }, - } + analyzed_links.sort(key=lambda x: x["distance_km"], reverse=True) + analyzed_links = analyzed_links[:max_results] - summary_duration = time.time() - summary_start - total_duration = time.time() - start_time + analyzed_paths.sort(key=lambda x: x["total_distance_km"], reverse=True) + analyzed_paths = analyzed_paths[:max_results] - logger.info(f"TIMING: Summary creation took {summary_duration:.3f}s") - logger.info(f"TIMING: Total function duration: {total_duration:.3f}s") - logger.info( - f"TIMING: Breakdown - Fetch: {fetch_duration:.3f}s ({fetch_duration / total_duration * 100:.1f}%), " - f"Prefetch: {prefetch_duration:.3f}s ({prefetch_duration / total_duration * 100:.1f}%), " - f"Process: {process_duration:.3f}s ({process_duration / total_duration * 100:.1f}%), " - f"Build: {build_duration:.3f}s ({build_duration / total_duration * 100:.1f}%)" - ) + build_duration = time.time() - build_start + logger.info(f"TIMING: Result building took {build_duration:.3f}s") - return result_dict - - finally: - # Restore original implementation to avoid side-effects - _tru.get_node_location_at_timestamp = cast(Any, _orig_get_location) + longest_direct = f"{analyzed_links[0]['distance_km']:.2f} km" if analyzed_links else None + longest_path = f"{analyzed_paths[0]['total_distance_km']:.2f} km" if analyzed_paths else None + result_dict = { + "summary": { + "total_links": len(analyzed_links) + len(analyzed_paths), + "direct_links": len(analyzed_links), + "longest_direct": longest_direct, + "longest_path": longest_path, + }, + "direct_links": analyzed_links, + "indirect_links": analyzed_paths, + "criteria": { + "min_distance_km": min_distance_km, + "min_snr": min_snr, + "max_results": max_results, + "analysis_period_days": 7, + }, + "cache_stats": { + "location_lookups_cached": len(location_cache), + }, + } + total_duration = time.time() - start_time + logger.info(f"TIMING: Total longest links duration: {total_duration:.3f}s") + return result_dict except Exception as e: logger.error(f"Error in longest links analysis: {e}") raise @@ -1021,7 +753,7 @@ def get_network_graph_data( min_snr: float = -200.0, include_indirect: bool = False, filters: dict | None = None, - limit_packets: int = DEFAULT_GRAPH_PACKET_LIMIT, + limit_packets: int = -1, ) -> dict[str, Any]: """ Extract RF links from traceroute data to build a network connectivity graph. @@ -1031,7 +763,7 @@ def get_network_graph_data( min_snr: Minimum SNR threshold for including links include_indirect: Whether to include indirect (multi-hop) connections filters: Optional filters dict with start_time, end_time, gateway_id, etc. - limit_packets: Maximum number of packets to analyze + limit_packets: Maximum number of packets to analyze (-1 for unlimited) Returns: Dictionary with nodes and links data for graph visualization @@ -1077,11 +809,11 @@ def get_network_graph_data( # Always filter for successfully processed packets filters["processed_successfully_only"] = True - # Get traceroute data - packets = TracerouteRepository.get_traceroute_packets_for_graph( - limit=limit_packets, - filters=filters, - ) + # Get traceroute hops directly from materialized traceroute_hops. + # The query returns complete hop sequences so path structure is + # preserved; SNR filtering happens per hop below and continuity is + # validated before any path-level aggregation. + hops = get_traceroute_hops_for_graph(filters=filters) # Track nodes and links nodes = {} # node_id -> node_data @@ -1090,138 +822,111 @@ def get_network_graph_data( # Statistics stats = { - "packets_analyzed": len(packets), - "packets_with_rf_hops": 0, - "total_rf_hops": 0, + "packets_analyzed": len({h["packet_id"] for h in hops}), + "packets_with_rf_hops": len({(h["packet_id"], h.get("direction", "forward")) for h in hops}), + "total_rf_hops": len(hops), "links_found": 0, "links_filtered_by_snr": 0, "links_filtered_due_to_snr_0": 0, } - # Process each traceroute packet - for tr_data in packets: - if not tr_data["raw_payload"]: - continue - - try: - # Create TraceroutePacket object for analysis - tr_packet = TraceroutePacket( - packet_data=tr_data, resolve_names=False - ) - - # Get RF hops (actual radio transmissions) - rf_hops = tr_packet.get_rf_hops() - if not rf_hops: + # Group hops by (packet_id, direction) to preserve path context + hops_by_path: dict[tuple[int, str], list[dict[str, Any]]] = {} + for hop in hops: + path_key = (hop["packet_id"], hop.get("direction", "forward")) + if path_key not in hops_by_path: + hops_by_path[path_key] = [] + hops_by_path[path_key].append(hop) + + for (packet_id, _direction), rf_hops in hops_by_path.items(): + ts = rf_hops[0]["timestamp"] + for hop in rf_hops: + snr = hop["snr"] + if not is_plausible_traceroute_snr(snr) or ( + min_snr != -200 and snr < min_snr + ): + stats["links_filtered_by_snr"] += 1 + continue + if snr == 0: + stats["links_filtered_due_to_snr_0"] += 1 + continue + from_id = hop["from_node_id"] + to_id = hop["to_node_id"] + if 4294967295 in (from_id, to_id): continue - stats["packets_with_rf_hops"] += 1 - stats["total_rf_hops"] += len(rf_hops) - - # Process direct RF links - for hop in rf_hops: - # Filter by SNR - if min_snr is -200, it means "no limit" so only filter missing/garbage values - if not is_plausible_traceroute_snr(hop.snr) or ( - min_snr != -200 and hop.snr < min_snr - ): - stats["links_filtered_by_snr"] += 1 - continue - # filter 0db links (MQTT or UDP) - if hop.snr == 0: - stats["links_filtered_due_to_snr_0"] += 1 + # Add nodes to the graph + for node_id in (from_id, to_id): + if node_id not in nodes: + nodes[node_id] = { + "id": node_id, + "name": f"!{node_id:08x}", + "packet_count": 0, + "total_snr": 0.0, + "snr_count": 0, + "connections": set(), + "last_seen": ts, + } + nodes[node_id]["packet_count"] += 1 + if ts > nodes[node_id]["last_seen"]: + nodes[node_id]["last_seen"] = ts + + link_key = tuple(sorted([from_id, to_id])) + if link_key not in direct_links: + direct_links[link_key] = { + "source": link_key[0], + "target": link_key[1], + "snr_values": [snr], + "packet_count": 1, + "last_seen": ts, + "last_packet_id": packet_id, + } + stats["links_found"] += 1 + else: + link = direct_links[link_key] + link["snr_values"].append(snr) + link["packet_count"] += 1 + if ts > link["last_seen"]: + link["last_seen"] = ts + link["last_packet_id"] = packet_id + + nodes[from_id]["connections"].add(to_id) + nodes[to_id]["connections"].add(from_id) + nodes[from_id]["total_snr"] += snr + nodes[from_id]["snr_count"] += 1 + + # Process indirect connections if requested. Only contiguous + # runs of qualifying hops count as a path: a route whose middle + # hop failed the SNR filters is two disconnected segments, not + # a shortcut between its endpoints. + if include_indirect: + for segment in _contiguous_path_segments( + rf_hops, lambda hop: _rf_hop_qualifies(hop, min_snr) + ): + first_from = segment[0]["from_node_id"] + last_to = segment[-1]["to_node_id"] + if 4294967295 in (first_from, last_to): continue - if 4294967295 in [hop.from_node_id, hop.to_node_id]: + indirect_key = tuple(sorted([first_from, last_to])) + if indirect_key in direct_links: continue - # Add nodes to the graph - for node_id, node_name in [ - (hop.from_node_id, hop.from_node_name), - (hop.to_node_id, hop.to_node_name), - ]: - if node_id not in nodes: - nodes[node_id] = { - "id": node_id, - "name": node_name or f"!{node_id:08x}", - "packet_count": 0, - "total_snr": 0.0, - "snr_count": 0, - "connections": set(), - "last_seen": tr_data["timestamp"], - } - - # Update node stats - nodes[node_id]["packet_count"] += 1 - if tr_data["timestamp"] > nodes[node_id]["last_seen"]: - nodes[node_id]["last_seen"] = tr_data["timestamp"] - - # Create bidirectional link key (sorted to ensure consistency) - link_key = tuple(sorted([hop.from_node_id, hop.to_node_id])) - - # Add/update direct link - if link_key not in direct_links: - direct_links[link_key] = { - "source": link_key[0], - "target": link_key[1], - "snr_values": [hop.snr], - "packet_count": 1, - "last_seen": tr_data["timestamp"], - "last_packet_id": tr_data["id"], + path_snrs = [h["snr"] for h in segment] + if indirect_key not in indirect_connections: + indirect_connections[indirect_key] = { + "source": indirect_key[0], + "target": indirect_key[1], + "hop_count": len(segment), + "path_count": 1, + "avg_snr": sum(path_snrs) / len(path_snrs), + "last_seen": ts, + "last_packet_id": packet_id, } - stats["links_found"] += 1 else: - link = direct_links[link_key] - link["snr_values"].append(hop.snr) - link["packet_count"] += 1 - if tr_data["timestamp"] > link["last_seen"]: - link["last_seen"] = tr_data["timestamp"] - link["last_packet_id"] = tr_data["id"] - - # Track connections for nodes - nodes[hop.from_node_id]["connections"].add(hop.to_node_id) - nodes[hop.to_node_id]["connections"].add(hop.from_node_id) - nodes[hop.from_node_id]["total_snr"] += hop.snr - nodes[hop.from_node_id]["snr_count"] += 1 - - # Process indirect connections if requested - if include_indirect and len(rf_hops) > 1: - # Find endpoints of multi-hop paths - first_hop = rf_hops[0] - last_hop = rf_hops[-1] - - # Create indirect connection key - indirect_key = tuple( - sorted([first_hop.from_node_id, last_hop.to_node_id]) - ) - - # Only add if it's not already a direct link - if indirect_key not in direct_links: - if indirect_key not in indirect_connections: - path_snrs = [ - h.snr - for h in rf_hops - if h.snr and is_plausible_traceroute_snr(h.snr) - ] - indirect_connections[indirect_key] = { - "source": indirect_key[0], - "target": indirect_key[1], - "hop_count": len(rf_hops), - "path_count": 1, - "avg_snr": (sum(path_snrs) / len(path_snrs)) - if path_snrs - else None, - "last_seen": tr_data["timestamp"], - "last_packet_id": tr_data["id"], - } - else: - conn = indirect_connections[indirect_key] - conn["path_count"] += 1 - if tr_data["timestamp"] > conn["last_seen"]: - conn["last_seen"] = tr_data["timestamp"] - conn["last_packet_id"] = tr_data["id"] - - except Exception as e: - logger.warning( - f"Error processing traceroute packet {tr_data['id']}: {e}" - ) - continue + conn = indirect_connections[indirect_key] + conn["path_count"] += 1 + if ts > conn["last_seen"]: + conn["last_seen"] = ts + conn["last_packet_id"] = packet_id node_ids = list(nodes.keys()) node_names = get_bulk_node_names(node_ids) if node_ids else {} diff --git a/src/malla/static/js/modern-table.js b/src/malla/static/js/modern-table.js index f726a3f7..3554b960 100644 --- a/src/malla/static/js/modern-table.js +++ b/src/malla/static/js/modern-table.js @@ -25,7 +25,8 @@ class ModernTable { data: [], totalCount: 0, totalPages: 0, - isGrouped: false + isGrouped: false, + totalCountExact: false }; this.searchTimeout = null; @@ -249,10 +250,14 @@ class ModernTable { this.state.totalCount = data.total_count || 0; this.state.totalPages = Math.ceil(this.state.totalCount / this.state.pageSize); this.state.isGrouped = params.get('group_packets') === 'true'; + this.state.totalCountExact = true; this.renderTableBody(); this.updatePagination(); - this.emit('dataLoaded', { data: this.state.data, totalCount: this.state.totalCount }); + this.emit('dataLoaded', { + data: this.state.data, + totalCount: this.state.totalCount, + }); } catch (error) { console.error('Error loading table data:', error); this.showError(error.message); @@ -386,7 +391,7 @@ class ModernTable { document.getElementById(`${this.container.id}-end`).textContent = end; const totalElement = document.getElementById(`${this.container.id}-total`); - if (this.state.isGrouped && this.state.data.length === this.state.pageSize) { + if (this.hasEstimatedCount()) { totalElement.textContent = `${this.state.totalCount}+`; totalElement.title = 'Estimated count (optimized for performance)'; } else { @@ -415,6 +420,14 @@ class ModernTable { paginationContainer.addEventListener('click', this.paginationClickHandler); } + hasEstimatedCount() { + return Boolean( + this.state.isGrouped && + !this.state.totalCountExact && + this.state.data.length === this.state.pageSize + ); + } + renderPaginationButtons() { const { page, totalPages } = this.state; const nodes = []; @@ -425,7 +438,7 @@ class ModernTable { disabled: page <= 1 }, icon('bi bi-chevron-left'), textNode(' Previous'))); - if (this.state.isGrouped && this.state.data.length === this.state.pageSize) { + if (this.hasEstimatedCount()) { const maxDisplayPages = Math.min(totalPages, page + 5); for (let i = Math.max(1, page - 2); i <= Math.min(maxDisplayPages, page + 2); i++) { nodes.push(this.createPageButton(i, i === page)); @@ -456,7 +469,7 @@ class ModernTable { } } - const hasNextPage = this.state.isGrouped ? + const hasNextPage = this.hasEstimatedCount() ? this.state.data.length === this.state.pageSize : page < totalPages; diff --git a/src/malla/static/js/url-filter-manager.js b/src/malla/static/js/url-filter-manager.js index 9f96b0c9..9fad8e17 100644 --- a/src/malla/static/js/url-filter-manager.js +++ b/src/malla/static/js/url-filter-manager.js @@ -12,6 +12,29 @@ class URLFilterManager { this.form = document.querySelector(this.options.formSelector); this.groupingCheckbox = document.querySelector(this.options.groupingSelector); + this._userEditedFields = new Set(); + this._trackUserEdits(); + } + + /** + * Track fields the user has interacted with, so that the async + * applyURLParameters() restoration never overwrites user input + * (it can still be in flight when a user starts typing after a reload). + */ + _trackUserEdits() { + const markEdited = (event) => { + const target = event.target; + if (!target) return; + // Prefer the field name; picker search inputs only carry an id + const field = target.name || target.id; + if (field) { + this._userEditedFields.add(field); + } + }; + // Listen at document level: some controls (e.g. the grouping + // checkbox) may live outside the form element + document.addEventListener('input', markEdited, true); + document.addEventListener('change', markEdited, true); } /** @@ -58,6 +81,7 @@ class URLFilterManager { // Apply simple form field values Object.entries(urlParams).forEach(([key, value]) => { + if (this._userEditedFields.has(key)) return; // never clobber user input if (key === 'group_packets') { if (this.groupingCheckbox) { this.groupingCheckbox.checked = value === 'true'; @@ -102,6 +126,8 @@ class URLFilterManager { * Set node picker value by node ID */ async setNodePickerValue(fieldName, nodeId) { + // Skip if the user edited this field while we were initialising + if (this._userEditedFields.has(fieldName)) return; try { let displayName = `Node ${nodeId}`; // Default fallback @@ -124,6 +150,9 @@ class URLFilterManager { document.querySelector(`input[data-field="${fieldName}"]`) || document.querySelector(`.node-picker-input[data-field="${fieldName}"]`); + // Re-check: the user may have edited this field while we fetched + if (this._userEditedFields.has(fieldName)) return; + if (hiddenField && visibleField) { hiddenField.value = nodeId; visibleField.value = displayName; @@ -141,6 +170,8 @@ class URLFilterManager { * Set gateway picker value by gateway ID */ async setGatewayPickerValue(fieldName, gatewayId) { + // Skip if the user edited this field while we were initialising + if (this._userEditedFields.has(fieldName)) return; try { // Normalise: convert !hex to decimal string if applicable let nodeId = gatewayId; @@ -197,6 +228,9 @@ class URLFilterManager { document.querySelector(`input[data-field="${fieldName}"]`) || document.querySelector(`.gateway-picker-input[data-field="${fieldName}"]`); + // Re-check: the user may have edited this field while we fetched + if (this._userEditedFields.has(fieldName)) return; + if (hiddenField && visibleField) { hiddenField.value = nodeId; visibleField.value = displayName; diff --git a/src/malla/templates/components/table_layout_macros.html b/src/malla/templates/components/table_layout_macros.html index 6f9703d9..d392a670 100644 --- a/src/malla/templates/components/table_layout_macros.html +++ b/src/malla/templates/components/table_layout_macros.html @@ -125,6 +125,9 @@
{{ title }}
.table-content { flex: 1; + display: flex; + flex-direction: column; + min-height: 0; overflow: hidden; padding: 0; background: var(--bs-body-bg); @@ -197,18 +200,37 @@
{{ title }}
margin-bottom: 0.25rem; } +/* Compact status banner within full-screen layout */ +.table-content .table-status-bar { + flex-shrink: 0; + padding: 0.35rem 1.25rem; + font-size: 0.8125rem; + line-height: 1.35; + margin: 0; + border-radius: 0; + border-left: 0; + border-right: 0; + border-top: 0; + border-bottom: 1px solid var(--bs-border-color); +} + /* Modern table within full-screen layout */ .table-content .modern-table-container { + flex: 1; + min-height: 0; height: 100%; display: flex; flex-direction: column; } -.table-content .modern-table { +.table-content .table-wrapper { flex: 1; + min-height: 0; + max-height: none; overflow: auto; } +.table-content .modern-pagination, .table-content .modern-table-pagination { flex-shrink: 0; background: var(--bs-tertiary-bg); diff --git a/src/malla/templates/map.html b/src/malla/templates/map.html index 426d5967..7ab62c46 100644 --- a/src/malla/templates/map.html +++ b/src/malla/templates/map.html @@ -441,6 +441,8 @@
Hop Depth
let firstDisplay = true; // prevent recenter after initial load let currentHopDepth = 1; // max hop depth to display around selected node let currentOverlay = null; +let locationsAbortController = null; // cancels superseded /api/locations fetches +let lastLoadedTimeKey = null; // time window the currently loaded data covers // Update map theme based on current theme setting function updateMapTheme() { @@ -565,7 +567,7 @@
Hop Depth
// Handle filter form submission $('#locationFilterForm').on('submit', function(e) { e.preventDefault(); - applyClientSideFilters(); + refreshMapData(); }); // Link type checkbox handlers @@ -613,7 +615,9 @@
Hop Depth
$('#endDateTime').val(''); } } - applyClientSideFilters(); + // Time window changed – reload so the server recomputes link + // aggregates for the selected period + refreshMapData(); }); // Clear dates button @@ -621,7 +625,7 @@
Hop Depth
$('#startDateTime').val(''); $('#endDateTime').val(''); $('#maxAge').val(''); - applyClientSideFilters(); + refreshMapData(); }); // Move hop depth section just below the selected details container @@ -703,14 +707,55 @@
Hop Depth
loadNodeLocations(); } -// Build filter parameters from form -function buildFilterParams() { +// Extract the selected time window (epoch seconds) from the date fields. +// The #maxAge select syncs into #startDateTime, so these inputs are the +// single source of truth for the server-side aggregation window. +function getTimeFilterParams() { const form = document.getElementById('locationFilterForm'); const formData = new FormData(form); + const params = {}; + const startDateTime = formData.get('startDateTime'); + const endDateTime = formData.get('endDateTime'); + if (startDateTime) { + const ts = new Date(startDateTime).getTime() / 1000; + if (!isNaN(ts)) params.start_time = Math.floor(ts); + } + if (endDateTime) { + const ts = new Date(endDateTime).getTime() / 1000; + if (!isNaN(ts)) params.end_time = Math.floor(ts); + } + return params; +} + +function timeFilterKey(timeParams) { + return `${timeParams.start_time ?? ''}-${timeParams.end_time ?? ''}`; +} + +// Reload from the server whenever the selected time window changed, so link +// aggregates (observation counts, averages) are recomputed for exactly that +// window. When the window is unchanged, fast categorical filters (role, +// channel, min contacts) are re-applied client-side without a round trip. +function refreshMapData() { + if (timeFilterKey(getTimeFilterParams()) !== lastLoadedTimeKey) { + loadNodeLocations(); + } else { + applyClientSideFilters(); + } +} + +// Build filter parameters from form +function buildFilterParams() { const params = new URLSearchParams(); - // Only send server-side filters (gateway, search, etc.) - // Age and role filtering is now done client-side + // Server-side filters: the selected time window drives server-side link + // aggregation so link metrics match the visible period. + const timeParams = getTimeFilterParams(); + if (timeParams.start_time !== undefined) { + params.set('start_time', timeParams.start_time); + } + if (timeParams.end_time !== undefined) { + params.set('end_time', timeParams.end_time); + } return params; } @@ -720,8 +765,19 @@
Hop Depth
try { showLoading(); + // Cancel any in-flight request so a stale response can never + // overwrite the results of a newer time window. + if (locationsAbortController) { + locationsAbortController.abort(); + } + locationsAbortController = new AbortController(); + const params = buildFilterParams(); - const response = await fetch(`/api/locations?${params.toString()}`); + lastLoadedTimeKey = timeFilterKey(getTimeFilterParams()); + + const response = await fetch(`/api/locations?${params.toString()}`, { + signal: locationsAbortController.signal + }); const data = await response.json(); if (data.error) { @@ -744,6 +800,10 @@
Hop Depth
hideLoading(); } catch (error) { + if (error.name === 'AbortError') { + // Superseded by a newer request – keep the loading state for it + return; + } console.error('Error loading node locations:', error); showError('Failed to load node locations'); } @@ -790,12 +850,30 @@
Hop Depth
let endTimestamp = null; if (startDateTime) { startTimestamp = new Date(startDateTime).getTime() / 1000; - filteredNodes = filteredNodes.filter(node => node.timestamp >= startTimestamp); } - if (endDateTime) { endTimestamp = new Date(endDateTime).getTime() / 1000; - filteredNodes = filteredNodes.filter(node => node.timestamp <= endTimestamp); + } + + // A node belongs in a historical range when any of its activity sources + // (position broadcast, traceroute hop, direct packet) falls inside the + // interval. Comparing the global-latest timestamp instead would hide a + // node from an earlier view merely because it stayed active later. + if (startTimestamp !== null || endTimestamp !== null) { + filteredNodes = filteredNodes.filter(node => { + const activityTs = [ + node.position_timestamp, + node.last_seen_network, + node.last_seen_packet, + ].filter(Boolean); + if (activityTs.length === 0) { + activityTs.push(node.timestamp || 0); + } + return activityTs.some(ts => + (startTimestamp === null || ts >= startTimestamp) && + (endTimestamp === null || ts <= endTimestamp) + ); + }); } // Filter links by minimum number of direct contacts (total_hops_seen) @@ -846,6 +924,12 @@
Hop Depth
markerClusterGroup.clearLayers(); nodeMarkers = []; + // Clear existing links unconditionally so no ghost links persist + tracerouteLinks.forEach(link => map.removeLayer(link)); + tracerouteLinks = []; + packetLinks.forEach(link => map.removeLayer(link)); + packetLinks = []; + // Add markers for each filtered node nodeData.forEach(node => { addNodeMarker(node); @@ -923,6 +1007,14 @@
Hop Depth
const button = el('button', { className: 'btn btn-sm btn-primary', type: 'button' }, 'View Details'); button.addEventListener('click', () => viewNodeDetails(node.node_id)); + let gpsNote = null; + if (node.position_timestamp && (node.timestamp - node.position_timestamp) > 3600) { + const gpsAgeHours = (Date.now() / 1000 - node.position_timestamp) / 3600; + gpsNote = el('div', { className: 'text-muted small' }, + el('span', null, `GPS fix: ${formatAge(gpsAgeHours)}`) + ); + } + return el('div', { className: 'node-marker-info' }, el('div', { className: 'node-marker-title' }, node.display_name), el('div', null, el('strong', null, 'ID:'), textNode(` !${nodeIdHex}`)), @@ -933,10 +1025,11 @@
Hop Depth
el('span', { className: 'badge', style: { backgroundColor: roleColor } }, node.role) ) : null, el('div', null, el('strong', null, 'Location:'), textNode(` ${node.latitude.toFixed(6)}, ${node.longitude.toFixed(6)}`)), + gpsNote, node.altitude ? el('div', null, el('strong', null, 'Altitude:'), textNode(` ${node.altitude}m`)) : null, node.hw_model ? el('div', null, el('strong', null, 'Hardware:'), textNode(` ${node.hw_model}`)) : null, el('div', null, - el('strong', null, 'Age:'), + el('strong', null, 'Last Active:'), textNode(' '), el('span', { className: `age-indicator ${ageClass}` }, formatAge(ageHours)) ), @@ -1079,6 +1172,12 @@
Hop Depth
text: 'View Node Details' }); + let locText = `${node.latitude.toFixed(6)}, ${node.longitude.toFixed(6)}`; + if (node.position_timestamp && (node.timestamp - node.position_timestamp) > 3600) { + const gpsAgeHours = (Date.now() / 1000 - node.position_timestamp) / 3600; + locText += ` (fix: ${formatAge(gpsAgeHours)})`; + } + setChildren(document.getElementById('selectedDetailsContent'), el('div', { className: 'row' }, el('div', { className: 'col-12' }, @@ -1088,8 +1187,8 @@
Hop Depth
) ), el('div', { className: 'row' }, - el('div', { className: 'col-6' }, el('strong', null, 'Location:'), el('br'), el('small', null, `${node.latitude.toFixed(6)}, ${node.longitude.toFixed(6)}`)), - el('div', { className: 'col-6' }, el('strong', null, 'Age:'), el('br'), el('span', { className: `age-indicator ${ageClass}` }, formatAge(ageHours))) + el('div', { className: 'col-6' }, el('strong', null, 'Location:'), el('br'), el('small', null, locText)), + el('div', { className: 'col-6' }, el('strong', null, 'Last Active:'), el('br'), el('span', { className: `age-indicator ${ageClass}` }, formatAge(ageHours))) ), node.role || node.altitude ? el('div', { className: 'row mt-2' }, node.role ? el('div', { className: 'col-6' }, el('strong', null, 'Role:'), el('br'), el('span', { className: 'badge', style: { backgroundColor: roleColor } }, node.role)) : null, diff --git a/src/malla/templates/traceroute.html b/src/malla/templates/traceroute.html index f5d6f5a6..d85f1f8b 100644 --- a/src/malla/templates/traceroute.html +++ b/src/malla/templates/traceroute.html @@ -15,7 +15,6 @@ {% call fullscreen_table_container("tracerouteTable", "Traceroute Analysis", "bi bi-map", "toggleSidebar") %} -
{% endcall %} diff --git a/src/malla/utils/geo_utils.py b/src/malla/utils/geo_utils.py index 45f128d3..b6ab1df7 100644 --- a/src/malla/utils/geo_utils.py +++ b/src/malla/utils/geo_utils.py @@ -4,6 +4,39 @@ import math +# Positions within this radius of (0, 0) are treated as firmware garbage. +# (0, 0) lies in the Atlantic Ocean ~600 km from the nearest coast, so real +# nodes are never affected. Firmware bugs emit near-zero coordinates (e.g. +# 0.00012, -0.003) instead of exactly (0, 0), which dodge simple == 0 checks. +NULL_ISLAND_RADIUS_KM = 50.0 + + +def is_valid_position(latitude: float | None, longitude: float | None) -> bool: + """ + Check whether decoded coordinates represent a plausible real location. + + Rejects missing values, non-finite numbers, out-of-range coordinates, and + positions near "null island" (0, 0) produced by firmware bugs. + + Args: + latitude: Latitude in decimal degrees (or None if unset) + longitude: Longitude in decimal degrees (or None if unset) + + Returns: + True if the position is plausible and safe to use. + """ + if latitude is None or longitude is None: + return False + if not (math.isfinite(latitude) and math.isfinite(longitude)): + return False + if not -90.0 <= latitude <= 90.0: + return False + if not -180.0 <= longitude <= 180.0: + return False + if calculate_distance(latitude, longitude, 0.0, 0.0) < NULL_ISLAND_RADIUS_KM: + return False + return True + def calculate_distance(lat1: float, lon1: float, lat2: float, lon2: float) -> float: """ diff --git a/src/malla/utils/traceroute_utils.py b/src/malla/utils/traceroute_utils.py index adf54d62..7b94b435 100644 --- a/src/malla/utils/traceroute_utils.py +++ b/src/malla/utils/traceroute_utils.py @@ -19,6 +19,18 @@ class RouteData(TypedDict): snr_back: list[float] +def decode_traceroute_payload(raw_payload: bytes) -> RouteData: + """Decode RouteDiscovery, raising on malformed payloads instead of hiding them.""" + route_discovery = mesh_pb2.RouteDiscovery() + route_discovery.ParseFromString(raw_payload) + return RouteData( + route_nodes=list(route_discovery.route), + snr_towards=[snr / 4.0 for snr in route_discovery.snr_towards], + route_back=list(route_discovery.route_back), + snr_back=[snr / 4.0 for snr in route_discovery.snr_back], + ) + + def parse_traceroute_payload(raw_payload: bytes) -> RouteData: """ Parse traceroute payload from raw bytes using protobuf parsing. @@ -44,18 +56,7 @@ def parse_traceroute_payload(raw_payload: bytes) -> RouteData: return RouteData(route_nodes=[], snr_towards=[], route_back=[], snr_back=[]) try: - # Try protobuf parsing - route_discovery = mesh_pb2.RouteDiscovery() - route_discovery.ParseFromString(raw_payload) - - result = RouteData( - route_nodes=[int(node_id) for node_id in route_discovery.route], - # Convert SNR from scaled integer to actual dB (divide by 4) - snr_towards=[float(snr) / 4.0 for snr in route_discovery.snr_towards], - route_back=[int(node_id) for node_id in route_discovery.route_back], - # Convert SNR from scaled integer to actual dB (divide by 4) - snr_back=[float(snr) / 4.0 for snr in route_discovery.snr_back], - ) + result = decode_traceroute_payload(raw_payload) logger.debug( f"Protobuf parsing successful: {len(result['route_nodes'])} nodes, " diff --git a/tests/e2e/test_map_filters.py b/tests/e2e/test_map_filters.py index c98b9f7b..9fa180a0 100644 --- a/tests/e2e/test_map_filters.py +++ b/tests/e2e/test_map_filters.py @@ -88,8 +88,9 @@ def test_client_side_role_filtering(self, page: Page, test_server_url): assert len(location_requests) == 0, "Role filtering should be client-side only" @pytest.mark.e2e - def test_client_side_age_filtering(self, page: Page, test_server_url): - """Test that age filtering works on client-side without server requests.""" + def test_age_filter_triggers_window_reload(self, page: Page, test_server_url): + """Age filtering reloads /api/locations so the server aggregates links + over the selected window instead of showing multi-day totals.""" page.goto(f"{test_server_url}/map") # Wait for loading to complete @@ -116,9 +117,14 @@ def test_client_side_age_filtering(self, page: Page, test_server_url): # Wait for filtering to complete page.wait_for_timeout(2000) - # Check that no new API requests were made to /api/locations - location_requests = [req for req in requests if "/api/locations" in req] - assert len(location_requests) == 0, "Age filtering should be client-side only" + # A new /api/locations request carrying the 1-hour window must have + # been made so link metrics are recomputed server-side + location_requests = [ + req for req in requests if "/api/locations" in req and "start_time=" in req + ] + assert len(location_requests) >= 1, ( + "Age filtering must reload /api/locations with a start_time parameter" + ) # Verify filtering worked filtered_count = page.locator("#nodeCount").text_content() @@ -129,13 +135,13 @@ def test_client_side_age_filtering(self, page: Page, test_server_url): @pytest.mark.e2e def test_filter_reset_functionality(self, page: Page, test_server_url): - """Test that filters can be reset to show all data.""" + """Test that filters can be reset and re-applied to show the same data.""" page.goto(f"{test_server_url}/map") # Wait for loading to complete page.wait_for_selector("#mapLoading", state="hidden", timeout=DEFAULT_TIMEOUT) - # Get initial counts + # Get initial counts (default max age is 24 hours) initial_node_count = page.locator("#nodeCount").text_content() initial_link_count = page.locator("#statsLinks").text_content() @@ -150,13 +156,18 @@ def test_filter_reset_functionality(self, page: Page, test_server_url): apply_button.click() page.wait_for_timeout(2000) - # Reset filters + # Reset filters (No Limit covers the full 14-day server window) age_filter.select_option("") role_filter.select_option("") apply_button.click() page.wait_for_timeout(2000) - # Verify counts return to initial values + # Re-apply the original 24-hour window; counts should return to the + # initial values now that the same aggregation window is loaded + age_filter.select_option("24") + apply_button.click() + page.wait_for_timeout(2000) + final_node_count = page.locator("#nodeCount").text_content() final_link_count = page.locator("#statsLinks").text_content() diff --git a/tests/e2e/test_traceroute_filters_e2e.py b/tests/e2e/test_traceroute_filters_e2e.py index 60e3e0f5..8acc3d0f 100644 --- a/tests/e2e/test_traceroute_filters_e2e.py +++ b/tests/e2e/test_traceroute_filters_e2e.py @@ -1543,13 +1543,21 @@ def test_traceroute_route_node_filter_e2e(self, page: Page, test_server_url: str api_data = response.json() assert "data" in api_data, "API response should contain data field" - # Verify API results contain the route node + # Verify API results contain the route node anywhere in the path for item in api_data["data"]: route_nodes = item.get("route_nodes", []) - assert int(route_node_id) in route_nodes, ( - f"API result should contain route_node {route_node_id}, got {route_nodes}" + all_path_nodes = set(route_nodes) | { + item.get("from_node_id"), + item.get("to_node_id"), + } + assert int(route_node_id) in all_path_nodes, ( + f"API result should contain route_node {route_node_id} anywhere in route path, got {all_path_nodes}" ) + # Verify frontend table shows filtered results + filtered_rows = page.locator("#tracerouteTable tbody tr").count() + assert filtered_rows > 0, "Should have filtered results" + def test_traceroute_gateway_filter_e2e(self, page: Page, test_server_url: str): """Test that the gateway filter works end-to-end.""" page.goto(f"{test_server_url}/traceroute") @@ -2209,4 +2217,24 @@ def test_traceroute_url_manager_initialization( """) assert has_methods, "URL manager should have required methods" - assert has_methods, "URL manager should have required methods" + def test_traceroute_pagination_visibility( + self, page: Page, test_server_url: str + ): + """Test that pagination controls remain visible in viewport.""" + page.set_viewport_size({"width": 1280, "height": 800}) + page.goto(f"{test_server_url}/traceroute") + + # Wait for table to load + page.wait_for_selector("#tracerouteTable .modern-pagination", timeout=10000) + + # Verify pagination is fully inside the viewport + pagination_locator = page.locator("#tracerouteTable .modern-pagination") + expect(pagination_locator).to_be_visible() + pag_box = pagination_locator.bounding_box() + assert pag_box is not None + assert pag_box["y"] + pag_box["height"] <= 800.5, ( + f"Pagination bottom ({pag_box['y'] + pag_box['height']}) must not exceed viewport height (800)" + ) + + # Verify pagination navigation buttons exist and are visible + expect(page.locator(".pagination-btn").first).to_be_visible() diff --git a/tests/fixtures/database_fixtures.py b/tests/fixtures/database_fixtures.py index f455aa28..4e23c183 100644 --- a/tests/fixtures/database_fixtures.py +++ b/tests/fixtures/database_fixtures.py @@ -36,6 +36,7 @@ def create_test_database(self, db_path: str): logger.info(f"Creating test database at {db_path}") with sqlite3.connect(db_path) as conn: + conn.row_factory = sqlite3.Row cursor = conn.cursor() # Create the schema @@ -45,6 +46,18 @@ def create_test_database(self, db_path: str): self._insert_node_info(cursor) self._insert_packets(cursor) + # Materialize the fixture rows so tests exercise the materialized reader path. + from malla.database.schema import ensure_startup_schema + from malla.database.traceroutes import write_traceroute + + ensure_startup_schema(cursor) + traceroutes = cursor.execute( + "SELECT * FROM packet_history " + "WHERE portnum = 70 OR portnum_name = 'TRACEROUTE_APP'" + ).fetchall() + for traceroute in traceroutes: + write_traceroute(cursor, dict(traceroute)) + conn.commit() logger.info( diff --git a/tests/unit/test_api_locations_time_window.py b/tests/unit/test_api_locations_time_window.py new file mode 100644 index 00000000..239d50b7 --- /dev/null +++ b/tests/unit/test_api_locations_time_window.py @@ -0,0 +1,190 @@ +""" +Unit tests for the /api/locations time-window handling. + +The endpoint must aggregate link metrics (traceroute hops, direct packet +receptions) over the exact window selected on the map, while keeping the +node position lookup on a wide window so actively routing nodes whose last +GPS report is older than the selection stay visible. +""" + +import time as time_module +from unittest.mock import patch + +import pytest + +EMPTY_NETWORK_DATA = {"nodes": [], "links": []} + + +@pytest.fixture +def mocked_location_services(): + """Patch the expensive service calls used by /api/locations.""" + with ( + patch( + "src.malla.routes.api_routes.TracerouteService.get_network_graph_data", + return_value=EMPTY_NETWORK_DATA, + ) as graph_mock, + patch( + "src.malla.routes.api_routes.LocationService.get_packet_links", + return_value=[], + ) as packet_links_mock, + patch( + "src.malla.routes.api_routes.LocationService.get_node_locations", + return_value=[], + ) as node_locations_mock, + patch( + "src.malla.routes.api_routes.LocationService.get_traceroute_links", + return_value=[], + ) as traceroute_links_mock, + ): + yield { + "graph": graph_mock, + "packet_links": packet_links_mock, + "node_locations": node_locations_mock, + "traceroute_links": traceroute_links_mock, + } + + +class TestApiLocationsTimeWindow: + """Test /api/locations server-side time window resolution.""" + + @pytest.mark.unit + def test_hours_param_scopes_link_aggregation( + self, client, mocked_location_services + ): + """hours=1 aggregates links over the last hour only.""" + before = time_module.time() + response = client.get("/api/locations?hours=1") + after = time_module.time() + + assert response.status_code == 200 + + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + assert graph_filters["start_time"] >= before - 3600 + assert graph_filters["end_time"] <= after + + packet_filters = mocked_location_services["packet_links"].call_args.args[0] + assert packet_filters["start_time"] == graph_filters["start_time"] + + traceroute_filters = mocked_location_services[ + "traceroute_links" + ].call_args.args[0] + assert traceroute_filters["start_time"] == graph_filters["start_time"] + + @pytest.mark.unit + def test_start_end_params_scope_link_aggregation( + self, client, mocked_location_services + ): + """Explicit epoch start/end bound the link aggregation window.""" + end = time_module.time() - 7200 + start = end - 3600 + + response = client.get(f"/api/locations?start_time={start}&end_time={end}") + + assert response.status_code == 200 + + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + assert graph_filters["start_time"] == start + assert graph_filters["end_time"] == end + + packet_filters = mocked_location_services["packet_links"].call_args.args[0] + assert packet_filters["start_time"] == start + assert packet_filters["end_time"] == end + + traceroute_filters = mocked_location_services[ + "traceroute_links" + ].call_args.args[0] + assert traceroute_filters["start_time"] == start + assert traceroute_filters["end_time"] == end + + @pytest.mark.unit + def test_max_age_hours_alias_supported(self, client, mocked_location_services): + """max_age_hours is accepted as an alias for hours.""" + before = time_module.time() + response = client.get("/api/locations?max_age_hours=6") + assert response.status_code == 200 + + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + assert graph_filters["start_time"] >= before - 6 * 3600 - 5 + assert graph_filters["start_time"] <= before - 6 * 3600 + 3600 + + @pytest.mark.unit + def test_default_window_is_14_days_without_time_params( + self, client, mocked_location_services + ): + """No time parameters keeps the historical 14-day default window.""" + before = time_module.time() + + response = client.get("/api/locations") + + assert response.status_code == 200 + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + # Start must be roughly 14 days ago (allow small execution slack) + assert before - 14 * 24 * 3600 - 5 <= graph_filters["start_time"] + assert graph_filters["start_time"] <= before - 14 * 24 * 3600 + 5 + + @pytest.mark.unit + def test_window_capped_at_14_days(self, client, mocked_location_services): + """Windows larger than 14 days are clamped for performance.""" + end = time_module.time() + start = end - 20 * 24 * 3600 + + response = client.get(f"/api/locations?start_time={start}&end_time={end}") + + assert response.status_code == 200 + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + assert graph_filters["start_time"] >= end - 14 * 24 * 3600 - 5 + assert graph_filters["start_time"] <= end - 14 * 24 * 3600 + 5 + + @pytest.mark.unit + def test_invalid_time_range_returns_400(self, client, mocked_location_services): + """start_time >= end_time is rejected.""" + now = time_module.time() + response = client.get(f"/api/locations?start_time={now}&end_time={now - 10}") + assert response.status_code == 400 + + @pytest.mark.unit + def test_position_lookup_keeps_wide_window_despite_short_link_window( + self, client, mocked_location_services + ): + """Regression guard: a 1-hour link window must not narrow the GPS lookup. + + Nodes actively routing now whose last position broadcast is hours or + days old must remain visible at their last known good position, so + the position lookup keeps the wide 14-day window even when link + aggregates are computed for the last hour only. + """ + before = time_module.time() + + response = client.get("/api/locations?hours=1") + + assert response.status_code == 200 + + graph_filters = mocked_location_services["graph"].call_args.kwargs["filters"] + position_filters = mocked_location_services["node_locations"].call_args.args[0] + + # Link aggregation is scoped to the last hour... + assert graph_filters["start_time"] >= before - 3600 - 5 + + # ...while the position lookup keeps the wide window. + assert before - 14 * 24 * 3600 - 5 <= position_filters["start_time"] + assert position_filters["start_time"] <= before - 13 * 24 * 3600 + assert position_filters["start_time"] < graph_filters["start_time"] + + @pytest.mark.unit + def test_gateway_filter_applied_to_both_windows( + self, client, mocked_location_services + ): + """gateway_id filters both link aggregation and position lookup.""" + response = client.get("/api/locations?hours=1&gateway_id=42") + + assert response.status_code == 200 + + packet_filters = mocked_location_services["packet_links"].call_args.args[0] + position_filters = mocked_location_services["node_locations"].call_args.args[0] + assert packet_filters["gateway_id"] == 42 + assert position_filters["gateway_id"] == 42 + + @pytest.mark.unit + def test_invalid_gateway_id_returns_400(self, client, mocked_location_services): + response = client.get("/api/locations?gateway_id=not-a-number") + assert response.status_code == 400 diff --git a/tests/unit/test_gateway_sorting.py b/tests/unit/test_gateway_sorting.py index cfd965fc..a2f9bde5 100644 --- a/tests/unit/test_gateway_sorting.py +++ b/tests/unit/test_gateway_sorting.py @@ -337,77 +337,12 @@ def test_packet_repository_gateway_sorting_desc(self): assert packets[2]["mesh_packet_id"] == "abc123" def test_traceroute_repository_gateway_sorting_asc(self): - """Test that TracerouteRepository sorts by gateway_count in ascending order when requested.""" - # Mock database connection and cursor - mock_conn = MagicMock() - mock_cursor = MagicMock() - mock_conn.cursor.return_value = mock_cursor - - # Mock the database query results - # First call: total count query (returns count as tuple/row) - # Second call: no longer used (removed sample count estimation) - mock_cursor.fetchone.side_effect = [ - (2,), # Total count query result - ] - - # Mock the main query results - individual packets that will be grouped in memory - # These represent individual packet records, not pre-grouped results - mock_cursor.fetchall.return_value = [ - # First group: mesh_packet_id="trace123" with 1 gateway - { - "id": 1, - "timestamp": 1000, - "from_node_id": 123, - "to_node_id": 456, - "mesh_packet_id": "trace123", - "gateway_id": "!433d0c24", - "hop_start": 3, - "hop_limit": 1, - "rssi": -80, - "snr": 5, - "payload_length": 50, - "processed_successfully": 1, - "timestamp_str": "2024-01-01 12:00:00", - "raw_payload": b"test", - }, - # Second group: mesh_packet_id="trace456" with 2 gateways (2 individual records) - { - "id": 2, - "timestamp": 2000, - "from_node_id": 789, - "to_node_id": 456, - "mesh_packet_id": "trace456", - "gateway_id": "!433d0c24", - "hop_start": 4, - "hop_limit": 3, - "rssi": -75, - "snr": 8, - "payload_length": 75, - "processed_successfully": 1, - "timestamp_str": "2024-01-01 12:01:00", - "raw_payload": b"test2", - }, - { - "id": 3, - "timestamp": 2001, - "from_node_id": 789, - "to_node_id": 456, - "mesh_packet_id": "trace456", # Same mesh_packet_id as above - "gateway_id": "!da73e9cc", # Different gateway - "hop_start": 2, - "hop_limit": 1, - "rssi": -70, - "snr": 10, - "payload_length": 75, - "processed_successfully": 1, - "timestamp_str": "2024-01-01 12:01:01", - "raw_payload": b"test2_longer", # Longer payload to test best selection - }, - ] - + """The public repository delegates grouped sorting to the saved-row reader.""" + expected = {"packets": [], "total_count": 0} with patch( - "src.malla.database.repositories.get_db_connection", return_value=mock_conn - ): + "src.malla.database.traceroute_read_repository.get_traceroute_packets", + return_value=expected, + ) as reader: result = TracerouteRepository.get_traceroute_packets( limit=10, offset=0, @@ -416,32 +351,16 @@ def test_traceroute_repository_gateway_sorting_asc(self): group_packets=True, ) - # Verify results are correctly grouped and sorted by gateway count - packets = result["packets"] - assert len(packets) == 2 - - # First packet should have 1 gateway (ascending order) - assert packets[0]["gateway_count"] == 1 - assert packets[0]["mesh_packet_id"] == "trace123" - assert packets[0]["gateway_list"] == "!433d0c24" - - # Second packet should have 2 gateways - assert packets[1]["gateway_count"] == 2 - assert packets[1]["mesh_packet_id"] == "trace456" - # Gateway list order may vary due to set() usage, so check both gateways are present - gateway_list = packets[1]["gateway_list"] - assert "!433d0c24" in gateway_list - assert "!da73e9cc" in gateway_list - assert gateway_list.count(",") == 1 # Exactly 2 gateways - - # Verify aggregation worked correctly for the second group - assert packets[1]["min_rssi"] == -75 - assert packets[1]["max_rssi"] == -70 - assert packets[1]["min_snr"] == 8 - assert packets[1]["max_snr"] == 10 - - # Verify best payload was selected (longest one) - assert packets[1]["raw_payload"] == b"test2_longer" + assert result is expected + reader.assert_called_once_with( + limit=10, + offset=0, + filters=None, + order_by="gateway_id", + order_dir="asc", + search=None, + group_packets=True, + ) class TestGatewaySortingDataFormat: diff --git a/tests/unit/test_location_service.py b/tests/unit/test_location_service.py new file mode 100644 index 00000000..e3a2bad2 --- /dev/null +++ b/tests/unit/test_location_service.py @@ -0,0 +1,157 @@ +""" +Unit tests for LocationService node activity consolidation and location enrichment. +""" + +from unittest.mock import patch + +import pytest + +from malla.services.location_service import LocationService + + +class TestLocationServiceNodeLocations: + """Test LocationService.get_node_locations consolidation of timestamps.""" + + @pytest.mark.unit + def test_node_active_timestamp_consolidated_from_traceroute(self): + """When traceroute activity is newer than position, node.timestamp uses traceroute time.""" + pos_time = 1000.0 + tr_time = 5000.0 + + mock_raw_locations = [ + { + "node_id": 12345, + "hex_id": "!00003039", + "display_name": "Test Node", + "long_name": "Test Node Long", + "short_name": "TN", + "hw_model": "T-Beam", + "role": "ROUTER", + "latitude": 40.0, + "longitude": -95.0, + "altitude": 100, + "timestamp": pos_time, + "precision_bits": 16, + "precision_meters": 364.0, + "sats_in_view": 8, + } + ] + + mock_network_data = { + "nodes": [ + { + "id": 12345, + "name": "Test Node", + "packet_count": 5, + "avg_snr": 8.5, + "last_seen": tr_time, + } + ], + "links": [], + } + + with patch("malla.database.repositories.LocationRepository.get_node_locations", return_value=mock_raw_locations): + results = LocationService.get_node_locations( + filters={}, + network_data=mock_network_data, + packet_links=[], + ) + + assert len(results) == 1 + node = results[0] + # timestamp must be updated to the active (traceroute) timestamp + assert node["timestamp"] == tr_time + # position_timestamp must preserve original position packet timestamp + assert node["position_timestamp"] == pos_time + assert "position_timestamp_str" in node + assert node["last_seen_network"] == tr_time + + @pytest.mark.unit + def test_node_active_timestamp_consolidated_from_packet_link(self): + """When packet link activity is newer than position and traceroute, node.timestamp uses packet time.""" + pos_time = 1000.0 + pkt_time = 8000.0 + + mock_raw_locations = [ + { + "node_id": 12345, + "hex_id": "!00003039", + "display_name": "Test Node", + "long_name": "Test Node Long", + "short_name": "TN", + "hw_model": "T-Beam", + "role": "ROUTER", + "latitude": 40.0, + "longitude": -95.0, + "altitude": 100, + "timestamp": pos_time, + } + ] + + mock_packet_links = [ + { + "from_node_id": 12345, + "to_node_id": 99999, + "last_seen": pkt_time, + "total_hops_seen": 2, + } + ] + + with patch("malla.database.repositories.LocationRepository.get_node_locations", return_value=mock_raw_locations): + results = LocationService.get_node_locations( + filters={}, + network_data={"nodes": [], "links": []}, + packet_links=mock_packet_links, + ) + + assert len(results) == 1 + node = results[0] + assert node["timestamp"] == pkt_time + assert node["position_timestamp"] == pos_time + assert node["last_seen_packet"] == pkt_time + + @pytest.mark.unit + def test_node_active_timestamp_defaults_to_position_when_newest(self): + """When position is newest, node.timestamp remains the position timestamp.""" + pos_time = 10000.0 + tr_time = 5000.0 + + mock_raw_locations = [ + { + "node_id": 12345, + "hex_id": "!00003039", + "display_name": "Test Node", + "long_name": "Test Node Long", + "short_name": "TN", + "hw_model": "T-Beam", + "role": "ROUTER", + "latitude": 40.0, + "longitude": -95.0, + "altitude": 100, + "timestamp": pos_time, + } + ] + + mock_network_data = { + "nodes": [ + { + "id": 12345, + "name": "Test Node", + "packet_count": 1, + "last_seen": tr_time, + } + ], + "links": [], + } + + with patch("malla.database.repositories.LocationRepository.get_node_locations", return_value=mock_raw_locations): + results = LocationService.get_node_locations( + filters={}, + network_data=mock_network_data, + packet_links=[], + ) + + assert len(results) == 1 + node = results[0] + assert node["timestamp"] == pos_time + assert node["position_timestamp"] == pos_time diff --git a/tests/unit/test_position_validity.py b/tests/unit/test_position_validity.py new file mode 100644 index 00000000..e58efe6b --- /dev/null +++ b/tests/unit/test_position_validity.py @@ -0,0 +1,234 @@ +"""Tests for position validity filtering (null-island firmware bug). + +Firmware sometimes emits near-zero (~0, ~0) coordinates instead of an exact +(0, 0). Such fixes must never reach the map or the longest-link distance +calculations: readers fall back to the previous valid position instead. +""" + +import math +import sqlite3 +from contextlib import closing +from unittest.mock import patch + +import pytest +from meshtastic import mesh_pb2 + +from malla.database.repositories import LocationRepository +from malla.utils.geo_utils import is_valid_position + +pytestmark = pytest.mark.unit + +VALID_LAT = 52.37 +VALID_LON = 4.89 + + +@pytest.fixture +def database(tmp_path): + path = tmp_path / "positions.db" + with closing(sqlite3.connect(path)) as conn: + conn.row_factory = sqlite3.Row + conn.execute(""" + CREATE TABLE packet_history ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + timestamp REAL NOT NULL, + portnum INTEGER, + portnum_name TEXT, + from_node_id INTEGER, + raw_payload BLOB, + processed_successfully INTEGER DEFAULT 1 + ) + """) + conn.execute(""" + CREATE TABLE node_info ( + node_id INTEGER PRIMARY KEY, + long_name TEXT, + short_name TEXT, + hw_model TEXT, + role TEXT, + primary_channel TEXT + ) + """) + conn.commit() + return path + + +def _connection(path): + conn = sqlite3.connect(path) + conn.row_factory = sqlite3.Row + return conn + + +def _position_payload(lat, lon, altitude=42): + return mesh_pb2.Position( + latitude_i=int(lat * 1e7), longitude_i=int(lon * 1e7), altitude=altitude + ).SerializeToString() + + +def _insert_position(conn, node_id, timestamp, lat, lon): + conn.execute( + """ + INSERT INTO packet_history + (timestamp, portnum, portnum_name, from_node_id, raw_payload) + VALUES (?, 3, 'POSITION_APP', ?, ?) + """, + (timestamp, node_id, _position_payload(lat, lon)), + ) + conn.commit() + + +class TestIsValidPosition: + def test_accepts_normal_coordinates(self): + assert is_valid_position(VALID_LAT, VALID_LON) is True + + def test_accepts_equator_far_from_null_island(self): + assert is_valid_position(0.0, 36.8) is True + + def test_rejects_none(self): + assert is_valid_position(None, VALID_LON) is False + assert is_valid_position(VALID_LAT, None) is False + + def test_rejects_exact_zero(self): + assert is_valid_position(0.0, 0.0) is False + + def test_rejects_near_null_island_firmware_garbage(self): + assert is_valid_position(0.00012, -0.003) is False + assert is_valid_position(0.1, 0.1) is False + + def test_rejects_out_of_range(self): + assert is_valid_position(95.0, VALID_LON) is False + assert is_valid_position(VALID_LAT, 200.0) is False + assert is_valid_position(-91.0, VALID_LON) is False + + def test_rejects_non_finite(self): + assert is_valid_position(math.nan, VALID_LON) is False + assert is_valid_position(VALID_LAT, math.inf) is False + + +class TestGetNodeLocations: + def test_falls_back_to_previous_valid_fix(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, VALID_LAT, VALID_LON) + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + locations = LocationRepository.get_node_locations() + + assert len(locations) == 1 + entry = locations[0] + assert entry["node_id"] == 100 + assert entry["latitude"] == pytest.approx(VALID_LAT) + assert entry["longitude"] == pytest.approx(VALID_LON) + assert entry["timestamp"] == 100.0 + + def test_node_with_only_garbage_is_absent(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, 0.00012, -0.003) + _insert_position(conn, 100, 200.0, 0.0, 0.0) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + locations = LocationRepository.get_node_locations() + + assert locations == [] + + def test_uses_newest_valid_position(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, 51.9, 4.4) + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + _insert_position(conn, 100, 300.0, VALID_LAT, VALID_LON) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + locations = LocationRepository.get_node_locations() + + assert len(locations) == 1 + assert locations[0]["timestamp"] == 300.0 + assert locations[0]["latitude"] == pytest.approx(VALID_LAT) + + +class TestNodeLocationHistory: + def test_history_filters_garbage_rows(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, 51.9, 4.4) + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + _insert_position(conn, 100, 300.0, VALID_LAT, VALID_LON) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + history = LocationRepository.get_node_location_history(100) + batched = LocationRepository.get_nodes_location_history([100]) + + assert [h["timestamp"] for h in history] == [300.0, 100.0] + assert [h["timestamp"] for h in batched[100]] == [300.0, 100.0] + + +class TestGetLatestNodeLocation: + def test_falls_back_past_garbage_latest(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, VALID_LAT, VALID_LON) + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + location = LocationRepository.get_latest_node_location(100) + + assert location is not None + assert location["latitude"] == pytest.approx(VALID_LAT) + assert location["timestamp"] == 100.0 + + def test_returns_none_when_only_garbage(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + location = LocationRepository.get_latest_node_location(100) + + assert location is None + + +class TestGetNodeLocationAtTimestamp: + def test_skips_garbage_before_target(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, VALID_LAT, VALID_LON) + _insert_position(conn, 100, 200.0, 0.00012, -0.003) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + location = LocationRepository.get_node_location_at_timestamp(100, 250.0) + + assert location is not None + assert location["latitude"] == pytest.approx(VALID_LAT) + assert location["timestamp"] == 100.0 + assert "ago" in location["age_warning"] + + def test_falls_forward_to_valid_after_target(self, database): + with closing(_connection(database)) as conn: + _insert_position(conn, 100, 100.0, 0.00012, -0.003) + _insert_position(conn, 100, 200.0, VALID_LAT, VALID_LON) + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + location = LocationRepository.get_node_location_at_timestamp(100, 150.0) + + assert location is not None + assert location["latitude"] == pytest.approx(VALID_LAT) + assert location["timestamp"] == 200.0 + assert "later" in location["age_warning"] diff --git a/tests/unit/test_traceroute_link_endpoint_fix.py b/tests/unit/test_traceroute_link_endpoint_fix.py index 346c3327..4677e885 100644 --- a/tests/unit/test_traceroute_link_endpoint_fix.py +++ b/tests/unit/test_traceroute_link_endpoint_fix.py @@ -1,202 +1,186 @@ -""" -Unit tests for the traceroute link endpoint bug fix. +"""Regression tests for materialized traceroute-link pagination.""" -Tests that the endpoint properly handles RF hops without crashing on missing gateway_node_name. -""" - -import json +import sqlite3 import time -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest - -from src.malla.models.traceroute import TracerouteHop, TraceroutePacket - - -class TestTracerouteLinkEndpointFix: - """Test the fix for the traceroute link endpoint gateway_node_name bug.""" - - @pytest.mark.unit - def test_endpoint_returns_rf_hops_without_gateway_node_name_error(self): +from flask import Flask + +from src.malla.database.traceroute_read_repository import get_traceroute_link +from src.malla.routes.api_routes import register_api_routes + +pytestmark = pytest.mark.unit + + +class _NonClosingConnection: + def __init__(self, connection): + self.connection = connection + + def __getattr__(self, name): + return getattr(self.connection, name) + + def close(self): + pass + + +def _packet(packet_id, timestamp, gateway_id="!0000012c"): + return { + "id": packet_id, + "timestamp": timestamp, + "timestamp_str": "2026-09-10 10:00:00", + "from_node_id": 100, + "to_node_id": 200, + "gateway_id": gateway_id, + "route_nodes_json": "[]", + "snr_towards_json": "[-16.5]", + "route_back_json": "[]", + "snr_back_json": "[]", + "target_hop_snr": -16.5, + } + + +def test_endpoint_paginates_details_but_keeps_full_window_statistics(): + now = time.time() + link_result = { + "packets": [_packet(15, now), _packet(14, now - 1)], + "total_count": 25, + "total_attempts": 25, + "forward_count": 13, + "reverse_count": 12, + "avg_snr": -11.25, + } + + app = Flask(__name__) + register_api_routes(app) + with ( + patch( + "src.malla.routes.api_routes.get_traceroute_link", + return_value=link_result, + ) as query, + patch( + "src.malla.routes.api_routes.NodeRepository.get_bulk_node_names", + return_value={100: "Node A", 200: "Node B", 300: "Gateway"}, + ) as names, + app.test_client() as client, + ): + response = client.get("/api/traceroute/link/100/200?limit=10&page=2") + + assert response.status_code == 200 + data = response.get_json() + assert [row["id"] for row in data["traceroutes"]] == [15, 14] + assert data["page"] == 2 + assert data["limit"] == 10 + assert data["total_count"] == 25 + assert data["total_pages"] == 3 + assert data["total_attempts"] == 25 + assert data["avg_snr"] == -11.25 + assert data["direction_counts"] == { + "Node A → Node B": 13, + "Node B → Node A": 12, + } + assert data["traceroutes"][0]["complete_path_display"] == "Node A" + assert data["traceroutes"][0]["gateway_node_name"] == "Gateway" + assert names.call_count == 1 + assert query.call_args.kwargs["limit"] == 10 + assert query.call_args.kwargs["offset"] == 10 + + +def test_endpoint_returns_empty_paginated_result(): + link_result = { + "packets": [], + "total_count": 0, + "total_attempts": 0, + "forward_count": 0, + "reverse_count": 0, + "avg_snr": None, + } + app = Flask(__name__) + register_api_routes(app) + with ( + patch( + "src.malla.routes.api_routes.get_traceroute_link", + return_value=link_result, + ), + patch( + "src.malla.routes.api_routes.NodeRepository.get_bulk_node_names", + return_value={}, + ), + app.test_client() as client, + ): + response = client.get("/api/traceroute/link/100/200?limit=10&page=2") + + data = response.get_json() + assert response.status_code == 200 + assert data["traceroutes"] == [] + assert data["direction_counts"] == {"forward": 0, "reverse": 0} + assert data["total_count"] == 0 + assert data["total_pages"] == 0 + + +def test_materialized_link_query_aggregates_all_rows_and_pages_packet_details(): + connection = sqlite3.connect(":memory:") + connection.row_factory = sqlite3.Row + connection.executescript( """ - Test that the endpoint returns RF hops between nodes without crashing - on the missing gateway_node_name attribute. - - This is a regression test for the bug where the endpoint tried to access - tr_packet.gateway_node_name which doesn't exist on TraceroutePacket. - """ - - # Mock TracerouteRepository.get_traceroute_packets - mock_packets = [ - { - "id": 12345, - "timestamp": time.time(), - "timestamp_str": "2024-01-20 10:30:00", - "from_node_id": 2510468508, - "to_node_id": 1128074276, - "gateway_id": 3333333333, - "raw_payload": b"fake_payload", - } - ] - - # Mock TraceroutePacket with RF hops between target nodes - mock_traceroute_packet = MagicMock(spec=TraceroutePacket) - mock_traceroute_packet.from_node_name = "Test Node A" - mock_traceroute_packet.to_node_name = "Test Node B" - mock_traceroute_packet.gateway_id = 3333333333 - # Note: gateway_node_name is intentionally NOT set to test the bug fix - mock_traceroute_packet.format_path_display.return_value = "A -> B" - mock_traceroute_packet.get_display_hops.return_value = [] - - # Create a mock RF hop between the target nodes - mock_rf_hop = MagicMock(spec=TracerouteHop) - mock_rf_hop.from_node_id = 2510468508 - mock_rf_hop.to_node_id = 1128074276 - mock_rf_hop.snr = -16.5 - mock_rf_hop.from_node_name = "Test Node A" - mock_rf_hop.to_node_name = "Test Node B" - mock_rf_hop.direction = "forward_rf" - - mock_traceroute_packet.get_rf_hops.return_value = [mock_rf_hop] - - with patch( - "src.malla.routes.api_routes.TracerouteRepository" - ) as mock_repo_class: - mock_repo = mock_repo_class - mock_repo.get_traceroute_packets.return_value = {"packets": mock_packets} - - with patch("src.malla.routes.api_routes.NodeRepository") as mock_node_repo: - mock_node_repo.get_bulk_node_names.return_value = { - 2510468508: "Test Node A", - 1128074276: "Test Node B", - 3333333333: "Gateway Node", - } - - with patch( - "src.malla.routes.api_routes.TraceroutePacket" - ) as mock_traceroute_class: - mock_traceroute_class.return_value = mock_traceroute_packet - - # Import here to use the mocked dependencies - from flask import Flask - - from src.malla.routes.api_routes import register_api_routes - - app = Flask(__name__) - register_api_routes(app) - - with app.test_client() as client: - # Test the endpoint that was previously crashing - response = client.get( - "/api/traceroute/link/2510468508/1128074276" - ) - - # Should not crash and should return valid data - assert response.status_code == 200 - - data = json.loads(response.data) - - # Should have traceroutes (not empty due to the crash) - assert "traceroutes" in data - assert len(data["traceroutes"]) > 0 - - # Should have proper statistics - assert "avg_snr" in data - assert data["avg_snr"] == -16.5 - - # Should have direction counts - assert "direction_counts" in data - - # Verify the traceroute entry structure (with gateway_node_name) - traceroute = data["traceroutes"][0] - expected_fields = { - "id", - "timestamp", - "timestamp_str", - "from_node_id", - "to_node_id", - "from_node_name", - "to_node_name", - "gateway_id", - "gateway_node_name", - "hop_snr", - "route_hops", - "complete_path_display", - } - - for field in expected_fields: - assert field in traceroute, f"Missing field: {field}" - - # Ensure gateway_node_name is properly set - assert traceroute["gateway_node_name"] == "Gateway Node" - - # Verify specific values - assert traceroute["from_node_id"] == 2510468508 - assert traceroute["to_node_id"] == 1128074276 - assert traceroute["hop_snr"] == -16.5 - assert traceroute["gateway_id"] == 3333333333 - - @pytest.mark.unit - def test_endpoint_handles_no_rf_hops_gracefully(self): + CREATE TABLE packet_history ( + id INTEGER PRIMARY KEY, gateway_id TEXT, channel_id TEXT, + hop_start INTEGER, hop_limit INTEGER, rssi REAL, snr REAL, + payload_length INTEGER, processed_successfully INTEGER + ); + CREATE TABLE traceroute_routes ( + packet_id INTEGER PRIMARY KEY, timestamp REAL, from_node_id INTEGER, + to_node_id INTEGER, mesh_packet_id INTEGER, route_nodes_json TEXT, + snr_towards_json TEXT, route_back_json TEXT, snr_back_json TEXT, + parse_status TEXT, parser_version INTEGER + ); + CREATE TABLE traceroute_hops ( + packet_id INTEGER, direction TEXT, hop_index INTEGER, timestamp REAL, + from_node_id INTEGER, to_node_id INTEGER, snr REAL + ); """ - Test that the endpoint returns empty results when no RF hops exist between nodes. - """ - - # Mock TracerouteRepository.get_traceroute_packets - mock_packets = [ - { - "id": 12346, - "timestamp": time.time(), - "timestamp_str": "2024-01-20 10:30:00", - "from_node_id": 1111111111, - "to_node_id": 2222222222, - "gateway_id": 3333333333, - "raw_payload": b"fake_payload", - } - ] - - # Mock TraceroutePacket with NO RF hops between target nodes - mock_traceroute_packet = MagicMock(spec=TraceroutePacket) - mock_traceroute_packet.get_rf_hops.return_value = [] # No RF hops - mock_traceroute_packet.get_display_hops.return_value = [] - - with patch( - "src.malla.routes.api_routes.TracerouteRepository" - ) as mock_repo_class: - mock_repo = mock_repo_class - mock_repo.get_traceroute_packets.return_value = {"packets": mock_packets} - - with patch("src.malla.routes.api_routes.NodeRepository") as mock_node_repo: - mock_node_repo.get_bulk_node_names.return_value = { - 1111111111: "Test Node A", - 2222222222: "Test Node B", - } - - with patch( - "src.malla.routes.api_routes.TraceroutePacket" - ) as mock_traceroute_class: - mock_traceroute_class.return_value = mock_traceroute_packet - - # Import here to use the mocked dependencies - from flask import Flask - - from src.malla.routes.api_routes import register_api_routes - - app = Flask(__name__) - register_api_routes(app) - - with app.test_client() as client: - # Test with nodes that have no RF hops between them - response = client.get( - "/api/traceroute/link/9999999999/8888888888" - ) - - # Should not crash - assert response.status_code == 200 - - data = json.loads(response.data) - - # Should return empty results - assert data["traceroutes"] == [] - assert data["avg_snr"] is None - assert data["direction_counts"] == {"forward": 0, "reverse": 0} + ) + now = time.time() + for packet_id in range(1, 26): + from_node_id, to_node_id = ( + (100, 200) if packet_id % 2 else (200, 100) + ) + connection.execute( + "INSERT INTO packet_history VALUES (?, ?, '', 5, 4, -80, 1, 10, 1)", + (packet_id, "!0000012c"), + ) + connection.execute( + "INSERT INTO traceroute_routes VALUES (?, ?, 100, 200, ?, '[]', ?, '[]', '[]', 'parsed', 1)", + (packet_id, now + packet_id, packet_id, "[-10]"), + ) + connection.execute( + "INSERT INTO traceroute_hops VALUES (?, 'forward', 0, ?, ?, ?, ?)", + (packet_id, now + packet_id, from_node_id, to_node_id, -packet_id), + ) + # Repeated occurrences stay in aggregate hop statistics but must not duplicate + # the packet in the paginated traceroute details. + connection.execute( + "INSERT INTO traceroute_hops VALUES (1, 'return', 1, ?, 200, 100, 5)", + (now + 1,), + ) + connection.commit() + + with patch( + "src.malla.database.traceroute_read_repository.get_db_connection", + return_value=_NonClosingConnection(connection), + ): + result = get_traceroute_link( + 100, + 200, + start_time=now, + end_time=now + 30, + limit=10, + offset=10, + ) + + assert result["total_count"] == 25 + assert result["total_attempts"] == 25 + assert result["forward_count"] == 13 + assert result["reverse_count"] == 13 + assert result["avg_snr"] == pytest.approx((sum(range(-1, -26, -1)) + 5) / 26) + assert [packet["id"] for packet in result["packets"]] == list(range(15, 5, -1)) diff --git a/tests/unit/test_traceroute_materialization.py b/tests/unit/test_traceroute_materialization.py new file mode 100644 index 00000000..21ff64f9 --- /dev/null +++ b/tests/unit/test_traceroute_materialization.py @@ -0,0 +1,471 @@ +"""Behavioral checks for capture, raw-only imports, and resumable preparation.""" + +import json +import sqlite3 +from contextlib import closing +from unittest.mock import patch + +import pytest +from meshtastic import mesh_pb2, mqtt_pb2, portnums_pb2 + +from malla import mqtt_capture +from malla.backfill_traceroutes import main, prepare_traceroutes +from malla.database.schema import ensure_startup_schema +from malla.database.traceroute_schema import ensure_traceroute_schema +from malla.database.traceroutes import ( + PARSER_VERSION, + decode_traceroute, + inspect_traceroutes, + write_traceroute, +) +from malla.models.traceroute import TraceroutePacket + +pytestmark = pytest.mark.unit + + +def payload(route=(), snr=(), back=(), snr_back=()): + return mesh_pb2.RouteDiscovery( + route=route, snr_towards=snr, route_back=back, snr_back=snr_back + ).SerializeToString() + + +@pytest.fixture +def database(tmp_path, monkeypatch): + path = tmp_path / "history.db" + monkeypatch.setattr(mqtt_capture, "DATABASE_FILE", str(path)) + monkeypatch.setattr( + mqtt_capture, "seed_query_planner_stats_async", lambda *_: False + ) + mqtt_capture.init_database() + with closing(sqlite3.connect(path)) as conn: + conn.row_factory = sqlite3.Row + conn.execute("PRAGMA foreign_keys=ON") + yield path, conn + + +def insert_raw(conn, raw, **overrides): + fields = { + "timestamp": 1000.0, + "topic": "msh/test/e/LongFast/!12345678", + "from_node_id": 100, + "to_node_id": 200, + "portnum": 70, + "portnum_name": "TRACEROUTE_APP", + "mesh_packet_id": 42, + "gateway_id": "!12345678", + "hop_start": 5, + "hop_limit": 3, + "raw_payload": raw, + **overrides, + } + cursor = conn.execute( + f"INSERT INTO packet_history ({', '.join(fields)}) " + f"VALUES ({', '.join('?' for _ in fields)})", + list(fields.values()), + ) + return dict( + conn.execute( + "SELECT * FROM packet_history WHERE id = ?", (cursor.lastrowid,) + ).fetchone() + ) + + +def capture(raw, gateway="!12345678"): + packet = mesh_pb2.MeshPacket(id=42, to=200, hop_start=5, hop_limit=3) + setattr(packet, "from", 100) + packet.decoded.portnum = portnums_pb2.PortNum.TRACEROUTE_APP + packet.decoded.payload = raw + envelope = mqtt_pb2.ServiceEnvelope(gateway_id=gateway, channel_id="LongFast") + mqtt_capture.log_packet_to_database("msh/test/e/LongFast", envelope, packet) + + +@pytest.mark.parametrize( + "raw", + [ + payload(snr=[-16]), + payload(route=[110], snr=[-20]), + payload(route=[110], snr=[-20, -32], back=[120], snr_back=[4, -12]), + payload(route=[110, 100, 110], snr=[4, 8, 12, 16]), + payload(route=[110, 120]), + payload(snr=[-128]), + b"", + ], +) +def test_decoder_preserves_existing_path_rules_without_lookups(raw): + packet = { + "from_node_id": 100, + "to_node_id": 200, + "hop_start": 5, + "hop_limit": 3, + "raw_payload": raw, + } + expected = TraceroutePacket(packet, resolve_names=False) + with patch( + "malla.models.traceroute.TraceroutePacket._resolve_node_names", + side_effect=AssertionError("no lookup"), + ): + decoded = decode_traceroute(packet) + assert decoded.route == expected.route_data + assert decoded.hops == tuple(expected.get_rf_hops()) + assert decoded.forward_complete == expected.is_complete() + assert decoded.return_complete == expected.is_return_complete() + + +def test_capture_preserves_repeated_hops_and_gateway_receptions(database): + _, conn = database + raw = payload(route=[110, 100, 110], snr=[4, 8, 12, 16]) + capture(raw) + capture(raw, gateway="!87654321") + routes = conn.execute( + "SELECT * FROM traceroute_routes ORDER BY packet_id" + ).fetchall() + assert len(routes) == 2 + assert routes[0]["mesh_packet_id"] == routes[1]["mesh_packet_id"] == 42 + for route in routes: + assert route["parse_status"] == "parsed" + assert json.loads(route["route_nodes_json"]) == [110, 100, 110] + hops = conn.execute( + "SELECT hop_index, from_node_id, to_node_id, snr FROM traceroute_hops " + "WHERE packet_id = ? ORDER BY hop_index", + (route["packet_id"],), + ).fetchall() + assert [tuple(hop) for hop in hops] == [ + (0, 200, 110, 1.0), + (1, 110, 100, 2.0), + (2, 100, 110, 3.0), + (3, 110, 100, 4.0), + ] + + +def test_capture_and_backfill_are_identical(database): + _, conn = database + raw = payload(route=[110], snr=[-20, -32], back=[120], snr_back=[4, -128]) + capture(raw) + before_route = dict(conn.execute("SELECT * FROM traceroute_routes").fetchone()) + before_hops = [ + tuple(row) + for row in conn.execute( + "SELECT * FROM traceroute_hops ORDER BY direction, hop_index" + ) + ] + with conn: + conn.execute("DELETE FROM traceroute_routes") + result = prepare_traceroutes(conn) + after_route = dict(conn.execute("SELECT * FROM traceroute_routes").fetchone()) + before_route.pop("materialized_at") + after_route.pop("materialized_at") + assert before_route == after_route + assert before_hops == [ + tuple(row) + for row in conn.execute( + "SELECT * FROM traceroute_hops ORDER BY direction, hop_index" + ) + ] + assert result["complete"] + assert result["processed_this_run"] == 1 + + +@pytest.mark.parametrize( + "raw,status", [(b"\xff", "invalid_payload"), (b"", "valid_empty")] +) +def test_capture_keeps_malformed_and_empty_packets_distinct(database, raw, status): + _, conn = database + capture(raw) + assert conn.execute("SELECT raw_payload FROM packet_history").fetchone()[0] == raw + route = conn.execute("SELECT * FROM traceroute_routes").fetchone() + assert route["parse_status"] == status + assert bool(route["parse_error"]) == (status == "invalid_payload") + assert conn.execute("SELECT COUNT(*) FROM traceroute_hops").fetchone()[0] == 0 + assert prepare_traceroutes(conn)["processed_this_run"] == 0 + + +def test_capture_storage_failure_rolls_back_raw_and_all_derived_rows(database): + _, conn = database + with conn: + conn.execute(""" + CREATE TRIGGER reject_second_hop BEFORE INSERT ON traceroute_hops + WHEN NEW.hop_index = 1 BEGIN SELECT RAISE(ABORT, 'test storage failure'); END + """) + with pytest.raises(sqlite3.IntegrityError, match="test storage failure"): + capture(payload(route=[110], snr=[4, 8])) + for table in ("packet_history", "traceroute_routes", "traceroute_hops"): + assert conn.execute(f"SELECT COUNT(*) FROM {table}").fetchone()[0] == 0 + # The failed capture connection was closed and its lock released. + capture(payload(snr=[4])) + assert conn.execute("SELECT COUNT(*) FROM traceroute_routes").fetchone()[0] == 1 + + +def test_pending_imports_and_decoder_versions_affect_completeness(database): + _, conn = database + assert inspect_traceroutes(conn.cursor())["complete"] + with conn: + packet = insert_raw(conn, payload(snr=[4])) + assert not inspect_traceroutes(conn.cursor())["complete"] + assert prepare_traceroutes(conn)["complete"] + assert inspect_traceroutes(conn.cursor())["complete"] + with conn: + packet = insert_raw(conn, payload(snr=[4])) + assert not inspect_traceroutes(conn.cursor())["complete"] + with conn: + conn.execute("BEGIN") + write_traceroute(conn.cursor(), packet) + assert inspect_traceroutes(conn.cursor())["complete"] + with patch("malla.database.traceroutes.PARSER_VERSION", PARSER_VERSION + 1): + assert not inspect_traceroutes(conn.cursor())["complete"] + capture(payload(snr=[8])) + assert inspect_traceroutes(conn.cursor())["complete"] + + +def test_raw_only_updates_and_replacements_remove_stale_hops(database): + _, conn = database + capture(payload(route=[110], snr=[4, 8])) + assert prepare_traceroutes(conn)["complete"] + with conn: + conn.execute( + "UPDATE packet_history SET raw_payload = ?", (payload(route=[120]),) + ) + assert not inspect_traceroutes(conn.cursor())["complete"] + assert conn.execute("SELECT COUNT(*) FROM traceroute_hops").fetchone()[0] == 0 + assert prepare_traceroutes(conn)["complete"] + assert json.loads( + conn.execute("SELECT route_nodes_json FROM traceroute_routes").fetchone()[0] + ) == [120] + with conn: + conn.execute( + "INSERT OR REPLACE INTO packet_history (id, timestamp, topic, portnum) VALUES (1, 1000, 'test', 1)" + ) + assert conn.execute("SELECT COUNT(*) FROM traceroute_routes").fetchone()[0] == 0 + + +@pytest.mark.parametrize("foreign_keys", [True, False]) +def test_raw_deletion_cleans_derived_rows_even_for_raw_only_importers( + database, foreign_keys +): + _, conn = database + capture(payload(snr=[4])) + conn.execute(f"PRAGMA foreign_keys={'ON' if foreign_keys else 'OFF'}") + with conn: + conn.execute("DELETE FROM packet_history") + assert conn.execute("SELECT COUNT(*) FROM traceroute_routes").fetchone()[0] == 0 + assert conn.execute("SELECT COUNT(*) FROM traceroute_hops").fetchone()[0] == 0 + + +def test_retention_deletes_route_and_hops(database, monkeypatch): + _, conn = database + with patch("malla.mqtt_capture.time.time", return_value=1000.0): + capture(payload(snr=[4])) + monkeypatch.setattr(mqtt_capture, "DATA_RETENTION_HOURS", 1) + mqtt_capture.cleanup_old_data() + for table in ("packet_history", "traceroute_routes", "traceroute_hops"): + assert conn.execute(f"SELECT COUNT(*) FROM {table}").fetchone()[0] == 0 + + +def test_backfill_batches_resume_without_retrying_invalid_packets(database): + _, conn = database + with conn: + for raw in ( + payload(snr=[4]), + b"\xff", + payload(route=[110]), + b"", + payload(snr=[8]), + ): + insert_raw(conn, raw) + insert_raw(conn, b"ordinary text", portnum=1, portnum_name="TEXT_MESSAGE_APP") + + def interrupt(_): + raise KeyboardInterrupt + + with pytest.raises(KeyboardInterrupt): + prepare_traceroutes(conn, batch_size=2, progress=interrupt) + assert ( + conn.execute( + "SELECT COUNT(*) FROM traceroute_routes WHERE parser_version = ?", + (PARSER_VERSION,), + ).fetchone()[0] + == 2 + ) + assert not inspect_traceroutes(conn.cursor())["complete"] + result = prepare_traceroutes(conn, batch_size=2) + assert result["raw_traceroutes"] == result["routes"] == 5 + assert result["invalid_payload"] == 1 + assert result["processed_this_run"] == 3 + assert result["complete"] + rows = [ + tuple(row) + for row in conn.execute("SELECT * FROM traceroute_routes ORDER BY packet_id") + ] + assert prepare_traceroutes(conn)["processed_this_run"] == 0 + assert rows == [ + tuple(row) + for row in conn.execute("SELECT * FROM traceroute_routes ORDER BY packet_id") + ] + + +def test_backfill_failed_batch_is_atomic_and_can_resume(database): + _, conn = database + with conn: + for _ in range(3): + insert_raw(conn, payload(snr=[4])) + conn.execute(""" + CREATE TRIGGER reject_packet BEFORE INSERT ON traceroute_hops + WHEN NEW.packet_id = 2 BEGIN SELECT RAISE(ABORT, 'batch failure'); END + """) + with pytest.raises(sqlite3.IntegrityError, match="batch failure"): + prepare_traceroutes(conn, batch_size=3) + assert ( + conn.execute( + "SELECT COUNT(*) FROM traceroute_routes WHERE parse_status = 'pending'" + ).fetchone()[0] + == 3 + ) + assert conn.execute("SELECT COUNT(*) FROM traceroute_hops").fetchone()[0] == 0 + with conn: + conn.execute("DROP TRIGGER reject_packet") + assert prepare_traceroutes(conn)["complete"] + + +def test_capture_can_commit_between_backfill_batches(database): + _, conn = database + with conn: + for _ in range(4): + insert_raw(conn, payload(snr=[4])) + result = prepare_traceroutes( + conn, batch_size=2, progress=lambda _: capture(payload(snr=[8])) + ) + assert result["processed_this_run"] == 4 + assert result["routes"] == result["raw_traceroutes"] == 6 + assert result["complete"] + + +def test_raw_only_import_during_backfill_cannot_look_complete(database): + path, conn = database + with conn: + insert_raw(conn, payload(snr=[4])) + + def raw_import(_): + with closing(sqlite3.connect(path)) as other, other: + other.row_factory = sqlite3.Row + insert_raw(other, payload(snr=[8])) + + result = prepare_traceroutes(conn, progress=raw_import) + assert not result["complete"] + assert result["pending"] == 1 + assert prepare_traceroutes(conn)["complete"] + + +def test_missing_endpoints_are_not_coerced_and_missing_payload_is_reported(database): + _, conn = database + with conn: + insert_raw(conn, None, mesh_packet_id=None, from_node_id=None, to_node_id=None) + insert_raw( + conn, + payload(route=[110]), + mesh_packet_id=None, + from_node_id=None, + to_node_id=None, + ) + result = prepare_traceroutes(conn) + assert result["invalid_payload"] == 1 + assert result["parsed"] == 1 + assert result["hops"] == 0 + for row in conn.execute("SELECT * FROM traceroute_routes"): + assert row["from_node_id"] is row["to_node_id"] is row["mesh_packet_id"] is None + + +def test_version_downgrade_is_rejected(database): + _, conn = database + capture(payload(snr=[4])) + with conn: + conn.execute( + "UPDATE traceroute_routes SET parser_version = ?", (PARSER_VERSION + 1,) + ) + with pytest.raises(ValueError, match="newer traceroute decoder"): + prepare_traceroutes(conn) + assert not inspect_traceroutes(conn.cursor())["complete"] + + +def test_cli_requires_explicit_database_and_rejects_missing_path(tmp_path): + with pytest.raises(SystemExit) as exc: + main([]) + assert exc.value.code == 2 + path = tmp_path / "missing.db" + assert main(["--database", str(path)]) == 1 + assert not path.exists() + + +def test_cli_check_does_not_create_schema_or_change_history(database, capsys): + path, conn = database + with conn: + for event in ("insert", "update", "delete"): + conn.execute(f"DROP TRIGGER traceroute_packet_{event}") + conn.execute("DROP TABLE traceroute_hops") + conn.execute("DROP TABLE traceroute_routes") + insert_raw(conn, payload(snr=[4])) + before = list(conn.iterdump()) + assert main(["--database", str(path), "--check"]) == 1 + assert list(conn.iterdump()) == before + output = capsys.readouterr().out + assert str(path.resolve()) in output + assert '"missing": 1' in output + # Shared startup creates tables and pending triggers, but does not backfill. + with conn: + ensure_startup_schema(conn.cursor()) + assert conn.execute("SELECT COUNT(*) FROM traceroute_routes").fetchone()[0] == 0 + assert not inspect_traceroutes(conn.cursor())["complete"] + assert main(["--database", str(path), "--batch-size", "1"]) == 0 + assert inspect_traceroutes(conn.cursor())["complete"] + + +def test_validation_reports_missing_and_orphaned_rows(database): + _, conn = database + capture(payload(snr=[4])) + conn.execute("PRAGMA foreign_keys=OFF") + with conn: + conn.execute("DELETE FROM traceroute_routes") + ensure_traceroute_schema(conn.cursor()) + with conn: + conn.execute("BEGIN") + result = inspect_traceroutes(conn.cursor()) + assert result["missing"] == 1 + assert result["orphan_hops"] == 1 + assert not result["complete"] + + +def test_clearing_derived_history_invalidates_completeness(database): + _, conn = database + capture(payload(snr=[4])) + assert prepare_traceroutes(conn)["complete"] + with conn: + conn.execute("DELETE FROM traceroute_routes") + assert not inspect_traceroutes(conn.cursor())["complete"] + assert prepare_traceroutes(conn)["complete"] + with conn: + conn.execute("DELETE FROM packet_history") + assert inspect_traceroutes(conn.cursor())["complete"] + + +def test_backfill_includes_explicit_nonpositive_ids_and_either_port_field(database): + _, conn = database + with conn: + insert_raw(conn, payload(snr=[4]), id=-1, portnum=None) + insert_raw(conn, payload(snr=[4]), id=0, portnum_name=None) + result = prepare_traceroutes(conn, batch_size=1) + assert result["processed_this_run"] == 2 + assert result["complete"] + + +def test_new_decoder_reprepares_old_records(database, monkeypatch): + _, conn = database + capture(payload(snr=[4])) + assert prepare_traceroutes(conn)["complete"] + monkeypatch.setattr("malla.database.traceroutes.PARSER_VERSION", PARSER_VERSION + 1) + monkeypatch.setattr("malla.backfill_traceroutes.PARSER_VERSION", PARSER_VERSION + 1) + assert not inspect_traceroutes(conn.cursor())["complete"] + result = prepare_traceroutes(conn) + assert result["complete"] + assert result["processed_this_run"] == 1 + assert ( + conn.execute("SELECT parser_version FROM traceroute_routes").fetchone()[0] + == PARSER_VERSION + 1 + ) diff --git a/tests/unit/test_traceroute_read_repository.py b/tests/unit/test_traceroute_read_repository.py new file mode 100644 index 00000000..afd56f3f --- /dev/null +++ b/tests/unit/test_traceroute_read_repository.py @@ -0,0 +1,272 @@ +"""Behavioral tests for PR2's saved-traceroute list reader.""" + +import sqlite3 +from contextlib import closing +from unittest.mock import patch + +import pytest +from meshtastic import mesh_pb2 + +from malla.database.repositories import LocationRepository +from malla.database.traceroute_read_repository import ( + get_node_traceroute_statistics, + get_route_patterns_data, + get_traceroute_hops_for_graph, + get_traceroute_hops_for_longest_links, + get_traceroute_packets, +) +from malla.database.traceroute_schema import ensure_traceroute_schema +from malla.database.traceroutes import write_traceroute + +pytestmark = pytest.mark.unit + + +@pytest.fixture +def database(tmp_path): + path = tmp_path / "reader.db" + with closing(sqlite3.connect(path)) as conn: + conn.row_factory = sqlite3.Row + conn.execute(""" + CREATE TABLE packet_history ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + timestamp REAL NOT NULL, + portnum INTEGER, + portnum_name TEXT, + mesh_packet_id INTEGER, + from_node_id INTEGER, + to_node_id INTEGER, + gateway_id TEXT, + channel_id TEXT, + hop_start INTEGER, + hop_limit INTEGER, + rssi REAL, + snr REAL, + payload_length INTEGER, + raw_payload BLOB, + processed_successfully INTEGER DEFAULT 1 + ) + """) + ensure_traceroute_schema(conn.cursor()) + conn.commit() + return path + + +def _connection(path): + conn = sqlite3.connect(path) + conn.row_factory = sqlite3.Row + conn.execute("PRAGMA foreign_keys=ON") + return conn + + +def _insert(conn, *, packet_id, timestamp, mesh_id, gateway, route=(900,), snr_towards=None): + if snr_towards is None: + snr_towards = [-40] * (len(route) + 1) + raw = mesh_pb2.RouteDiscovery( + route=route, snr_towards=snr_towards + ).SerializeToString() + cursor = conn.execute( + """ + INSERT INTO packet_history ( + id, timestamp, portnum, portnum_name, mesh_packet_id, + from_node_id, to_node_id, gateway_id, channel_id, + hop_start, hop_limit, rssi, snr, payload_length, raw_payload + ) VALUES (?, ?, 70, 'TRACEROUTE_APP', ?, 100, 200, ?, 'LongFast', + 5, 3, -80, -10, ?, ?) + """, + (packet_id, timestamp, mesh_id, gateway, len(raw), raw), + ) + packet = dict( + conn.execute("SELECT * FROM packet_history WHERE id = ?", (cursor.lastrowid,)).fetchone() + ) + write_traceroute(conn.cursor(), packet) + + +def test_empty_database_returns_empty_packets(database): + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + result = get_traceroute_packets() + + assert result["total_count"] == 0 + assert result["packets"] == [] + + +def test_grouping_filters_full_history_before_exact_pagination(database): + with closing(_connection(database)) as conn: + # More receptions than the old grouped-reader scan cap, with an older + # matching route that must remain reachable on a later page. + for packet_id in range(1, 101): + _insert( + conn, + packet_id=packet_id, + timestamp=packet_id, + mesh_id=packet_id, + gateway="!00000001", + route=(900 if packet_id == 1 else 901,), + ) + # A second reception of one mesh packet must change its reception count, + # not the number of grouped rows. + _insert( + conn, + packet_id=101, + timestamp=101, + mesh_id=1, + gateway="!00000002", + route=(900,), + ) + conn.commit() + + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + page = get_traceroute_packets( + limit=10, + offset=90, + filters={"start_time": 0.0, "end_time": 200.0}, + group_packets=True, + ) + route_match = get_traceroute_packets( + limit=10, + filters={"start_time": 0.0, "end_time": 200.0, "route_node": 900}, + group_packets=True, + ) + gateway_sorted = get_traceroute_packets( + limit=-1, + filters={"start_time": 0.0, "end_time": 200.0}, + order_by="gateway_id", + order_dir="asc", + group_packets=True, + ) + + assert page["total_count"] == 100 + assert len(page["packets"]) == 10 + assert route_match["total_count"] == 1 + assert route_match["packets"][0]["mesh_packet_id"] == 1 + assert route_match["packets"][0]["reception_count"] == 2 + assert route_match["packets"][0]["gateway_count"] == 2 + assert gateway_sorted["packets"][0]["gateway_count"] == 1 + assert gateway_sorted["packets"][-1]["gateway_count"] == 2 + + +def test_get_traceroute_hops_for_graph_and_longest_links(database): + with closing(_connection(database)) as conn: + _insert(conn, packet_id=1, timestamp=10.0, mesh_id=101, gateway="!00000001", route=(901, 902)) + conn.commit() + + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + hops = get_traceroute_hops_for_graph(filters={"start_time": 0.0, "end_time": 20.0}) + longest_hops = get_traceroute_hops_for_longest_links(start_time=0.0, end_time=20.0) + + assert len(hops) == 3 + assert [h["from_node_id"] for h in hops] == [100, 901, 902] + assert [h["to_node_id"] for h in hops] == [901, 902, 200] + assert len(longest_hops) == 3 + + +def test_hop_query_preserves_zero_snr_hops_in_path_order(database): + """A zero-SNR middle hop stays in the sequence so paths keep their structure. + + For A->B->C->D the stored hops are A->B, B->C, C->D. Dropping B->C (SNR 0) + in SQL would leave two disconnected segments that downstream consumers + would splice into a bogus two-hop A->D path. + """ + with closing(_connection(database)) as conn: + _insert( + conn, + packet_id=1, + timestamp=10.0, + mesh_id=101, + gateway="!00000001", + route=(901, 902), + snr_towards=[-160, 0, -160], + ) + conn.commit() + + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + hops = get_traceroute_hops_for_graph(filters={"start_time": 0.0, "end_time": 20.0}) + longest_hops = get_traceroute_hops_for_longest_links(start_time=0.0, end_time=20.0) + + for query_hops in (hops, longest_hops): + assert [h["from_node_id"] for h in query_hops] == [100, 901, 902] + assert [h["to_node_id"] for h in query_hops] == [901, 902, 200] + assert [h["snr"] for h in query_hops] == [-40.0, 0.0, -40.0] + + +def test_get_route_patterns_data(database): + with closing(_connection(database)) as conn: + for pid in range(1, 4): + _insert(conn, packet_id=pid, timestamp=float(pid), mesh_id=pid, gateway="!00000001", route=(905,)) + conn.commit() + + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + patterns_result = get_route_patterns_data(start_time=0.0, end_time=10.0, limit=5) + + assert patterns_result["total_patterns"] == 1 + assert patterns_result["analyzed_traceroutes"] == 3 + assert len(patterns_result["sorted_patterns"]) == 1 + pattern_key, pattern_data = patterns_result["sorted_patterns"][0] + assert pattern_data["count"] == 3 + assert pattern_key[1] == (905,) + + +def test_get_node_traceroute_statistics(database): + with closing(_connection(database)) as conn: + _insert(conn, packet_id=1, timestamp=10.0, mesh_id=101, gateway="!00000001", route=(900,)) + conn.commit() + + with patch( + "malla.database.traceroute_read_repository.get_db_connection", + side_effect=lambda: _connection(database), + ): + source_stats = get_node_traceroute_statistics(node_id=100) + dest_stats = get_node_traceroute_statistics(node_id=200) + intermediate_stats = get_node_traceroute_statistics(node_id=900) + + assert source_stats["as_source"]["total"] == 1 + assert source_stats["as_source"]["successful"] == 1 + assert dest_stats["as_destination"]["total"] == 1 + assert dest_stats["as_destination"]["successful"] == 1 + assert intermediate_stats["as_intermediate_hop"]["participation_count"] == 1 + + +def test_get_nodes_location_history(database): + pos = mesh_pb2.Position(latitude_i=400000000, longitude_i=-300000000, altitude=150) + pos_bytes = pos.SerializeToString() + + with closing(_connection(database)) as conn: + conn.execute( + """ + INSERT INTO packet_history ( + id, timestamp, portnum, portnum_name, from_node_id, to_node_id, + gateway_id, payload_length, raw_payload + ) VALUES (500, 100.0, 3, 'POSITION_APP', 100, 4294967295, '!00000001', ?, ?) + """, + (len(pos_bytes), pos_bytes), + ) + conn.commit() + + with patch( + "malla.database.repositories.get_db_connection", + side_effect=lambda: _connection(database), + ): + locs = LocationRepository.get_nodes_location_history([100, 200], limit_per_node=5) + + assert 100 in locs + assert len(locs[100]) == 1 + assert locs[100][0]["latitude"] == 40.0 + assert locs[100][0]["longitude"] == -30.0 + assert locs[100][0]["altitude"] == 150 + assert 200 in locs + assert len(locs[200]) == 0 + diff --git a/tests/unit/test_traceroute_service.py b/tests/unit/test_traceroute_service.py index 768262a2..cd444a8b 100644 --- a/tests/unit/test_traceroute_service.py +++ b/tests/unit/test_traceroute_service.py @@ -5,7 +5,7 @@ """ from datetime import datetime -from unittest.mock import Mock, patch +from unittest.mock import patch from src.malla.services.traceroute_service import TracerouteService @@ -13,53 +13,51 @@ class TestTracerouteServiceLongestLinks: """Test TracerouteService longest links analysis functionality.""" - @patch("src.malla.services.traceroute_service.TracerouteRepository") - @patch("src.malla.services.traceroute_service.TraceroutePacket") - def test_longest_links_analysis_basic(self, mock_traceroute_packet, mock_repo): + @patch("src.malla.services.traceroute_service.get_bulk_node_names") + @patch("src.malla.services.traceroute_service.LocationRepository.get_nodes_location_history") + @patch("src.malla.services.traceroute_service.get_traceroute_hops_for_longest_links") + def test_longest_links_analysis_basic( + self, mock_get_hops, mock_get_locs, mock_get_names + ): """Test basic longest links analysis functionality.""" - # Mock repository response - mock_packet_data = { - "id": 1, + now_ts = datetime.now().timestamp() + mock_hop = { + "packet_id": 1, + "direction": "forward", + "hop_index": 0, + "timestamp": now_ts, "from_node_id": 100, "to_node_id": 200, - "timestamp": datetime.now().timestamp(), - "gateway_id": "!12345678", - "raw_payload": b"mock_payload", - "processed_successfully": True, + "snr": -5.0, } - - mock_repo.get_traceroute_packets.return_value = {"packets": [mock_packet_data]} - - # Mock TraceroutePacket - mock_packet = Mock() - mock_packet.from_node_id = 100 - mock_packet.to_node_id = 200 - - # Mock RF hop - mock_hop = Mock() - mock_hop.from_node_id = 100 - mock_hop.to_node_id = 200 - mock_hop.from_node_name = "Node100" - mock_hop.to_node_name = "Node200" - mock_hop.distance_km = 5.0 # 5km - mock_hop.snr = -5.0 - - mock_packet.get_rf_hops.return_value = [mock_hop] - mock_packet.get_display_hops.return_value = [mock_hop] - mock_packet.calculate_hop_distances = Mock() - - mock_traceroute_packet.return_value = mock_packet + mock_get_hops.return_value = [mock_hop] + mock_get_locs.return_value = { + 100: [ + { + "from_node_id": 100, + "latitude": 40.0, + "longitude": -3.0, + "altitude": 100, + "timestamp": now_ts, + } + ], + 200: [ + { + "from_node_id": 200, + "latitude": 40.045, + "longitude": -3.0, + "altitude": 100, + "timestamp": now_ts, + } + ], + } + mock_get_names.return_value = {100: "Node100", 200: "Node200"} # Call the method result = TracerouteService.get_longest_links_analysis( min_distance_km=1.0, min_snr=-10.0, max_results=10 ) - # Verify TraceroutePacket was called with correct arguments - mock_traceroute_packet.assert_called_with( - packet_data=mock_packet_data, resolve_names=True - ) - # Verify structure assert "summary" in result assert "direct_links" in result @@ -67,25 +65,24 @@ def test_longest_links_analysis_basic(self, mock_traceroute_packet, mock_repo): # Verify summary summary = result["summary"] - assert "total_links" in summary - assert "direct_links" in summary - assert "longest_direct" in summary - assert "longest_path" in summary + assert summary["total_links"] == 1 + assert summary["direct_links"] == 1 + assert summary["longest_direct"] is not None + assert summary["longest_path"] is None # Verify direct links assert len(result["direct_links"]) == 1 direct_link = result["direct_links"][0] assert direct_link["from_node_id"] == 100 assert direct_link["to_node_id"] == 200 - assert direct_link["distance_km"] == 5.0 + assert direct_link["distance_km"] > 4.0 assert direct_link["avg_snr"] == -5.0 assert direct_link["traceroute_count"] == 1 - @patch("src.malla.services.traceroute_service.TracerouteRepository") - def test_longest_links_analysis_empty_data(self, mock_repo): + @patch("src.malla.services.traceroute_service.get_traceroute_hops_for_longest_links") + def test_longest_links_analysis_empty_data(self, mock_get_hops): """Test analysis with no traceroute data.""" - # Mock empty repository response - mock_repo.get_traceroute_packets.return_value = {"packets": []} + mock_get_hops.return_value = [] # Call the method result = TracerouteService.get_longest_links_analysis() @@ -97,3 +94,146 @@ def test_longest_links_analysis_empty_data(self, mock_repo): assert result["summary"]["longest_path"] is None assert len(result["direct_links"]) == 0 assert len(result["indirect_links"]) == 0 + + @staticmethod + def _hop(packet_id, hop_index, from_node, to_node, snr, timestamp): + return { + "packet_id": packet_id, + "direction": "forward", + "hop_index": hop_index, + "timestamp": timestamp, + "from_node_id": from_node, + "to_node_id": to_node, + "snr": snr, + } + + @staticmethod + def _linear_locations(node_positions, timestamp): + return { + node_id: [ + { + "from_node_id": node_id, + "latitude": lat, + "longitude": lon, + "altitude": 100, + "timestamp": timestamp, + } + ] + for node_id, (lat, lon) in node_positions.items() + } + + @patch("src.malla.services.traceroute_service.get_bulk_node_names") + @patch("src.malla.services.traceroute_service.LocationRepository.get_nodes_location_history") + @patch("src.malla.services.traceroute_service.get_traceroute_hops_for_longest_links") + def test_longest_links_broken_path_not_aggregated( + self, mock_get_hops, mock_get_locs, mock_get_names + ): + """A zero-SNR middle hop splits A->B->C->D; segments must not join as A->D.""" + now_ts = datetime.now().timestamp() + mock_get_hops.return_value = [ + self._hop(1, 0, 100, 200, -5.0, now_ts), + self._hop(1, 1, 200, 300, 0.0, now_ts), + self._hop(1, 2, 300, 400, -5.0, now_ts), + ] + mock_get_locs.return_value = self._linear_locations( + { + 100: (40.000, -3.0), + 200: (40.045, -3.0), + 300: (40.090, -3.0), + 400: (40.135, -3.0), + }, + now_ts, + ) + mock_get_names.return_value = {nid: f"Node{nid}" for nid in (100, 200, 300, 400)} + + result = TracerouteService.get_longest_links_analysis( + min_distance_km=1.0, min_snr=-30.0, max_results=10 + ) + + direct = {(link["from_node_id"], link["to_node_id"]) for link in result["direct_links"]} + assert direct == {(100, 200), (300, 400)} + assert result["indirect_links"] == [] + assert result["summary"]["longest_path"] is None + assert result["summary"]["longest_direct"] is not None + + @patch("src.malla.services.traceroute_service.get_bulk_node_names") + @patch("src.malla.services.traceroute_service.LocationRepository.get_nodes_location_history") + @patch("src.malla.services.traceroute_service.get_traceroute_hops_for_longest_links") + def test_longest_links_contiguous_path_aggregated( + self, mock_get_hops, mock_get_locs, mock_get_names + ): + """A fully evidenced A->B->C path still aggregates with correct hops/preview.""" + now_ts = datetime.now().timestamp() + mock_get_hops.return_value = [ + self._hop(1, 0, 100, 200, -5.0, now_ts), + self._hop(1, 1, 200, 300, -5.0, now_ts), + ] + mock_get_locs.return_value = self._linear_locations( + {100: (40.000, -3.0), 200: (40.045, -3.0), 300: (40.090, -3.0)}, + now_ts, + ) + mock_get_names.return_value = {nid: f"Node{nid}" for nid in (100, 200, 300)} + + result = TracerouteService.get_longest_links_analysis( + min_distance_km=1.0, min_snr=-10.0, max_results=10 + ) + + assert len(result["indirect_links"]) == 1 + path = result["indirect_links"][0] + assert (path["from_node_id"], path["to_node_id"]) == (100, 300) + assert path["hop_count"] == 2 + assert path["route_preview"] == ["Node100", "Node200", "Node300"] + assert path["total_distance_km"] > 9.0 + assert path["avg_snr"] == -5.0 + + @patch("src.malla.services.traceroute_service.LocationRepository.get_node_locations") + @patch("src.malla.services.traceroute_service.get_bulk_node_names") + @patch("src.malla.services.traceroute_service.get_traceroute_hops_for_graph") + def test_network_graph_broken_path_no_indirect( + self, mock_get_hops, mock_get_names, mock_get_locs + ): + """Graph indirect connections require continuity: no fake A->D shortcut.""" + from src.malla.services.traceroute_service import _NETWORK_GRAPH_CACHE + + _NETWORK_GRAPH_CACHE.clear() + now_ts = datetime.now().timestamp() + mock_get_hops.return_value = [ + self._hop(1, 0, 100, 200, -5.0, now_ts), + self._hop(1, 1, 200, 300, 0.0, now_ts), + self._hop(1, 2, 300, 400, -5.0, now_ts), + ] + mock_get_names.return_value = {nid: f"Node{nid}" for nid in (100, 200, 300, 400)} + mock_get_locs.return_value = [] + + try: + result = TracerouteService.get_network_graph_data( + hours=24, + min_snr=-200.0, + include_indirect=True, + filters={"start_time": now_ts - 60, "end_time": now_ts + 60}, + ) + finally: + _NETWORK_GRAPH_CACHE.clear() + + direct = {(link["source"], link["target"]) for link in result["links"]} + assert direct == {(100, 200), (300, 400)} + assert result["indirect_connections"] == [] + assert result["stats"]["links_filtered_due_to_snr_0"] == 1 + + @patch("src.malla.services.traceroute_service.get_bulk_node_names") + @patch("src.malla.services.traceroute_service.get_node_traceroute_statistics") + def test_node_traceroute_stats(self, mock_get_stats, mock_get_names): + """Test node traceroute stats delegates to SQL statistics.""" + mock_get_stats.return_value = { + "node_id": 12345, + "as_source": {"total": 10, "successful": 8, "success_rate": 80.0}, + "as_destination": {"total": 5, "successful": 4, "success_rate": 80.0}, + "as_intermediate_hop": {"participation_count": 3}, + "total_involvement": 18, + } + mock_get_names.return_value = {12345: "TestNode"} + + stats = TracerouteService.get_node_traceroute_stats(12345) + assert stats["node_id"] == 12345 + assert stats["node_name"] == "TestNode" + assert stats["total_involvement"] == 18