#!/usr/bin/env python3
"""Minimal plaintext Photon (eNet/LoadBalancing) server for SvT revival.

Speaks enough of Photon's UDP protocol to get the client past the
"connecting to game server" wall: eNet handshake, Init, and the
LoadBalancing operations (GetRegions on the nameserver role; Authenticate +
JoinLobby on the master role). Encryption is expected to be DISABLED on the
client side (patched), so all operations are plaintext Protocol16.

Run:  photon_srv.py --role master --port 5055   (or --role nameserver --port 5058)
"""
import socket, struct, sys, time, threading, argparse

def now_ms(): return int(time.time()*1000) & 0xffffffff

# ---- eNet command types ----
ACK=1; CONNECT=2; VERIFYCONNECT=3; DISCONNECT=4; PING=5; RELIABLE=6; UNRELIABLE=7; FRAGMENT=8

# ---- Protocol16 (GpBinary) types ----
class Gp:
    Null=42; Dictionary=68; StringArray=97; Byte=98; Double=100; EventData=101
    Float=102; Hashtable=104; Integer=105; Short=107; Long=108; IntArray=110
    Boolean=111; OperationResponse=112; OperationRequest=113; String=115
    ByteArray=120; Array=121; ObjectArray=122

def w_u16(v): return struct.pack(">H",v & 0xffff)
def w_u32(v): return struct.pack(">I",v & 0xffffffff)

# ---------- Protocol16 serialize ----------
def ser_value(v):
    if v is None: return bytes([Gp.Null])
    if isinstance(v,bool): return bytes([Gp.Boolean,1 if v else 0])
    if isinstance(v,int):
        # default ints to Integer(4)
        return bytes([Gp.Integer])+struct.pack(">i",v)
    if isinstance(v,Short):
        return bytes([Gp.Short])+struct.pack(">h",v.v)
    if isinstance(v,Byte):
        return bytes([Gp.Byte,v.v & 0xff])
    if isinstance(v,float):
        return bytes([Gp.Float])+struct.pack(">f",v)
    if isinstance(v,str):
        b=v.encode('utf-8'); return bytes([Gp.String])+w_u16(len(b))+b
    if isinstance(v,bytes):
        return bytes([Gp.ByteArray])+w_u32(len(v))+v
    if isinstance(v,dict):
        return bytes([Gp.Hashtable])+ser_hashtable_body(v)
    if isinstance(v,StrArray):
        out=bytes([Gp.StringArray])+w_u16(len(v.items))
        for s in v.items:
            bb=s.encode('utf-8'); out+=w_u16(len(bb))+bb
        return out
    if isinstance(v,ByteArr):
        return bytes([Gp.ByteArray])+w_u32(len(v.data))+v.data
    if isinstance(v,BoolArray):
        out=bytes([Gp.Array])+w_u16(len(v.items))+bytes([Gp.Boolean])
        for x in v.items: out+=bytes([1 if x else 0])
        return out
    raise ValueError("cannot serialize %r"%type(v))

class Short:
    def __init__(self,v): self.v=v
class Byte:
    def __init__(self,v): self.v=v
class StrArray:
    def __init__(self,items): self.items=items
class ByteArr:
    def __init__(self,data): self.data=data
class BoolArray:
    def __init__(self,items): self.items=items

def ser_hashtable_body(d):
    out=w_u16(len(d))
    for k,v in d.items():
        out+=ser_value(k); out+=ser_value(v)
    return out

def ser_params(params):
    """params: dict {byteKey: value}. Returns short count + entries(key byte + typed value)."""
    out=w_u16(len(params))
    for k,v in params.items():
        out+=bytes([k & 0xff])+ser_value(v)
    return out

def make_op_response(opcode, ret_code=0, debug=None, params=None):
    params=params or {}
    body=bytes([0xf3,0x03,opcode & 0xff])+struct.pack(">h",ret_code)
    # debug message: typed value (Null or String)
    body+=ser_value(debug)
    body+=ser_params(params)
    return body

def make_event(evcode, params=None):
    params=params or {}
    return bytes([0xf3,0x04,evcode & 0xff])+ser_params(params)

def make_init_response():
    return bytes([0xf3,0x01,0x00])

# A real (plaintext) f3 07 DH key-exchange response captured from Photon's live
# nameserver. Replayed to make the client consider encryption "established".
# The derived key is irrelevant because we force IsEncryptionAvailable=false so
# the client sends operations in plaintext.
CANNED_DH_RESPONSE=bytes.fromhex(
  "f3070000002a0001017800000060"
  "e137d543fe7e7e440277f9561cf586aa95a807121a9f2e68f61474618bc96e47"
  "9d1549065b8916bda8eec0875db77974862fa87b57564d47e6d66dc16e695155"
  "d6e5d25b6f3c1ea47d23fa01f73a884fec3284586c9637494b8d38f481291bc2")

# ---------- Protocol16 deserialize ----------
def deser_value(b,o):
    t=b[o]; o+=1
    if t==Gp.Null: return None,o
    if t==Gp.Boolean: return (b[o]!=0),o+1
    if t==Gp.Byte: return Byte(b[o]),o+1
    if t==Gp.Short: return struct.unpack_from(">h",b,o)[0],o+2
    if t==Gp.Integer: return struct.unpack_from(">i",b,o)[0],o+4
    if t==Gp.Long: return struct.unpack_from(">q",b,o)[0],o+8
    if t==Gp.Float: return struct.unpack_from(">f",b,o)[0],o+4
    if t==Gp.Double: return struct.unpack_from(">d",b,o)[0],o+8
    if t==Gp.String:
        ln=struct.unpack_from(">H",b,o)[0]; o+=2
        return b[o:o+ln].decode('utf-8','replace'),o+ln
    if t==Gp.ByteArray:
        ln=struct.unpack_from(">I",b,o)[0]; o+=4
        return b[o:o+ln],o+ln
    if t==Gp.Hashtable:
        cnt=struct.unpack_from(">H",b,o)[0]; o+=2; d={}
        for _ in range(cnt):
            k,o=deser_value(b,o); v,o=deser_value(b,o); d[k]=v
        return d,o
    if t==Gp.StringArray:
        cnt=struct.unpack_from(">H",b,o)[0]; o+=2; arr=[]
        for _ in range(cnt):
            ln=struct.unpack_from(">H",b,o)[0]; o+=2; arr.append(b[o:o+ln].decode('utf-8','replace')); o+=ln
        return arr,o
    if t==Gp.IntArray:
        cnt=struct.unpack_from(">I",b,o)[0]; o+=4
        arr=list(struct.unpack_from(">%di"%cnt,b,o)); return arr,o+4*cnt
    if t==Gp.Array:  # typed array: count(2), elementType(1), elements
        cnt=struct.unpack_from(">H",b,o)[0]; o+=2; et=b[o]; o+=1; arr=[]
        for _ in range(cnt):
            # element has no per-item type byte; decode by et
            v,o=deser_typed(b,o,et); arr.append(v)
        return arr,o
    if t==Gp.ObjectArray:
        cnt=struct.unpack_from(">H",b,o)[0]; o+=2; arr=[]
        for _ in range(cnt):
            v,o=deser_value(b,o); arr.append(v)
        return arr,o
    if t==Gp.Dictionary:  # keyType(1),valType(1),count(2),entries
        kt=b[o]; vt=b[o+1]; o+=2; cnt=struct.unpack_from(">H",b,o)[0]; o+=2; d={}
        for _ in range(cnt):
            k,o=deser_typed(b,o,kt) if kt!=0 else deser_value(b,o)
            v,o=deser_typed(b,o,vt) if vt!=0 else deser_value(b,o)
            d[k]=v
        return d,o
    raise ValueError("deser type %d @%d"%(t,o-1))

def deser_typed(b,o,t):
    """Deserialize a value whose type is already known (no leading type byte)."""
    if t==Gp.Byte: return b[o],o+1
    if t==Gp.Boolean: return (b[o]!=0),o+1
    if t==Gp.Short: return struct.unpack_from(">h",b,o)[0],o+2
    if t==Gp.Integer: return struct.unpack_from(">i",b,o)[0],o+4
    if t==Gp.Long: return struct.unpack_from(">q",b,o)[0],o+8
    if t==Gp.Float: return struct.unpack_from(">f",b,o)[0],o+4
    if t==Gp.Double: return struct.unpack_from(">d",b,o)[0],o+8
    if t==Gp.String:
        ln=struct.unpack_from(">H",b,o)[0]; o+=2; return b[o:o+ln].decode('utf-8','replace'),o+ln
    # fall back to full typed value (has its own type byte)
    return deser_value(b,o)

def parse_op_request(payload):
    # payload starts after f3 02
    opcode=payload[0]; o=1
    cnt=struct.unpack_from(">H",payload,o)[0]; o+=2
    params={}
    for _ in range(cnt):
        k=payload[o]; o+=1
        v,o=deser_value(payload,o)
        params[k]=v
    return opcode,params

# ---------------- eNet peer/server ----------------
class Peer:
    def __init__(self,addr,challenge):
        self.addr=addr; self.challenge=challenge
        self.peerId=0x0001
        self.out_seq={0:0,0xff:0}   # per-channel outgoing reliable seq
        self.last_client_sent=0
    def next_seq(self,chan):
        self.out_seq[chan]=self.out_seq.get(chan,0)+1
        return self.out_seq[chan]

class PhotonServer:
    def __init__(self,role,port,master_addr=None,verbose=True):
        self.role=role; self.port=port; self.master_addr=master_addr; self.v=verbose
        self.peers={}
        self.sock=socket.socket(socket.AF_INET,socket.SOCK_DGRAM)
        self.sock.setsockopt(socket.SOL_SOCKET,socket.SO_REUSEADDR,1)
        self.sock.bind(("0.0.0.0",port))

    def log(self,*a):
        if self.v: print("[%s:%d]"%(self.role,self.port),*a,flush=True)

    def send_commands(self,peer,cmds):
        """cmds: list of (type,channel,flags,reliable_bool,payload)."""
        pkt=w_u16(peer.peerId)+bytes([0])+bytes([len(cmds)])+w_u32(now_ms())+w_u32(peer.challenge)
        body=b""
        for (ct,chan,flags,reliable,payload) in cmds:
            if reliable:
                seq=peer.next_seq(chan)
            else:
                seq=0
            size=12+len(payload)
            body+=bytes([ct,chan,flags,0])+w_u32(size)+w_u32(seq)+payload
        self.sock.sendto(pkt+body,peer.addr)

    def send_ack(self,peer,acked_seq,acked_sent,chan=0):
        payload=w_u32(acked_seq)+w_u32(acked_sent)
        # ack is one command; build directly (ack seq is 0)
        self.send_commands(peer,[(ACK,chan,0,False,payload)])

    def send_reliable_msg(self,peer,msg,chan=0):
        self.send_commands(peer,[(RELIABLE,chan,1,True,msg)])

    def handle(self,data,addr):
        if len(data)<12: return
        peerId=struct.unpack_from(">H",data,0)[0]
        cc=data[3]; sent=struct.unpack_from(">I",data,4)[0]; ch=struct.unpack_from(">I",data,8)[0]
        peer=self.peers.get(addr)
        o=12
        for _ in range(cc):
            if o+12>len(data): break
            ct=data[o]; chan=data[o+1]; cf=data[o+2]; size=struct.unpack_from(">I",data,o+4)[0]
            seq=struct.unpack_from(">I",data,o+8)[0]
            payload=data[o+12:o+size]; o+=size
            if ct not in (RELIABLE,UNRELIABLE,ACK):
                self.log("  cmd type=%d chan=%d seq=%d"%(ct,chan,seq))
            if ct==CONNECT:
                peer=Peer(addr,ch); self.peers[addr]=peer
                self.log("CONNECT from",addr,"challenge=0x%08x"%ch)
                # VerifyConnect payload = assignedPeerId(2) + echo connect[2:16] + zero pad to 32
                vc=w_u16(peer.peerId)+payload[2:16]+b"\x00"*(32-2-14)
                self.send_commands(peer,[
                    (ACK,0xff,0,False,w_u32(seq)+w_u32(sent)),
                    (VERIFYCONNECT,0xff,1,True,vc),
                ])
            elif ct==DISCONNECT:
                self.log("DISCONNECT",addr); self.peers.pop(addr,None); return
            elif ct==PING:
                if peer: self.send_ack(peer,seq,sent,chan)
            elif ct==ACK:
                pass  # client acking our reliable; ignore (no retransmit)
            elif ct in (RELIABLE,UNRELIABLE):
                if not peer: continue
                if ct==RELIABLE:
                    self.send_ack(peer,seq,sent,chan)
                self.on_message(peer,payload,chan)
            elif peer and (cf & 1):
                # Any other reliable command (e.g. type 12 keepalive/timestamp
                # ping) must be acknowledged or the client retransmits then drops.
                self.send_ack(peer,seq,sent,chan)

    def on_message(self,peer,payload,chan):
        if not payload: return
        if payload[0]!=0xf3:
            self.log("non-f3 msg:",payload[:8].hex()); return
        mtype=payload[1]
        if mtype==0x00:   # Init request
            self.log("Init request raw=",payload.hex())
            self.send_reliable_msg(peer,make_init_response(),chan=0)
        elif mtype==0x02: # OperationRequest (plaintext)
            try:
                opcode,params=parse_op_request(payload[2:])
            except Exception as e:
                self.log("op parse err",e,payload.hex()); return
            self.log("OP req opcode=%d params=%s"%(opcode,fmt_params(params)))
            self.on_operation(peer,opcode,params)
        elif mtype==0x06:  # InitEncryption / DH key exchange request
            self.log("DH init (f3 06); replaying canned f3 07 response")
            self.send_reliable_msg(peer,CANNED_DH_RESPONSE,chan=0)
        elif mtype in (0x82,0x83):
            self.log("ENCRYPTED op msg type=0x%02x (IsEncryptionAvailable not false?) %s"%(mtype,payload[:24].hex()))
        else:
            self.log("msg type=0x%02x %s"%(mtype,payload[:16].hex()))

    def on_operation(self,peer,opcode,params):
        # Photon LoadBalancing ParameterCodes
        P_ADDRESS=230; P_SECRET=221; P_USERID=225
        if opcode==230:  # OpAuthenticate
            if self.role=="nameserver":
                # Return the master server address for the requested region.
                resp={P_ADDRESS:self.master_addr, P_SECRET:"svt-token"}
                self.log("  -> auth OK, master=%s"%self.master_addr)
                self.send_reliable_msg(peer,make_op_response(230,0,None,resp))
            else:  # master/game: auth success -> OnConnectedToMaster
                resp={P_SECRET:"svt-token", P_USERID:params.get(225,"user")}
                self.log("  -> master auth OK (OnConnectedToMaster)")
                self.send_reliable_msg(peer,make_op_response(230,0,None,resp))
        elif opcode==0:  # ChatOp Subscribe (chat peer shares this server)
            chans=params.get(0,[]) or []
            if isinstance(chans,str): chans=[chans]
            chans=list(chans)
            self.log("  -> chat subscribe OK: %s"%chans)
            # OperationResponse for the OpSubscribe
            self.send_reliable_msg(peer,make_op_response(0,0,None,
                {0:StrArray(chans),15:BoolArray([True]*len(chans))}))
            # Subscribe EVENT (ChatEventCode.Subscribe=5) — this is what actually
            # drives ChatClient.HandleSubscribeEvent -> OnSubscribed.
            self.send_reliable_msg(peer,make_event(5,
                {0:StrArray(chans),15:BoolArray([True]*len(chans))}))
        elif opcode==229:  # OpJoinLobby -> OnJoinedLobby
            self.log("  -> JoinLobby OK (OnJoinedLobby)")
            self.send_reliable_msg(peer,make_op_response(229,0))
        elif opcode==220:  # OpGetRegions
            # StringArray of region codes (210) + StringArray of addresses (230)
            resp={210:StrArray(["eu"]),230:StrArray([self.master_addr])}
            self.log("  -> GetRegions -> %s"%self.master_addr)
            self.send_reliable_msg(peer,make_op_response(220,0,None,resp))
        else:
            self.log("  -> generic OK for opcode %d"%opcode)
            self.send_reliable_msg(peer,make_op_response(opcode,0))

    def serve(self,seconds=None):
        self.log("listening")
        self.sock.settimeout(1.0)
        t0=time.time()
        while True:
            try:
                data,addr=self.sock.recvfrom(4096)
                self.handle(data,addr)
            except socket.timeout:
                pass
            if seconds and time.time()-t0>seconds: break

def fmt_params(p):
    out={}
    for k,v in p.items():
        if isinstance(v,(bytes,)): v="bytes[%d]"%len(v)
        elif isinstance(v,Byte): v="B(%d)"%v.v
        out[k]=v
    return out

if __name__=="__main__":
    ap=argparse.ArgumentParser()
    ap.add_argument("--role",default="master")
    ap.add_argument("--port",type=int,default=5055)
    ap.add_argument("--seconds",type=int,default=0)
    ap.add_argument("--master-addr",default="10.0.2.2:5055")
    a=ap.parse_args()
    srv=PhotonServer(a.role,a.port,master_addr=a.master_addr)
    srv.serve(a.seconds or None)
