import struct, collections
path='/tmp/r443.pcap'
f=open(path,'rb'); gh=f.read(24)
if len(gh)<24: print('EMPTY PCAP'); raise SystemExit
b=gh[:4]
end={'\xd4\xc3\xb2\xa1':'<','\xa1\xb2\xc3\xd4':'>','\x4d\x3c\xb2\xa1':'<','\xa1\xb2\x3c\x4d':'>'}.get(b.decode('latin1'))
if not end: print('bad magic',b.hex()); raise SystemExit
link=struct.unpack(end+'I',gh[20:24])[0]
print('linktype=%d' % link)
KNOWN=bytes.fromhex('b2c987a58c492585')
recs=[]
np=0
while True:
    ph=f.read(16)
    if len(ph)<16: break
    _,_,incl,_=struct.unpack(end+'IIII',ph)
    d=f.read(incl)
    if len(d)<incl: break
    np+=1
    if link==1:
        if len(d)<14: continue
        et=struct.unpack('>H',d[12:14])[0]; o=14
        if et==0x8100: et=struct.unpack('>H',d[16:18])[0]; o=18
    elif link==101: et=0x0800; o=0
    else: continue
    if et==0x0800:
        if len(d)<o+20: continue
        ihl=(d[o]&0xf)*4
        if d[o+9]!=6: continue
        src='.'.join(str(x) for x in d[o+12:o+16]); t=o+ihl
    elif et==0x86DD:
        if len(d)<o+40: continue
        if d[o+6]!=6: continue
        import ipaddress; src=str(ipaddress.IPv6Address(d[o+8:o+24])); t=o+40
    else: continue
    if len(d)<t+20: continue
    sport,dport=struct.unpack('>HH',d[t:t+4])
    doff=(d[t+12]>>4)*4
    p=d[t+doff:]
    if len(p)<6 or p[0]!=0x16 or p[5]!=0x01: continue
    hslen=int.from_bytes(p[6:9],'big'); ch=p[9:9+hslen]
    i=34
    if i>=len(ch): continue
    sl=ch[i]; i+=1; sid=ch[i:i+sl]; i+=sl
    csl=int.from_bytes(ch[i:i+2],'big'); i+=2; cs=ch[i:i+csl]; i+=csl
    cl=ch[i]; i+=1+cl
    if i+2>len(ch): continue
    el=int.from_bytes(ch[i:i+2],'big'); i+=2; endp=min(i+el,len(ch))
    exts=[]; sni=''; groups=[]; ks=[]; alpn=[]; vers=[]
    while i+4<=endp:
        etp=int.from_bytes(ch[i:i+2],'big'); edl=int.from_bytes(ch[i+2:i+4],'big'); i+=4
        ed=ch[i:i+edl]; i+=edl
        exts.append(etp)
        if etp==0 and len(ed)>=5:
            nl=int.from_bytes(ed[3:5],'big'); sni=ed[5:5+nl].decode('latin1',errors='replace')
        elif etp==10 and len(ed)>=2:
            n=int.from_bytes(ed[:2],'big'); groups=[int.from_bytes(ed[2+j*2:4+j*2],'big') for j in range(n//2)]
        elif etp==51 and len(ed)>=2:
            n=int.from_bytes(ed[:2],'big'); j=2
            while j+4<=2+n:
                g=int.from_bytes(ed[j:j+2],'big'); L=int.from_bytes(ed[j+2:j+4],'big'); j+=4
                ks.append((g,ed[j:j+L].hex())); j+=L
        elif etp==16 and len(ed)>=2:
            n=int.from_bytes(ed[:2],'big'); j=2
            while j+1<=1+n:
                L=ed[j]; j+=1; alpn.append(ed[j:j+L].decode('latin1',errors='replace')); j+=L
        elif etp==43:
            vers=[ed[k] for k in range(0,len(ed),1) if ed[k]!=0] if len(ed)<8 else [int.from_bytes(ed[k:k+2],'big') for k in range(1,len(ed),2)]
    pos=sid.find(KNOWN)
    recs.append(dict(src=src,sport=sport,sni=sni,sid=sid.hex(),pos=pos,groups=groups,ks=ks,
                     ncipher=len(cs)//2,exts=exts,alpn=alpn,vers=vers))
print('packets=%d  clienthellos=%d' % (np,len(recs)))
groups=collections.Counter((r['sni'],r['sid'],tuple(r['ks'])) for r in recs)
print('\n=== unique (SNI | session_id | key_share) ===')
for k,c in groups.most_common():
    sni,sid,ks=k
    print('count=%d' % c)
    print('  SNI          = %s' % sni)
    print('  session_id   = %s' % sid)
    if ks:
        for g,hx in ks: print('  key_share    = group 0x%04x len=%d %s' % (g,len(hx)//2,hx[:64]))
    r=next(r for r in recs if (r['sni'],r['sid'],tuple(r['ks']))==k)
    print('  groups       = %s' % [hex(x) for x in r['groups']])
    print('  ciphers=%d alpn=%s exts=%s' % (r['ncipher'],r['alpn'],r['exts']))
    ports=[str(x['sport']) for x in recs if (x['sni'],x['sid'],tuple(x['ks']))==k][:8]
    print('  src ports    = %s' % ','.join(ports))
    raw=bytes.fromhex(sid)
    print('  known short_id b2c987a58c492585 at offset = %s' % (raw.find(KNOWN) if raw.find(KNOWN)>=0 else 'NOT FOUND'))
    print()
