You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 

285 lines
14 KiB

// stcp_client.c — STCP client implementation
#include "stcp_client.h"
#include "secure_channel.h"
#include "crc32.h"
#include "../lib/u_async.h"
#include "../lib/socket_compat.h"
#include "../lib/ll_queue.h"
#include "../lib/mem.h"
#include "../lib/debug_config.h"
#include "../lib/platform_compat.h"
#include <stdlib.h>
#include <string.h>
#include <errno.h>
#ifndef _WIN32
#include <unistd.h>
#include <fcntl.h>
#include <netinet/tcp.h>
#endif
struct stcp_client {
struct stcp_conn conn;
struct UASYNC *ua;
stcp_ready_cb ready_cb;
void *ready_arg;
uint8_t peer_pubkey[SC_PUBKEY_SIZE];
};
static void client_connect_write_cb(socket_t sock, void *arg);
static void client_conn_read_cb(socket_t sock, void *arg);
static void client_conn_write_cb(socket_t sock, void *arg);
static void client_tx_queue_cb(struct ll_queue *q, void *arg);
static void stcp_conn_process_recv(struct stcp_conn *c);
static int client_try_send(struct stcp_conn *c, const uint8_t *data, size_t len);
static void client_do_close(struct stcp_conn *c, int err);
static int encrypt_and_crc(struct stcp_conn *c, const uint8_t *data, size_t data_len,
uint8_t *output, size_t *output_len) {
uint32_t crc = crc32_calc(data, data_len);
output[0] = (uint8_t)(data_len >> 0);
output[1] = (uint8_t)(data_len >> 8);
memcpy(output + 2, data, data_len);
output[2 + data_len + 0] = (uint8_t)(crc >> 0);
output[2 + data_len + 1] = (uint8_t)(crc >> 8);
output[2 + data_len + 2] = (uint8_t)(crc >> 16);
output[2 + data_len + 3] = (uint8_t)(crc >> 24);
if (sc_stream_xor(&c->stream_send, output + 2, data_len + 4) != SC_OK) {
DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stream_xor failed len=%zu", data_len);
return -1;
}
*output_len = 2 + data_len + 4;
return 0;
}
static int decrypt_and_check(uint8_t *data, size_t len, struct sc_stream_state *stream, size_t *out_len) {
if (sc_stream_xor(stream, data, len) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stream_xor failed"); return -1; }
size_t data_len = len - SC_CRC32_SIZE;
uint32_t recv_crc = ((uint32_t)data[data_len]) | ((uint32_t)data[data_len + 1] << 8) |
((uint32_t)data[data_len + 2] << 16) | ((uint32_t)data[data_len + 3] << 24);
uint32_t calc_crc = crc32_calc(data, data_len);
if (recv_crc != calc_crc) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "CRC mismatch"); return -1; }
*out_len = data_len;
return 0;
}
static void client_do_close(struct stcp_conn *c, int err) {
if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return;
c->state = STCP_STATE_CLOSED;
if (c->socket_id) { uasync_remove_socket_t(c->ua, c->sock); c->socket_id = NULL; }
if (c->sock != SOCKET_INVALID) { socket_close_wrapper(c->sock); c->sock = SOCKET_INVALID; }
if (c->recv_buf) { u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; }
if (c->send_buf) { u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; }
if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; cb(c, err, c->close_arg); }
}
static int client_try_send(struct stcp_conn *c, const uint8_t *data, size_t len) {
if (c->sock == SOCKET_INVALID) return -1;
if (c->send_buf) return -1;
ssize_t sent = send(c->sock, data, len, 0);
if (sent < 0) {
int e = socket_get_error();
if (e == ERR_AGAIN || e == ERR_WOULDBLOCK) { c->send_buf = (uint8_t *)data; c->send_len = len; c->send_offset = 0; uasync_set_socket_write(c->ua, c->socket_id, 1); return 0; }
DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "send failed err=%d", e);
u_free((uint8_t *)data); client_do_close(c, e); return -1;
}
if ((size_t)sent < len) { c->send_buf = (uint8_t *)data; c->send_len = len; c->send_offset = (size_t)sent; uasync_set_socket_write(c->ua, c->socket_id, 1); return 0; }
u_free((uint8_t *)data);
return 0;
}
static void client_conn_write_cb(socket_t sock, void *arg) {
struct stcp_conn *c = (struct stcp_conn *)arg;
if (!c->send_buf) { uasync_set_socket_write(c->ua, c->socket_id, 0); return; }
ssize_t sent = send(sock, c->send_buf + c->send_offset, c->send_len - c->send_offset, 0);
if (sent < 0) {
int e = socket_get_error();
if (e == ERR_AGAIN || e == ERR_WOULDBLOCK) return;
client_do_close(c, e); return;
}
c->send_offset += (size_t)sent;
if (c->send_offset >= c->send_len) { u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; c->send_offset = 0; uasync_set_socket_write(c->ua, c->socket_id, 0); }
}
static void client_tx_queue_cb(struct ll_queue *q, void *arg) {
struct stcp_conn *c = (struct stcp_conn *)arg;
struct ll_entry *e = queue_data_get(q);
if (!e) { queue_resume_callback(q); return; }
if (c->state == STCP_STATE_DATA && e->dgram) {
size_t need;
uint8_t *buf = u_malloc(2 + e->len + SC_CRC32_SIZE);
if (buf && !encrypt_and_crc(c, e->dgram, e->len, buf, &need)) client_try_send(c, buf, need);
else if (buf) u_free(buf);
}
queue_dgram_free(e);
queue_entry_free(e);
queue_resume_callback(q);
}
static int client_derive_session(struct stcp_conn *c, const uint8_t *peer_pubkey) {
struct secure_channel sc;
sc_init_ctx(&sc, &c->my_keys);
if (sc_set_peer_public_key(&sc, peer_pubkey, SC_PEER_PUBKEY_BIN) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "client ECDH failed"); return -1; }
memcpy(c->session_key, sc.session_key, SC_SESSION_KEY_SIZE);
if (sc_stream_init(&sc, &c->stream_send, STCP_STREAM_CLIENT_SEND) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "client stream_send init failed"); return -1; }
if (sc_stream_init(&sc, &c->stream_recv, STCP_STREAM_SERVER_SEND) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "client stream_recv init failed"); return -1; }
return 0;
}
static void client_send_handshake(struct stcp_conn *c, const uint8_t *server_pubkey) {
uint8_t salt[SC_PUBKEY_ENC_SALT_SIZE];
if (random_bytes(salt, SC_PUBKEY_ENC_SALT_SIZE) != 0) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "client random_bytes failed"); client_do_close(c, 1); return; }
uint16_t padding = 8;
size_t total = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT + padding;
uint8_t *hs = u_malloc(total);
if (!hs) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "client hs malloc failed"); client_do_close(c, 1); return; }
memcpy(hs, salt, SC_PUBKEY_ENC_SALT_SIZE);
sc_obfuscate_pubkey(salt, server_pubkey, c->my_keys.public_key, hs + SC_PUBKEY_ENC_SALT_SIZE);
uint8_t plain[2] = {(uint8_t)padding, (uint8_t)(padding >> 8)};
uint32_t crc = crc32_calc(plain, 2);
uint8_t *enc_dst = hs + SC_PUBKEY_ENC_SIZE;
enc_dst[0] = plain[0]; enc_dst[1] = plain[1];
enc_dst[2] = (uint8_t)(crc >> 0); enc_dst[3] = (uint8_t)(crc >> 8);
enc_dst[4] = (uint8_t)(crc >> 16); enc_dst[5] = (uint8_t)(crc >> 24);
if (sc_stream_xor(&c->stream_send, enc_dst, STCP_HS_ENC_CLIENT) != SC_OK) { u_free(hs); client_do_close(c, 2); return; }
for (int i = 0; i < padding; i++) hs[SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT + i] = (uint8_t)(salt[0] ^ i);
c->state = STCP_STATE_HS_CLIENT_SENT;
client_try_send(c, hs, total);
}
static void process_server_response(struct stcp_conn *c) {
if (!c->hs_key_processed) {
const uint8_t *salt = c->recv_buf;
const uint8_t *enc_pubkey = salt + SC_PUBKEY_ENC_SALT_SIZE;
uint8_t server_pubkey[SC_PUBKEY_SIZE];
sc_obfuscate_pubkey(salt, c->my_keys.public_key, enc_pubkey, server_pubkey);
c->hs_key_processed = 1;
}
uint8_t enc_hs[STCP_HS_ENC_SERVER];
memcpy(enc_hs, c->recv_buf + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_SERVER);
size_t hs_data_len;
if (decrypt_and_check(enc_hs, STCP_HS_ENC_SERVER, &c->stream_recv, &hs_data_len)) { client_do_close(c, 3); return; }
uint8_t status = enc_hs[0];
uint16_t padding_size = (uint16_t)enc_hs[1] | ((uint16_t)enc_hs[2] << 8);
c->hs_expected_len = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + padding_size;
if (status != 0) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server handshake status=%d", status); client_do_close(c, 4); return; }
}
static void finish_server_response(struct stcp_conn *c) {
memmove(c->recv_buf, c->recv_buf + c->hs_expected_len, c->recv_buf_len - c->hs_expected_len);
c->recv_buf_len -= c->hs_expected_len;
c->hs_expected_len = 0;
c->hs_key_processed = 0;
c->state = STCP_STATE_DATA;
DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_client: handshake OK, entering DATA state");
if (c->on_ready) { void (*cb)(struct stcp_conn*, void*) = c->on_ready; c->on_ready = NULL; cb(c, c->ready_arg); }
}
static void stcp_conn_process_recv(struct stcp_conn *c) {
while (c->recv_buf_len > 0) {
switch (c->state) {
case STCP_STATE_HS_CLIENT_SENT:
if (!c->hs_key_processed) {
if (c->recv_buf_len < SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER) return;
process_server_response(c);
if (c->state != STCP_STATE_HS_CLIENT_SENT) return;
}
if (c->hs_key_processed && c->recv_buf_len >= c->hs_expected_len) { finish_server_response(c); return; }
return;
case STCP_STATE_DATA: {
if (c->recv_buf_len < 2) return;
uint16_t msg_size = (uint16_t)c->recv_buf[0] | ((uint16_t)c->recv_buf[1] << 8);
if (msg_size > STCP_MAX_MSG_SIZE) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "client msg too large: %u", msg_size); client_do_close(c, 5); return; }
size_t total = 2 + msg_size + SC_CRC32_SIZE;
if (c->recv_buf_len < total) return;
uint8_t *enc_data = c->recv_buf + 2;
size_t data_len;
if (decrypt_and_check(enc_data, msg_size + SC_CRC32_SIZE, &c->stream_recv, &data_len)) { client_do_close(c, 6); return; }
if (c->rx_queue) {
struct ll_entry *e = queue_entry_new(0);
if (e) { e->dgram = u_malloc(data_len); if (e->dgram) { if (data_len) memcpy(e->dgram, enc_data, data_len); e->len = (uint16_t)data_len; queue_data_put(c->rx_queue, e); } else { queue_entry_free(e); } }
}
memmove(c->recv_buf, c->recv_buf + total, c->recv_buf_len - total);
c->recv_buf_len -= total;
break;
}
default: return;
}
}
}
static void client_conn_read_cb(socket_t sock, void *arg) {
struct stcp_conn *c = (struct stcp_conn *)arg;
if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return;
if (!c->recv_buf) { c->recv_buf_cap = STCP_RECV_BUF_INIT; c->recv_buf = u_malloc(c->recv_buf_cap); if (!c->recv_buf) { client_do_close(c, ENOMEM); return; } }
if (c->recv_buf_len + 4096 > c->recv_buf_cap) {
size_t nc = c->recv_buf_cap * 2; if (nc > STCP_RECV_BUF_MAX) nc = STCP_RECV_BUF_MAX;
if (nc <= c->recv_buf_cap) { client_do_close(c, ENOBUFS); return; }
uint8_t *nb = u_realloc(c->recv_buf, nc); if (!nb) { client_do_close(c, ENOMEM); return; }
c->recv_buf = nb; c->recv_buf_cap = nc;
}
ssize_t n = recv(sock, c->recv_buf + c->recv_buf_len, c->recv_buf_cap - c->recv_buf_len, 0);
if (n < 0) { int e = socket_get_error(); if (e == ERR_AGAIN || e == ERR_WOULDBLOCK) return; client_do_close(c, e); return; }
if (n == 0) { client_do_close(c, 0); return; }
c->recv_buf_len += (size_t)n;
stcp_conn_process_recv(c);
}
static void client_connect_write_cb(socket_t sock, void *arg) {
struct stcp_client *cli = (struct stcp_client *)arg;
struct stcp_conn *c = &cli->conn;
int err = 0;
socklen_t len = sizeof(err);
if (getsockopt(sock, SOL_SOCKET, SO_ERROR, (char *)&err, &len) < 0 || err != 0) {
DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_client connect failed err=%d", err);
client_do_close(c, err); return;
}
uasync_remove_socket_t(cli->ua, sock);
c->socket_id = uasync_add_socket_t(cli->ua, sock, client_conn_read_cb, client_conn_write_cb, NULL, c);
if (!c->socket_id) { client_do_close(c, ENOMEM); return; }
int opt = 1; setsockopt(sock, IPPROTO_TCP, TCP_NODELAY, (const char *)&opt, sizeof(opt));
if (client_derive_session(c, cli->peer_pubkey)) { client_do_close(c, 1); return; }
client_send_handshake(c, cli->peer_pubkey);
}
struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uint16_t port,
struct SC_MYKEYS *keys, const uint8_t *peer_pubkey,
stcp_ready_cb ready_cb, void *arg,
stcp_close_cb close_cb, void *close_arg) {
if (!ua || !addr || !keys || !peer_pubkey || !ready_cb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "invalid args"); return NULL; }
struct stcp_client *cli = u_calloc(1, sizeof(struct stcp_client));
if (!cli) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "calloc failed"); return NULL; }
cli->ua = ua; cli->ready_cb = ready_cb; cli->ready_arg = arg;
memcpy(cli->peer_pubkey, peer_pubkey, SC_PUBKEY_SIZE);
struct stcp_conn *c = &cli->conn;
c->ua = ua; c->state = STCP_STATE_INIT; c->is_server = 0; c->my_keys = *keys;
c->on_ready = ready_cb; c->ready_arg = arg;
c->on_close = close_cb; c->close_arg = close_arg;
c->tx_cb = client_tx_queue_cb;
c->sock = socket(AF_INET, SOCK_STREAM, 0);
if (c->sock == SOCKET_INVALID) { u_free(cli); return NULL; }
socket_set_nonblocking(c->sock);
struct sockaddr_in saddr; memset(&saddr, 0, sizeof(saddr));
saddr.sin_family = AF_INET; saddr.sin_port = htons(port);
if (inet_pton(AF_INET, addr, &saddr.sin_addr) != 1) { socket_close_wrapper(c->sock); u_free(cli); return NULL; }
if (connect(c->sock, (struct sockaddr *)&saddr, sizeof(saddr)) < 0 && socket_get_error() != EINPROGRESS) {
DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_client connect to %s:%u failed err=%d", addr, port, socket_get_error());
socket_close_wrapper(c->sock); u_free(cli); return NULL;
}
c->socket_id = uasync_add_socket_t(ua, c->sock, NULL, client_connect_write_cb, NULL, cli);
if (!c->socket_id) { socket_close_wrapper(c->sock); u_free(cli); return NULL; }
return cli;
}
void stcp_client_destroy(struct stcp_client *cli) {
if (!cli) return;
client_do_close(&cli->conn, 0);
stcp_conn_free(&cli->conn);
u_free(cli);
}
struct stcp_conn *stcp_client_get_conn(struct stcp_client *cli) {
return cli ? &cli->conn : NULL;
}