#!/usr/bin/env python3
"""
DVeProto wire 1.3 — standalone reference (pip install cryptography).
Handshake: dve_hello / dve_client_ack (JSON, v=1). HKDF labels DVeProto-v1/c2s|s2c.
This file implements ONLY wire 1.3 after handshake.
"""
from __future__ import annotations
import base64, json, os, struct
from dataclasses import dataclass
from typing import Any, Dict, Optional, Tuple, Union
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.3"
INFO_C2S = b"DVeProto-v1/c2s"
INFO_S2C = b"DVeProto-v1/s2c"
NONCE_LEN = 12
GCM_TAG_LEN = 16

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

BIN_VER = 0x13
HEADER_LEN = 16
PTYPE_APP_JSON = 0x01
PTYPE_PING, PTYPE_PONG, PTYPE_CLOSE, PTYPE_WINDOW_UPDATE = 0x30, 0x31, 0x32, 0x33
STREAM_DEFAULT = 0
DEFAULT_STREAM_WINDOW = 256 * 1024

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

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

class DVeClientSession:
    def __init__(self, c2s: bytes, s2c: bytes):
        self._c2s, self._s2c = AESGCM(c2s), AESGCM(s2c)
        self.wire_version = WIRE
        self._send_prefix = os.urandom(4)
        self._send_counter = 0
        self._send_credit = {0: DEFAULT_STREAM_WINDOW}
    @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)
        ack = json.dumps({"type":"dve_client_ack","proto":PROTO_NAME,"v":1,"client_pk":_b64e(pub),"select":"1.3"}, separators=(",",":"))
        return cls(c2s, s2c), 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
        if ptype not in (PTYPE_PING, PTYPE_PONG, PTYPE_CLOSE, PTYPE_WINDOW_UPDATE):
            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)
        nonce=self._next_nonce()
        aad=make_aad(ptype, sid)
        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; out[4:16]=nonce; out[16:]=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]
        aad=make_aad(raw[1], sid)
        pt=aead.decrypt(raw[4:16], raw[16:], aad)
        if raw[1]==PTYPE_WINDOW_UPDATE and len(pt)>=4:
            credit=struct.unpack_from(">I", pt, 0)[0]
            self._send_credit[sid]=self._send_credit.get(sid,0)+credit
        return DecodedFrame(raw[1], pt, sid)
    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(self, payload=b"", stream_id=0):
        return self._seal(self._c2s, PTYPE_PING, payload, stream_id)

if __name__ == "__main__":
    print("DVeProto reference wire", WIRE)
