#!/usr/bin/env python3
"""BackOne WebHub — ephemeral in-memory WebSocket signaling (no DB, no history)."""
import asyncio
import json
import os
import secrets
import time
import urllib.parse

try:
    import websockets
    from websockets.server import serve
except ImportError:
    raise SystemExit("python3-websockets is required")

BACKONE_HOME = os.environ.get("BACKONE_HOME", "/var/lib/backone")
AUTH_PATH = os.path.join(BACKONE_HOME, "authtoken.secret")
CONF_PATH = os.path.join(BACKONE_HOME, "local.conf")
DEFAULT_PORT = 9994
RING_TIMEOUT_SEC = 30

peers = {}  # peer_id -> {ws, network_id, registered_at}
calls = {}  # call_id -> call dict
ice_queues = {}  # call_id -> list of {from, candidate, ...}
peer_to_calls = {}  # peer_id -> set(call_id)


def load_auth_token():
    try:
        with open(AUTH_PATH, encoding="utf-8") as f:
            t = f.read().strip()
            return t if t else None
    except OSError:
        return None


def load_port():
    try:
        with open(CONF_PATH, encoding="utf-8") as f:
            cfg = json.load(f)
        settings = cfg.get("settings") or {}
        return int(settings.get("webhubPort", DEFAULT_PORT))
    except (OSError, ValueError, TypeError, json.JSONDecodeError):
        return DEFAULT_PORT


def valid_token(token):
    expected = load_auth_token()
    if not expected or not token:
        return False
    return secrets.compare_digest(token.strip(), expected.strip())


def new_call_id():
    return f"call_{int(time.time() * 1000)}_{secrets.token_hex(4)}"


def peer_key(peer_id):
    return (peer_id or "").lower()


def send_json(ws, payload):
    return ws.send(json.dumps(payload))


async def push_to_peer(target_id, payload):
    key = peer_key(target_id)
    entry = peers.get(key)
    if not entry:
        return False
    try:
        await send_json(entry["ws"], payload)
        return True
    except Exception:
        return False


def track_call_for_peer(pid, call_id):
    key = peer_key(pid)
    if key not in peer_to_calls:
        peer_to_calls[key] = set()
    peer_to_calls[key].add(call_id)


def untrack_call(call_id, call):
    for pid in (call.get("callerId"), call.get("calleeId")):
        key = peer_key(pid)
        if key in peer_to_calls:
            peer_to_calls[key].discard(call_id)
            if not peer_to_calls[key]:
                del peer_to_calls[key]


def purge_call(call_id):
    call = calls.pop(call_id, None)
    ice_queues.pop(call_id, None)
    if call:
        untrack_call(call_id, call)


def end_calls_for_peer(peer_id):
    key = peer_key(peer_id)
    for call_id in list(peer_to_calls.get(key, set())):
        call = calls.get(call_id)
        if not call:
            continue
        other = call["calleeId"] if peer_key(call["callerId"]) == key else call["callerId"]
        asyncio.create_task(push_to_peer(other, {
            "type": "call_state",
            "callId": call_id,
            "state": "ended",
            "reason": "peer_disconnect"
        }))
        purge_call(call_id)


async def ring_timeout_watch():
    while True:
        await asyncio.sleep(5)
        now = time.time()
        for call_id, call in list(calls.items()):
            if call.get("state") != "ringing":
                continue
            if now - call.get("createdAt", now) > RING_TIMEOUT_SEC:
                await push_to_peer(call["callerId"], {
                    "type": "call_state",
                    "callId": call_id,
                    "state": "rejected",
                    "reason": "timeout"
                })
                purge_call(call_id)


async def handle_message(ws, state, msg):
    mtype = msg.get("type")
    if mtype == "auth":
        token = msg.get("token", "")
        if not valid_token(token):
            await send_json(ws, {"type": "error", "error": "unauthorized"})
            await ws.close()
            return False
        state["authed"] = True
        await send_json(ws, {"type": "auth_ok"})
        return True

    if not state.get("authed"):
        await send_json(ws, {"type": "error", "error": "not authenticated"})
        return True

    if mtype == "register":
        peer_id = peer_key(msg.get("peerId", ""))
        network_id = (msg.get("networkId") or "").lower()
        if not peer_id:
            await send_json(ws, {"type": "error", "error": "peerId required"})
            return True
        peers[peer_id] = {"ws": ws, "networkId": network_id, "registeredAt": time.time()}
        state["peerId"] = peer_id
        await send_json(ws, {"type": "registered", "peerId": peer_id})
        return True

    peer_id = state.get("peerId")
    if not peer_id:
        await send_json(ws, {"type": "error", "error": "register first"})
        return True

    if mtype == "call_offer":
        target_id = peer_key(msg.get("targetId", ""))
        call_type = msg.get("callType") or "voice"
        sdp = msg.get("sdp", "")
        network_id = (msg.get("networkId") or peers[peer_id].get("networkId") or "").lower()
        if not target_id or not sdp:
            await send_json(ws, {"type": "error", "error": "targetId and sdp required"})
            return True
        call_id = new_call_id()
        call = {
            "callId": call_id,
            "callType": call_type,
            "networkId": network_id,
            "callerId": peer_id,
            "calleeId": target_id,
            "state": "ringing",
            "offerSdp": sdp,
            "answerSdp": "",
            "createdAt": time.time()
        }
        calls[call_id] = call
        ice_queues[call_id] = []
        track_call_for_peer(peer_id, call_id)
        track_call_for_peer(target_id, call_id)
        delivered = await push_to_peer(target_id, {
            "type": "incoming_call",
            "callId": call_id,
            "callType": call_type,
            "networkId": network_id,
            "callerId": peer_id,
            "callerName": msg.get("callerName", peer_id),
            "sdp": sdp,
            "overlayIps": msg.get("overlayIps", [])
        })
        await send_json(ws, {
            "type": "call_state",
            "callId": call_id,
            "state": "ringing",
            "delivered": delivered
        })
        return True

    if mtype == "call_answer":
        call_id = msg.get("callId", "")
        sdp = msg.get("sdp", "")
        call = calls.get(call_id)
        if not call or peer_key(call["calleeId"]) != peer_id:
            await send_json(ws, {"type": "error", "error": "call not found"})
            return True
        call["answerSdp"] = sdp
        call["state"] = "accepted"
        await push_to_peer(call["callerId"], {
            "type": "call_state",
            "callId": call_id,
            "state": "accepted",
            "sdp": sdp
        })
        await send_json(ws, {"type": "call_state", "callId": call_id, "state": "accepted"})
        return True

    if mtype == "call_reject":
        call_id = msg.get("callId", "")
        call = calls.get(call_id)
        if call and peer_key(call["calleeId"]) == peer_id:
            await push_to_peer(call["callerId"], {
                "type": "call_state",
                "callId": call_id,
                "state": "rejected"
            })
            purge_call(call_id)
        return True

    if mtype == "call_end":
        call_id = msg.get("callId", "")
        call = calls.get(call_id)
        if call and peer_id in (peer_key(call["callerId"]), peer_key(call["calleeId"])):
            other = call["calleeId"] if peer_key(call["callerId"]) == peer_id else call["callerId"]
            await push_to_peer(other, {
                "type": "call_state",
                "callId": call_id,
                "state": "ended"
            })
            purge_call(call_id)
        return True

    if mtype == "ice_candidate":
        call_id = msg.get("callId", "")
        call = calls.get(call_id)
        if not call:
            return True
        entry = {
            "from": peer_id,
            "candidate": msg.get("candidate"),
            "sdpMid": msg.get("sdpMid"),
            "sdpMLineIndex": msg.get("sdpMLineIndex"),
            "at": time.time()
        }
        ice_queues.setdefault(call_id, []).append(entry)
        other = call["calleeId"] if peer_key(call["callerId"]) == peer_id else call["callerId"]
        await push_to_peer(other, {
            "type": "ice_candidate",
            "callId": call_id,
            **{k: entry[k] for k in ("candidate", "sdpMid", "sdpMLineIndex")}
        })
        return True

    if mtype == "ice_fetch":
        call_id = msg.get("callId", "")
        since = float(msg.get("since", 0))
        queued = [e for e in ice_queues.get(call_id, []) if e.get("at", 0) > since and e.get("from") != peer_id]
        await send_json(ws, {"type": "ice_candidates", "callId": call_id, "candidates": queued})
        return True

    if mtype == "chat_message":
        target_id = peer_key(msg.get("targetId", ""))
        if target_id:
            await push_to_peer(target_id, {
                "type": "chat_message",
                "networkId": msg.get("networkId", ""),
                "senderId": peer_id,
                "senderName": msg.get("senderName", peer_id),
                "body": msg.get("body", ""),
                "at": int(time.time() * 1000)
            })
        return True

    if mtype == "typing":
        target_id = peer_key(msg.get("targetId", ""))
        if target_id:
            await push_to_peer(target_id, {
                "type": "typing",
                "networkId": msg.get("networkId", ""),
                "senderId": peer_id,
                "typing": bool(msg.get("typing", True))
            })
        return True

    await send_json(ws, {"type": "error", "error": f"unknown type: {mtype}"})
    return True


async def connection_handler(ws, path):
    state = {"authed": False, "peerId": None}
    query = urllib.parse.parse_qs(urllib.parse.urlparse(path).query)
    if "auth" in query and valid_token(query["auth"][0]):
        state["authed"] = True
        await send_json(ws, {"type": "auth_ok"})

    try:
        async for raw in ws:
            try:
                msg = json.loads(raw)
            except json.JSONDecodeError:
                await send_json(ws, {"type": "error", "error": "invalid json"})
                continue
            if not await handle_message(ws, state, msg):
                break
    finally:
        pid = state.get("peerId")
        if pid and peers.get(pid, {}).get("ws") is ws:
            del peers[pid]
            end_calls_for_peer(pid)


async def main():
    port = load_port()
    host = "127.0.0.1"
    asyncio.create_task(ring_timeout_watch())
    async with serve(connection_handler, host, port, ping_interval=20, ping_timeout=20):
        print(f"BackOne WebHub listening on ws://{host}:{port}", flush=True)
        await asyncio.Future()


if __name__ == "__main__":
    try:
        asyncio.run(main())
    except KeyboardInterrupt:
        pass
