#!/usr/bin/env python3
"""
TSPL Local Printer Bridge
--------------------------
Lightweight WebSocket & HTTP bridge script to send raw TSPL data from browser to local COM / USB label printers.

Runs on: http://127.0.0.1:8765 / ws://127.0.0.1:8765
"""

import sys
import os
import json
import glob
import time
import base64
import argparse
from http.server import HTTPServer, BaseHTTPRequestHandler
import socketserver
import struct, hashlib

DEFAULT_PORT = 8765

def list_available_ports():
    ports = []
    if sys.platform.startswith('win'):
        import winreg
        try:
            key = winreg.OpenKey(winreg.HKEY_LOCAL_MACHINE, r"HARDWARE\DEVICEMAP\SERIALCOMM")
            for i in range(256):
                try:
                    val = winreg.EnumValue(key, i)
                    ports.append(val[1])
                except OSError:
                    break
        except Exception:
            pass
        if not ports:
            ports = [f"COM{i}" for i in range(1, 10)]
    elif sys.platform.startswith('linux'):
        ports = glob.glob('/dev/ttyUSB*') + glob.glob('/dev/ttyACM*') + glob.glob('/dev/usb/lp*')
    elif sys.platform.startswith('darwin'):
        ports = glob.glob('/dev/tty.usbserial*') + glob.glob('/dev/tty.usbmodem*') + glob.glob('/dev/cu.*')
    
    if not ports:
        ports = ["COM1", "COM2", "COM3", "COM4", "/dev/ttyUSB0", "/dev/usb/lp0"]
    return sorted(list(set(ports)))

def print_to_serial_port(port_name, tspl_text, baud_rate=9600):
    data = tspl_text.encode('utf-8', errors='replace')
    try:
        import serial
        ser = serial.Serial(port_name, baudrate=baud_rate, timeout=3)
        ser.write(data)
        ser.flush()
        ser.close()
        return len(data), None
    except ImportError:
        pass
    
    # Fallback to OS file writing if pyserial is not installed
    try:
        if sys.platform.startswith('win'):
            full_path = f"\\\\.\\{port_name}"
            with open(full_path, "wb", buffering=0) as f:
                f.write(data)
        else:
            with open(port_name, "wb", buffering=0) as f:
                f.write(data)
        return len(data), None
    except Exception as e:
        return 0, str(e)

def perform_websocket_handshake(handler):
    key = handler.headers.get("Sec-WebSocket-Key")
    if not key:
        return False
    guid = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
    sha1 = hashlib.sha1((key + guid).encode('utf-8')).digest()
    accept_key = base64.b64encode(sha1).decode('utf-8')
    
    response = (
        "HTTP/1.1 101 Switching Protocols\r\n"
        "Upgrade: websocket\r\n"
        "Connection: Upgrade\r\n"
        f"Sec-WebSocket-Accept: {accept_key}\r\n\r\n"
    )
    handler.wfile.write(response.encode('utf-8'))
    handler.wfile.flush()
    return True

def decode_websocket_frame(rfile):
    first_byte = rfile.read(1)
    if not first_byte:
        return None
    second_byte = rfile.read(1)
    if not second_byte:
        return None
    
    length = second_byte[0] & 127
    if length == 126:
        length = struct.unpack(">H", rfile.read(2))[0]
    elif length == 127:
        length = struct.unpack(">Q", rfile.read(8))[0]
        
    masks = rfile.read(4)
    payload = rfile.read(length)
    decoded = bytearray(b ^ masks[i % 4] for i, b in enumerate(payload))
    return decoded.decode('utf-8', errors='ignore')

def encode_websocket_frame(text):
    payload = text.encode('utf-8')
    length = len(payload)
    header = bytearray()
    header.append(0x81) # FIN + text frame
    if length <= 125:
        header.append(length)
    elif length <= 65535:
        header.append(126)
        header.extend(struct.pack(">H", length))
    else:
        header.append(127)
        header.extend(struct.pack(">Q", length))
    return bytes(header) + payload

class PrinterBridgeHandler(BaseHTTPRequestHandler):
    def do_OPTIONS(self):
        self.send_response(200)
        self.send_header('Access-Control-Allow-Origin', '*')
        self.send_header('Access-Control-Allow-Methods', 'GET, POST, OPTIONS')
        self.send_header('Access-Control-Allow-Headers', 'Content-Type')
        self.end_headers()

    def do_GET(self):
        if self.headers.get("Upgrade", "").lower() == "websocket":
            if perform_websocket_handshake(self):
                self.handle_websocket()
            return
        
        self.send_response(200)
        self.send_header('Content-Type', 'application/json')
        self.send_header('Access-Control-Allow-Origin', '*')
        res = {
            "ok": True,
            "service": "TSPL Printer Bridge",
            "ports": list_available_ports(),
            "vmTimeMs": int(time.time() * 1000),
            "vmTimeIso": time.strftime("%Y-%m-%d %H:%M:%S")
        }
        self.wfile.write(json.dumps(res).encode('utf-8'))

    def do_POST(self):
        length = int(self.headers.get('Content-Length', 0))
        body_bytes = self.rfile.read(length)
        self.send_response(200)
        self.send_header('Content-Type', 'application/json')
        self.send_header('Access-Control-Allow-Origin', '*')
        self.end_headers()
        
        try:
            req = json.loads(body_bytes.decode('utf-8'))
            action = req.get("action", "print")
            if action == "list_ports":
                resp = {"ok": True, "ports": list_available_ports()}
            else:
                port = req.get("port", "COM1")
                baud = int(req.get("baudRate", 9600))
                tspl = req.get("tspl", "")
                bytes_sent, err = print_to_serial_port(port, tspl, baud)
                if err:
                    resp = {"ok": False, "error": err}
                else:
                    resp = {"ok": True, "bytesSent": bytes_sent, "port": port}
        except Exception as e:
            resp = {"ok": False, "error": str(e)}
        self.wfile.write(json.dumps(resp).encode('utf-8'))

    def handle_websocket(self):
        while True:
            try:
                msg_str = decode_websocket_frame(self.rfile)
                if msg_str is None:
                    break
                req = json.loads(msg_str)
                action = req.get("action", "print")
                if action == "status" or action == "list_ports":
                    resp = {"ok": True, "action": action, "ports": list_available_ports()}
                elif action == "print":
                    port = req.get("port", "COM1")
                    baud = int(req.get("baudRate", 9600))
                    tspl = req.get("tspl", "")
                    bytes_sent, err = print_to_serial_port(port, tspl, baud)
                    if err:
                        resp = {"ok": False, "action": "print", "error": err}
                    else:
                        resp = {"ok": True, "action": "print", "bytesSent": bytes_sent, "port": port}
                else:
                    resp = {"ok": True, "action": action}
                
                frame = encode_websocket_frame(json.dumps(resp))
                self.wfile.write(frame)
                self.wfile.flush()
            except Exception:
                break

def main():
    parser = argparse.ArgumentParser(description="TSPL Local Printer Bridge")
    parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Port to listen on (default 8765)")
    args = parser.parse_args()
    
    print(f"=====================================================")
    print(f"  TSPL Local Printer Bridge Running on Port {args.port}")
    print(f"  Listening on ws://127.0.0.1:{args.port} & http://127.0.0.1:{args.port}")
    print(f"  Available Ports: {', '.join(list_available_ports())}")
    print(f"=====================================================")
    
    server = socketserver.TCPServer(("0.0.0.0", args.port), PrinterBridgeHandler)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        print("\nShutting down TSPL Printer Bridge...")
        server.shutdown()

if __name__ == "__main__":
    main()
