#!/usr/bin/env python3
"""
DVeProto wire 1.4 — standalone reference (FINAL 1.x). pip install cryptography.
Handshake JSON v=1. After select 1.4: binary frames with pkt_seq + AAD + resume.
Next major: 2.0 (DVeNet / DVeVPN / UDP+TCP carriers). Frame body is carrier-agnostic.
"""
from __future__ import annotations
import base64, json, os, struct, time
from dataclasses import dataclass
from typing import Any, Dict, Optional, Tuple
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey, X25519PublicKey
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF

PROTO_NAME = "DVeProto"
PROTO_HANDSHAKE_VERSION = 1
WIRE = "1.4"
INFO_C2S, INFO_S2C = b"DVeProto-v1/c2s", b"DVeProto-v1/s2c"
INFO_RESUME_C2S, INFO_RESUME_S2C = b"DVeProto-v1/resume-c2s", b"DVeProto-v1/resume-s2c"
BIN_VER, HEADER_LEN, NONCE_LEN, GCM_TAG_LEN = 0x14, 20, 12, 16
DEFAULT_STREAM_WINDOW = 256 * 1024
DEFAULT_CAPS = 0x083F  # resume|priority|data_ack|cancel|rtt|pkt_seq|websocket

PTYPE_APP_JSON = 0x01
PTYPE_PING, PTYPE_PONG, PTYPE_CLOSE, PTYPE_WINDOW_UPDATE = 0x30, 0x31, 0x32, 0x33
PTYPE_STREAM_PRIORITY, PTYPE_STREAM_CANCEL = 0x34, 0x35
PTYPE_DATA_ACK, PTYPE_SESSION_TICKET, PTYPE_CAPS = 0x36, 0x37, 0x38

def _b64e(raw: bytes) -> str:
    return base64.b64encode(raw).decode("ascii")
def _b64d(s: str) -> bytes:
    return base64.b64decode(s.encode("ascii"), validate=True)

def derive_aes_keys(shared: bytes):
    c2s = HKDF(algorithm=hashes.SHA256(), length=32, salt=b"", info=INFO_C2S).derive(shared)
    s2c = HKDF(algorithm=hashes.SHA256(), length=32, salt=b"", info=INFO_S2C).derive(shared)
    return c2s, s2c

def derive_resume_keys(token: bytes):
    c2s = HKDF(algorithm=hashes.SHA256(), length=32, salt=b"", info=INFO_RESUME_C2S).derive(token)
    s2c = HKDF(algorithm=hashes.SHA256(), length=32, salt=b"", info=INFO_RESUME_S2C).derive(token)
    return c2s, s2c

def make_aad(ptype: int, stream_id: int, pkt_seq: int) -> bytes:
    return bytes([BIN_VER, ptype & 0xFF]) + struct.pack(">HI", stream_id & 0xFFFF, pkt_seq & 0xFFFFFFFF)

@dataclass
class DecodedFrame:
    ptype: int
    plaintext: bytes
    stream_id: int = 0
    pkt_seq: int = 0

class DVeClientSession:
    def __init__(self, c2s: bytes, s2c: bytes, caps: int = DEFAULT_CAPS):
        self._c2s, self._s2c = AESGCM(c2s), AESGCM(s2c)
        self.wire_version = WIRE
        self._send_prefix = os.urandom(4)
        self._send_counter = 0
        self._pkt_seq_send = 0
        self._send_credit = {0: DEFAULT_STREAM_WINDOW}
        self._caps = caps
        self._session_id = None
        self._resume_token = None
        self._rtt_ms = None

    @classmethod
    def from_server_hello_text(cls, hello_text: str):
        data = json.loads(hello_text)
        server_pub = X25519PublicKey.from_public_bytes(_b64d(data["server_pk"]))
        priv = X25519PrivateKey.generate()
        c2s, s2c = derive_aes_keys(priv.exchange(server_pub))
        pub = priv.public_key().public_bytes(
            encoding=serialization.Encoding.Raw, format=serialization.PublicFormat.Raw
        )
        caps = int(data.get("caps") or DEFAULT_CAPS) & DEFAULT_CAPS
        ack = json.dumps(
            {
                "type": "dve_client_ack",
                "proto": PROTO_NAME,
                "v": 1,
                "client_pk": _b64e(pub),
                "select": "1.4",
                "caps": caps,
            },
            separators=(",", ":"),
        )
        return cls(c2s, s2c, caps=caps), ack

    @classmethod
    def from_resume(cls, hello_text: str, session_id: bytes, token: bytes):
        data = json.loads(hello_text)
        caps = int(data.get("caps") or DEFAULT_CAPS) & DEFAULT_CAPS
        c2s, s2c = derive_resume_keys(token)
        ack = json.dumps(
            {
                "type": "dve_client_ack",
                "proto": PROTO_NAME,
                "v": 1,
                "select": "1.4",
                "resume": {"session_id": _b64e(session_id), "token": _b64e(token)},
                "caps": caps,
            },
            separators=(",", ":"),
        )
        sess = cls(c2s, s2c, caps=caps)
        sess._session_id, sess._resume_token = session_id, token
        return sess, ack

    def _next_nonce(self):
        self._send_counter += 1
        return self._send_prefix + struct.pack(">Q", self._send_counter)

    def _seal(self, aead, ptype, pt, stream_id=0):
        sid = stream_id & 0xFFFF
        exempt = {
            PTYPE_PING, PTYPE_PONG, PTYPE_CLOSE, PTYPE_WINDOW_UPDATE,
            PTYPE_STREAM_PRIORITY, PTYPE_STREAM_CANCEL, PTYPE_DATA_ACK,
            PTYPE_SESSION_TICKET, PTYPE_CAPS,
        }
        if ptype not in exempt:
            credit = self._send_credit.get(sid, DEFAULT_STREAM_WINDOW)
            if len(pt) > credit:
                raise ValueError("flow control: insufficient credit")
            self._send_credit[sid] = credit - len(pt)
        self._pkt_seq_send += 1
        pkt_seq = self._pkt_seq_send
        nonce = self._next_nonce()
        aad = make_aad(ptype, sid, pkt_seq)
        ct = aead.encrypt(nonce, pt, aad)
        out = bytearray(HEADER_LEN + len(ct))
        out[0] = BIN_VER
        out[1] = ptype
        out[2] = sid >> 8
        out[3] = sid & 0xFF
        struct.pack_into(">I", out, 4, pkt_seq)
        out[8:20] = nonce
        out[20:] = ct
        return bytes(out)

    def _open(self, aead, frame):
        raw = bytes(frame)
        if raw[0] != BIN_VER:
            raise ValueError("bad ver")
        sid = (raw[2] << 8) | raw[3]
        pkt_seq = struct.unpack_from(">I", raw, 4)[0]
        aad = make_aad(raw[1], sid, pkt_seq)
        pt = aead.decrypt(raw[8:20], raw[20:], aad)
        if raw[1] == PTYPE_WINDOW_UPDATE and len(pt) >= 4:
            self._send_credit[sid] = self._send_credit.get(sid, 0) + struct.unpack_from(">I", pt, 0)[0]
        if raw[1] == PTYPE_SESSION_TICKET and len(pt) == 52:
            self._session_id, self._resume_token = pt[:16], pt[20:52]
        if raw[1] == PTYPE_PONG and len(pt) >= 8:
            self._rtt_ms = max(0.0, time.time() * 1000 - struct.unpack_from(">Q", pt, 0)[0])
        return DecodedFrame(raw[1], pt, sid, pkt_seq)

    def pack_outgoing(self, obj, stream_id=0) -> bytes:
        pt = json.dumps(obj, separators=(",", ":"), ensure_ascii=False).encode()
        return self._seal(self._c2s, PTYPE_APP_JSON, pt, stream_id)

    def unpack_incoming(self, frame) -> dict:
        fr = self._open(self._s2c, frame)
        if fr.ptype != PTYPE_APP_JSON:
            raise ValueError("not APP_JSON")
        return json.loads(fr.plaintext.decode())

    def pack_window_update(self, credit: int, stream_id=0) -> bytes:
        return self._seal(self._c2s, PTYPE_WINDOW_UPDATE, struct.pack(">I", credit & 0xFFFFFFFF), stream_id)

    def pack_ping_rtt(self, stream_id=0) -> bytes:
        ts = struct.pack(">Q", int(time.time() * 1000) & 0xFFFFFFFFFFFFFFFF)
        return self._seal(self._c2s, PTYPE_PING, ts, stream_id)

    def pack_stream_priority(self, priority: int, stream_id=0) -> bytes:
        return self._seal(self._c2s, PTYPE_STREAM_PRIORITY, bytes([priority & 0xFF]), stream_id)

    def pack_data_ack(self, last_seq: int, stream_id=0) -> bytes:
        return self._seal(self._c2s, PTYPE_DATA_ACK, struct.pack(">I", last_seq & 0xFFFFFFFF), stream_id)

    def stats(self) -> Dict[str, Any]:
        return {
            "wire": WIRE,
            "rtt_ms": self._rtt_ms,
            "pkt_seq_send": self._pkt_seq_send,
            "caps": self._caps,
            "session_id": _b64e(self._session_id) if self._session_id else None,
        }

if __name__ == "__main__":
    # minimal round-trip smoke
    from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey as Priv
    srv = Priv.generate()
    hello = {
        "type": "dve_hello",
        "proto": PROTO_NAME,
        "v": 1,
        "server_pk": _b64e(
            srv.public_key().public_bytes(
                encoding=serialization.Encoding.Raw, format=serialization.PublicFormat.Raw
            )
        ),
        "offer": ["1.0", "1.4"],
        "caps": DEFAULT_CAPS,
    }
    client, ack = DVeClientSession.from_server_hello_text(json.dumps(hello))
    ack_obj = json.loads(ack)
    peer_pub = X25519PublicKey.from_public_bytes(_b64d(ack_obj["client_pk"]))
    shared = srv.exchange(peer_pub)
    c2s, s2c = derive_aes_keys(shared)
    server = AESGCM(s2c)
    # server seals APP_JSON like client _open expects (s2c)
    class S:
        pass
    # use client unpack by building frame with server keys via a tiny helper
    sess_srv_side = DVeClientSession(c2s, s2c)  # wrong dirs for pack; just verify client pack
    frame = client.pack_outgoing({"dve_op": "ping"})
    assert frame[0] == BIN_VER and len(frame) > HEADER_LEN
    print("1.4 smoke ok", client.stats())
