import os
import time
import json
import subprocess
import threading
from http.server import HTTPServer, BaseHTTPRequestHandler

def get_cpu_stats():
    try:
        with open('/proc/stat', 'r') as f:
            lines = f.readlines()
    except FileNotFoundError:
        return {}
    
    stats = {}
    for line in lines:
        if line.startswith('cpu'):
            parts = line.split()
            name = parts[0]
            # user, nice, system, idle, iowait, irq, softirq, steal, guest, guest_nice
            total_time = sum(float(x) for x in parts[1:])
            idle_time = float(parts[4]) + float(parts[5]) if len(parts) > 5 else float(parts[4])
            stats[name] = (total_time, idle_time)
    return stats

def calculate_cpu_percent(stats1, stats2):
    percents = {}
    for name in stats1:
        if name in stats2:
            t1, i1 = stats1[name]
            t2, i2 = stats2[name]
            total_diff = t2 - t1
            idle_diff = i2 - i1
            if total_diff > 0:
                percent = 100.0 * (total_diff - idle_diff) / total_diff
            else:
                percent = 0.0
            percents[name] = round(percent, 2)
    return percents

def get_mem_stats():
    mem = {}
    try:
        with open('/proc/meminfo', 'r') as f:
            for line in f:
                parts = line.split()
                mem[parts[0].strip(':')] = int(parts[1]) * 1024
    except FileNotFoundError:
        return {'total': 0, 'free': 0, 'used': 0}
    
    total = mem.get('MemTotal', 0)
    free = mem.get('MemFree', 0)
    buffers = mem.get('Buffers', 0)
    cached = mem.get('Cached', 0)
    available = mem.get('MemAvailable', free + buffers + cached)
    used = total - available
    return {
        'total': total,
        'free': available,
        'used': used
    }

def get_core_temp():
    try:
        out = subprocess.check_output(['vcgencmd', 'measure_temp']).decode('utf-8')
        # expected format: temp=45.0'C
        return float(out.replace('temp=', '').replace("'C\n", ''))
    except Exception:
        return 0.0

def get_throttled():
    try:
        out = subprocess.check_output(['vcgencmd', 'get_throttled']).decode('utf-8')
        # expected format: throttled=0x0
        return out.strip().split('=')[1]
    except Exception:
        return "0x0"

def get_net_stats():
    net = {}
    try:
        with open('/proc/net/dev', 'r') as f:
            lines = f.readlines()[2:]
        for line in lines:
            parts = line.split(':')
            if len(parts) == 2:
                iface = parts[0].strip()
                stats = parts[1].split()
                rx_bytes = int(stats[0])
                tx_bytes = int(stats[8])
                net[iface] = (rx_bytes, tx_bytes)
    except FileNotFoundError:
        pass
    return net

def calculate_net_mbps(net1, net2, dt):
    mbps = {}
    for iface in net1:
        if iface in net2:
            rx1, tx1 = net1[iface]
            rx2, tx2 = net2[iface]
            rx_mbps = (rx2 - rx1) * 8 / 1_000_000 / dt
            tx_mbps = (tx2 - tx1) * 8 / 1_000_000 / dt
            mbps[iface] = {'rx_mbps': round(rx_mbps, 2), 'tx_mbps': round(tx_mbps, 2)}
    return mbps

class Monitor:
    def __init__(self):
        self.last_cpu = get_cpu_stats()
        self.last_net = get_net_stats()
        self.last_time = time.time()
        self.current_stats = {}
    
    def update(self):
        curr_time = time.time()
        curr_cpu = get_cpu_stats()
        curr_net = get_net_stats()
        dt = curr_time - self.last_time
        if dt <= 0: dt = 0.001
        
        cpu_pct = calculate_cpu_percent(self.last_cpu, curr_cpu)
        net_mbps = calculate_net_mbps(self.last_net, curr_net, dt)
        
        try:
            load_avg = os.getloadavg()
        except AttributeError:
            load_avg = (0.0, 0.0, 0.0)
            
        self.current_stats = {
            'timestamp': curr_time,
            'cpu': cpu_pct,
            'memory': get_mem_stats(),
            'temp_c': get_core_temp(),
            'throttled': get_throttled(),
            'load_avg': load_avg,
            'network': net_mbps
        }
        
        self.last_cpu = curr_cpu
        self.last_net = curr_net
        self.last_time = curr_time

monitor = Monitor()

class APIHandler(BaseHTTPRequestHandler):
    def do_GET(self):
        if self.path == '/api/stats':
            self.send_response(200)
            self.send_header('Content-type', 'application/json')
            self.end_headers()
            monitor.update()
            self.wfile.write(json.dumps(monitor.current_stats).encode('utf-8'))
        else:
            self.send_response(404)
            self.end_headers()

def run_server(port=8080):
    server = HTTPServer(('', port), APIHandler)
    server.serve_forever()

if __name__ == '__main__':
    import argparse
    parser = argparse.ArgumentParser(description="Pi 5 Telemetry Monitor")
    parser.add_argument('--api', action='store_true', help='Run JSON API server on port 8080')
    parser.add_argument('--port', type=int, default=8080, help='API server port')
    parser.add_argument('--interval', type=float, default=1.0, help='CLI update interval')
    args = parser.parse_args()
    
    if args.api:
        print(f"Starting API server on port {args.port}")
        threading.Thread(target=run_server, args=(args.port,), daemon=True).start()
        try:
            while True:
                time.sleep(1)
        except KeyboardInterrupt:
            pass
    else:
        try:
            while True:
                monitor.update()
                print(json.dumps(monitor.current_stats, indent=2))
                time.sleep(args.interval)
        except KeyboardInterrupt:
            pass
