#!/usr/bin/env python3
"""
TODAY sync server — stores data in iCloud Drive JSON file.
Runs HTTPS on port 4443. Serves HTML + sync API.
"""
import http.server, ssl, json, os, threading, time, subprocess, sys

ICLOUD_FILE = os.path.expanduser(
    '~/Library/Mobile Documents/com~apple~CloudDocs/today-sync.json'
)
LOCAL_FILE = os.path.expanduser('~/.today/today-sync.json')
SERVE_DIR = os.path.expanduser('~/.today')
PORT = 4443
CERT = os.path.expanduser('~/.today/cert.pem')
KEY = os.path.expanduser('~/.today/key.pem')

lock = threading.Lock()
REPORT_SCRIPT = os.path.expanduser('~/.today/report.py')

def regenerate_report():
    """Run report.py in background after each sync."""
    try:
        subprocess.Popen(
            [sys.executable, REPORT_SCRIPT],
            stdout=subprocess.DEVNULL,
            stderr=subprocess.DEVNULL
        )
    except Exception:
        pass

def _atomic_write(path, data):
    """Atomic write: temp file + fsync + rename."""
    os.makedirs(os.path.dirname(path), exist_ok=True)
    tmp = path + '.tmp'
    with open(tmp, 'w') as f:
        json.dump(data, f)
        f.flush()
        os.fsync(f.fileno())
    os.replace(tmp, path)

def read_data():
    # Try local first (always writable), then iCloud
    for path in [LOCAL_FILE, ICLOUD_FILE]:
        try:
            with open(path, 'r') as f:
                return json.load(f)
        except:
            continue
    return {}

def write_data(data):
    # Always write to local (guaranteed writable)
    _atomic_write(LOCAL_FILE, data)
    # Try iCloud — if permission denied, silently skip
    try:
        _atomic_write(ICLOUD_FILE, data)
    except PermissionError:
        pass
    # Backup copy
    try:
        _atomic_write(LOCAL_FILE.replace('.json', '-backup.json'), data)
    except:
        pass

def merge(local, remote):
    """Merge two datasets. Remote (incoming POST) wins for booleans. MAX scores. Dedup logs."""
    result = {}
    all_dates = set(list(local.keys()) + list(remote.keys()))
    for date in all_dates:
        a = local.get(date, {})
        b = remote.get(date, {})
        result[date] = {}
        for cat in ['h', 'm', 'a', 'sd', 'sc', 'sk']:
            aa = a.get(cat, {})
            bb = b.get(cat, {})
            merged = {}
            for k in set(list(aa.keys()) + list(bb.keys())):
                va = aa.get(k, 0)
                vb = bb.get(k, 0)
                if cat == 'sc':
                    merged[k] = max(va or 0, vb or 0)
                else:
                    # Remote (incoming) wins — supports unchecking
                    if k in bb:
                        merged[k] = vb
                    else:
                        merged[k] = va
            if merged:
                result[date][cat] = merged
        # Merge log arrays with dedup
        a_log = a.get('log', [])
        b_log = b.get('log', [])
        if a_log or b_log:
            combined = a_log + b_log
            seen = set()
            deduped = []
            for entry in combined:
                key = str(entry.get('t','')) + '|' + str(entry.get('e','')) + '|' + str(entry.get('id', entry.get('idx', entry.get('tab', entry.get('cat', '')))))
                if key not in seen:
                    seen.add(key)
                    deduped.append(entry)
            result[date]['log'] = sorted(deduped, key=lambda x: x.get('t', 0))
        # Merge memo — longer string wins, locked state follows the memo
        a_memo = a.get('memo', '')
        b_memo = b.get('memo', '')
        if a_memo or b_memo:
            if len(b_memo) >= len(a_memo):
                result[date]['memo'] = b_memo
                if b.get('memoLocked'):
                    result[date]['memoLocked'] = True
            else:
                result[date]['memo'] = a_memo
                if a.get('memoLocked'):
                    result[date]['memoLocked'] = True
        # Merge night out — if either side marked it, keep it
        if b.get('no') and not a.get('no'):
            result[date]['no'] = b['no']
        elif a.get('no'):
            result[date]['no'] = a['no']
        # Merge closure record — earliest wins
        a_cl = a.get('cl')
        b_cl = b.get('cl')
        if a_cl and b_cl:
            result[date]['cl'] = a_cl if a_cl.get('at', 0) <= b_cl.get('at', 0) else b_cl
        elif a_cl:
            result[date]['cl'] = a_cl
        elif b_cl:
            result[date]['cl'] = b_cl
    return result

class Handler(http.server.SimpleHTTPRequestHandler):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, directory=SERVE_DIR, **kwargs)

    def end_headers(self):
        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.send_header('Cache-Control', 'no-cache, no-store, must-revalidate')
        self.send_header('Pragma', 'no-cache')
        self.send_header('Expires', '0')
        super().end_headers()

    def do_OPTIONS(self):
        self.send_response(200)
        self.end_headers()

    def do_GET(self):
        if self.path == '/sync':
            with lock:
                data = read_data()
            self.send_response(200)
            self.send_header('Content-Type', 'application/json')
            self.end_headers()
            self.wfile.write(json.dumps(data).encode())
        else:
            super().do_GET()

    def do_POST(self):
        if self.path == '/brief-2aaffc9f7ae7/refresh':
            # pull-to-refresh on the brief: flag a rebuild request; a watcher
            # (on the jarvis mini) picks it up. Capability URL is the guard;
            # the watcher enforces a cooldown so pulls can't stack builds.
            try:
                flag_dir = os.path.expanduser('~/brief-build')
                os.makedirs(flag_dir, exist_ok=True)
                with open(os.path.join(flag_dir, 'refresh-requested'), 'w') as f:
                    f.write(str(time.time()))
                sd = os.path.expanduser('~/.today/brief-2aaffc9f7ae7')
                tmp = os.path.join(sd, 'refresh-state.json.tmp')
                with open(tmp, 'w') as f:
                    json.dump({'state': 'building', 'note': 'requested', 'at': time.time()}, f)
                os.replace(tmp, os.path.join(sd, 'refresh-state.json'))
            except Exception:
                pass
            self.send_response(200)
            self.send_header('Content-Type', 'application/json')
            self.end_headers()
            self.wfile.write(b'{"ok": true}')
            return
        if self.path == '/sync':
            length = int(self.headers.get('Content-Length', 0))
            body = self.rfile.read(length)
            try:
                incoming = json.loads(body)
            except:
                incoming = {}
            with lock:
                existing = read_data()
                merged = merge(existing, incoming)
                write_data(merged)
            regenerate_report()
            self.send_response(200)
            self.send_header('Content-Type', 'application/json')
            self.end_headers()
            self.wfile.write(json.dumps(merged).encode())
        else:
            self.send_response(404)
            self.end_headers()

    def log_message(self, format, *args):
        pass  # silent

def ensure_cert():
    import subprocess, datetime
    need_new = not os.path.exists(CERT) or not os.path.exists(KEY)
    if not need_new:
        # Check if cert expires within 30 days
        try:
            r = subprocess.run(['openssl', 'x509', '-in', CERT, '-noout', '-enddate'],
                             capture_output=True, text=True)
            exp = r.stdout.strip().replace('notAfter=', '')
            from email.utils import parsedate_to_datetime
            exp_dt = datetime.datetime.strptime(exp, '%b %d %H:%M:%S %Y %Z')
            if exp_dt - datetime.datetime.utcnow() < datetime.timedelta(days=30):
                need_new = True
        except:
            pass
    if need_new:
        os.system(
            f"openssl req -x509 -newkey rsa:2048 -keyout {KEY} -out {CERT} "
            f"-days 730 -nodes -subj '/CN=today' 2>/dev/null"
        )

HTTP_PORT = 8443  # Plain HTTP for Tailscale Funnel

if __name__ == '__main__':
    ensure_cert()
    # Threaded: a slow /sync (report.py runs after each) must never stall the
    # page's own data fetches behind it.
    server = http.server.ThreadingHTTPServer(('0.0.0.0', PORT), Handler)
    ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    ctx.load_cert_chain(CERT, KEY)
    server.socket = ctx.wrap_socket(server.socket, server_side=True)
    # HTTP server (Tailscale Funnel proxy)
    http_server = http.server.ThreadingHTTPServer(('127.0.0.1', HTTP_PORT), Handler)
    t = threading.Thread(target=http_server.serve_forever, daemon=True)
    t.start()
    print(f'TODAY sync server: HTTPS on :{PORT}, HTTP on :{HTTP_PORT}')
    server.serve_forever()
