#!/opt/meshtastic-bridge/venv/bin/python3
import paho.mqtt.client as mqtt
import base64
import sys
import logging
import ssl
import signal
import struct

# Import Protobufs and Crypto primitives
from meshtastic.protobuf import mqtt_pb2, mesh_pb2, portnums_pb2
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from cryptography.hazmat.backends import default_backend

# ================= MQTT CONFIG =================
# Local Broker Settings (Gainesville Mesh)
LOCAL_BROKER = "mqtt.gville-swamp.com"
LOCAL_PORT = 8883
LOCAL_USER = "<User_Name>"
LOCAL_PASS = "<password>"
LOCAL_TOPIC = "msh/swamp/#"

# Florida Mesh Settings (FL Mesh)
FL_BROKER = "mqtt.areyoumeshingwith.us"
FL_PORT = 1883
FL_USER = "<User_Name>"
FL_PASS = "<password>"

# Local Channel PSK (Base64)
LOCAL_PSK_BASE64 = "AQ==" 

# ================= JournalCtl Logging =================
log_level = logging.INFO
# Unccoment the line below to enable debug logs
# log_level = logging.DEBUG

logging.basicConfig(level=log_level, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)

# ============= Cleanup ================
def cleanup(signum, frame):
    logger.info("Signal received, shutting down...")
    try:
        local_client.disconnect()
        fl_client.disconnect()
    except Exception as e:
        logger.error(f"Disconnect error: {e}")
    sys.exit(0)

# ================= Decrypt Function =================
def decrypt_packet(mp, key_base64):
    """Decrypts a Meshtastic packet using AES-CTR."""
    try:
        key_bytes = base64.b64decode(key_base64)
        
        # Get IDs to build number used once - force uint_32
        packet_id = getattr(mp, "id", 0) & 0xFFFFFFFF
        from_node = getattr(mp, "from", 0) & 0xFFFFFFFF
        
        # Docs CryptoEngine::initNonce(uint32_t fromNode, uint64_t packetId) 
        nonce = struct.pack('<Q', packet_id) 
        nonce += struct.pack('<I', from_node)
        
        cipher = Cipher(algorithms.AES(key_bytes), modes.CTR(nonce), backend=default_backend())
        decryptor = cipher.decryptor()
        decrypted_bytes = decryptor.update(mp.encrypted) + decryptor.finalize()
        
        data = mesh_pb2.Data()
        data.ParseFromString(decrypted_bytes)
        return data
    except Exception as e:
        logging.debug(f"Decryption failed: {e}")
        return None


# ================= Whitelisting =================
#Whitelist for Gateway packets
ALLOWED_GATEWAYS = [
    "!1bbece34", #SB0
    "!1bbf02c8", #SB1
]

#Whitelist for Node position data packets
ALLOWED_NODES = [
    "!1bbece34", #SB0
    "!1bbf02c8", #SB1
    "!1117d524", #KC4MHH repeater (awaiting permission)
]

#Whitelist for ALL approved nodes
WHITELIST = set(ALLOWED_GATEWAYS) | set(ALLOWED_NODES)

# ================= Callbacks =================
def on_message_local(client, userdata, msg):
    logger.debug(f"Raw packet on topic {msg.topic}")
    try:
        envelope = mqtt_pb2.ServiceEnvelope()
        envelope.ParseFromString(msg.payload)
        packet = envelope.packet
 
        #Determine if packet signed by allowed Gateway
        if envelope.gateway_id:
            gw_id = envelope.gateway_id if envelope.gateway_id.startswith('!') else f"!{envelope.gateway_id}"

            if gw_id not in ALLOWED_GATEWAYS:
                logger.warning(f"[DROPPED] Packet from unlisted gateway: {gw_id}")
                return
                
        else:
            logger.debug(f"Packet received with no gateway ID")

        # Encryption checking & decryption - Drops encrypted map reports with no preshared key.
        # Private decryption PSKs are not currently enabled for this script
        if packet.HasField("encrypted") and not packet.HasField("decoded"):
            
            decrypted_data = decrypt_packet(packet, LOCAL_PSK_BASE64)
            if decrypted_data:
                packet.decoded.CopyFrom(decrypted_data)
                logger.debug("[PASSED] Decrypted packet")
            else:
                logger.debug("[DROPPED] Decrypt failed")
                return # Drop if decryption fails
        else:
            p_num = getattr(packet.decoded, 'portnum', 'N/A') if packet.HasField('decoded') else 'N/A'
            logger.debug(f"Packet pre-decoded. Port: {p_num}")

        # Filter Logic
        if packet.HasField("decoded"):
            portnum = packet.decoded.portnum
            allowed_ports = [
                portnums_pb2.PortNum.POSITION_APP,
                #portnums_pb2.PortNum.NODEINFO_APP,
                portnums_pb2.PortNum.MAP_REPORT_APP
            ]

            if portnum in allowed_ports:

                node_id = getattr(packet, 'from')
                node_id_string = f"!{node_id:x}"
                
                if node_id_string not in WHITELIST:
                    logger.debug(f"[DROPPED] Packet originated from unlisted node {node_id_string}")
                    return

                port_name = portnums_pb2.PortNum.Name(portnum)
                
                # Topic rewrite for Fl Mesh Map - expects: msh/US/FL/<channel>/<modem>/<node_id>
                fl_topic = f"msh/US/FL/2/c/LongFast/{node_id_string}"
                logger.info(f"[FORWARD] {port_name} from {node_id_string} -> {fl_topic}")
                fl_client.publish(fl_topic, msg.payload)
        else:
            logger.debug(f"[DROPPED] Packet {portnum} not MAP_REPORT or POSITION_APP.")

    except Exception as e:
        logger.error(f"Error processing packet: {e}", exc_info=True)

# ================= CLIENTS =================

# Local Client
local_client = mqtt.Client(mqtt.CallbackAPIVersion.VERSION2, client_id="local_filter_script")
local_client.username_pw_set(LOCAL_USER, LOCAL_PASS)
local_client.on_message = on_message_local

# use tls for Local Client!
local_client.tls_set(
    ca_certs=None,
    certfile=None,
    keyfile=None,
    cert_reqs=ssl.CERT_REQUIRED,
    tls_version=ssl.PROTOCOL_TLS
)

local_client.tls_insecure_set(True)

try:
    local_client.connect(LOCAL_BROKER, LOCAL_PORT, keepalive=60)
    local_client.subscribe(LOCAL_TOPIC)
    logger.info(f"Connected to Local Broker ({LOCAL_BROKER}:{LOCAL_PORT})")
except Exception as e:
    logger.error(f"Failed to connect to Local Broker: {e}")
    sys.exit(1)

# Florida Client
fl_client = mqtt.Client(mqtt.CallbackAPIVersion.VERSION2, client_id="fl_uplink_script")
fl_client.username_pw_set(FL_USER, FL_PASS)

try:
    fl_client.connect_async(FL_BROKER, FL_PORT, keepalive=120)
    logger.info(f"Connected to Florida Mesh Broker ({FL_BROKER}:{FL_PORT})")
    fl_client.loop_start()
except Exception as e:
    logger.error(f"Failed to connect to Florida Broker: {e}")
    sys.exit(1)

# ============== SIGNAL HANDLERS ==============
signal.signal(signal.SIGINT, cleanup)
signal.signal(signal.SIGTERM, cleanup)

# ================= MAIN LOOP =================
logger.info("Starting Meshtastic Filter Bridge for Florida Mesh...")
logger.info(f"Local Decryption Configured: {'Yes' if LOCAL_PSK_BASE64 != 'AQ==' else 'No'}")
logger.info("Forwarding only whitelisted MAP_REPORT & POSITION_APP packets.")

local_client.loop_forever()   

