From f913bdf351dddcb53ec550d596a04d8f96a42c90 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Sun, 14 Jun 2026 10:00:02 +0300 Subject: [PATCH] stcp: streaming TCP module + stream cipher + comprehensive tests - secure_channel: add sc_stream_state + sc_stream_init/sc_stream_xor/sc_stream_cleanup AES-128-CTR streaming cipher for incremental encrypt/decrypt - stcp.h/c: shared stcp_conn lifecycle, set_on_close, rx/tx queue setup - stcp_server: listen/accept, handshake (obfuscated pubkey + ECDH + stream), DATA phase - stcp_client: connect, handshake, DATA phase - Protocol: key_block(72) + enc(data+crc) + padding for handshake size(2 LE) + enc(data+crc) for data messages - Tests: test_stream_cipher (AES-CTR correctness) test_stcp (6 cases: sizes 0..65535, 200 msgs, wrong key, close detection, multi-client, interleaved send/recv) - Fix: 0-length message handling in tx_queue callbacks --- src/Makefile.am | 3 + src/secure_channel.c | 100 ++++++++++ src/secure_channel.h | 19 ++ src/stcp.c | 35 ++++ src/stcp.h | 80 ++++++++ src/stcp_client.c | 282 ++++++++++++++++++++++++++++ src/stcp_client.h | 15 ++ src/stcp_server.c | 375 +++++++++++++++++++++++++++++++++++++ src/stcp_server.h | 14 ++ tests/Makefile.am | 10 + tests/test_stcp.c | 359 +++++++++++++++++++++++++++++++++++ tests/test_stream_cipher.c | 141 ++++++++++++++ 12 files changed, 1433 insertions(+) create mode 100644 src/stcp.c create mode 100644 src/stcp.h create mode 100644 src/stcp_client.c create mode 100644 src/stcp_client.h create mode 100644 src/stcp_server.c create mode 100644 src/stcp_server.h create mode 100644 tests/test_stcp.c create mode 100644 tests/test_stream_cipher.c diff --git a/src/Makefile.am b/src/Makefile.am index a3799146..5efc41f8 100644 --- a/src/Makefile.am +++ b/src/Makefile.am @@ -26,6 +26,9 @@ utun_CORE_SOURCES = \ etcp_dump.c \ secure_channel.c \ crc32.c \ + stcp.c \ + stcp_server.c \ + stcp_client.c \ pkt_normalizer.c \ packet_dump.c \ etcp_api.c \ diff --git a/src/secure_channel.c b/src/secure_channel.c index 95ae0128..c1e73dfe 100644 --- a/src/secure_channel.c +++ b/src/secure_channel.c @@ -35,6 +35,7 @@ #include "../tinycrypt/lib/include/tinycrypt/ecc_dh.h" #include "../tinycrypt/lib/include/tinycrypt/aes.h" #include "../tinycrypt/lib/include/tinycrypt/ccm_mode.h" +#include "../tinycrypt/lib/include/tinycrypt/ctr_mode.h" #include "../tinycrypt/lib/include/tinycrypt/constants.h" #include "../tinycrypt/lib/include/tinycrypt/ecc_platform_specific.h" #include "../tinycrypt/lib/include/tinycrypt/sha256.h" @@ -142,6 +143,23 @@ static void sc_derive_session_key(const uint8_t *shared_secret, uint8_t *session memcpy(session_key, hash, SC_SESSION_KEY_SIZE); } +// Derive 12-byte nonce for stream cipher from session key + stream_id +static void sc_stream_derive_nonce(const uint8_t *session_key, uint32_t stream_id, uint8_t *nonce_out) { + SC_SHA256_CTX sha_ctx; + uint8_t hash[SC_HASH_SIZE]; + uint8_t stream_id_le[4]; + stream_id_le[0] = (uint8_t)(stream_id >> 0); + stream_id_le[1] = (uint8_t)(stream_id >> 8); + stream_id_le[2] = (uint8_t)(stream_id >> 16); + stream_id_le[3] = (uint8_t)(stream_id >> 24); + sc_sha256_init(&sha_ctx); + sc_sha256_update(&sha_ctx, session_key, SC_SESSION_KEY_SIZE); + sc_sha256_update(&sha_ctx, stream_id_le, sizeof(stream_id_le)); + sc_sha256_update(&sha_ctx, (const uint8_t *)"uTun3-stream", 12); + sc_sha256_final(&sha_ctx, hash); + memcpy(nonce_out, hash, SC_STREAM_NONCE_SIZE); +} + #ifdef USE_OPENSSL // OpenSSL-specific implementations @@ -513,6 +531,50 @@ sc_status_t sc_compute_public_key_from_private(const uint8_t *private_key, uint8 return SC_OK; } +// --- OpenSSL streaming cipher (AES-128-CTR) --- + +sc_status_t sc_stream_init(sc_context_t *ctx, struct sc_stream_state *state, uint32_t stream_id) { + if (!ctx || !state) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: invalid args"); return SC_ERR_INVALID_ARG; } + if (!ctx->session_ready) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: session not ready"); return SC_ERR_NOT_INITIALIZED; } + uint8_t nonce[SC_STREAM_NONCE_SIZE]; + sc_stream_derive_nonce(ctx->session_key, stream_id, nonce); + uint8_t iv[16]; + memcpy(iv, nonce, SC_STREAM_NONCE_SIZE); + memset(iv + SC_STREAM_NONCE_SIZE, 0, 4); + EVP_CIPHER_CTX *ectx = EVP_CIPHER_CTX_new(); + if (!ectx) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: EVP_CIPHER_CTX_new failed"); return SC_ERR_CRYPTO; } + if (EVP_EncryptInit_ex(ectx, EVP_aes_128_ctr(), NULL, ctx->session_key, iv) != 1) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: EVP_EncryptInit_ex failed"); + EVP_CIPHER_CTX_free(ectx); + return SC_ERR_CRYPTO; + } + state->ectx = ectx; + state->initialized = 1; + DEBUG_INFO(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: stream_id=%u nonce=%02x%02x%02x%02x...", + stream_id, nonce[0], nonce[1], nonce[2], nonce[3]); + return SC_OK; +} + +sc_status_t sc_stream_xor(struct sc_stream_state *state, uint8_t *data, size_t data_len) { + if (data_len == 0) return SC_OK; + if (!state || !data) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: invalid args"); return SC_ERR_INVALID_ARG; } + if (!state->initialized) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: not initialized"); return SC_ERR_NOT_INITIALIZED; } + int outlen; + if (EVP_EncryptUpdate((EVP_CIPHER_CTX *)state->ectx, data, &outlen, data, (int)data_len) != 1 + || (size_t)outlen != data_len) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: EVP_EncryptUpdate failed outlen=%d datalen=%zu", outlen, data_len); + return SC_ERR_CRYPTO; + } + return SC_OK; +} + +void sc_stream_cleanup(struct sc_stream_state *state) { + if (!state || !state->ectx) return; + EVP_CIPHER_CTX_free((EVP_CIPHER_CTX *)state->ectx); + state->ectx = NULL; + state->initialized = 0; +} + #else // Original TinyCrypt implementations (unchanged logic) @@ -816,6 +878,44 @@ sc_status_t sc_compute_public_key_from_private(const uint8_t *private_key, uint8 return SC_OK; } +// --- TinyCrypt streaming cipher (AES-128-CTR) --- + +sc_status_t sc_stream_init(sc_context_t *ctx, struct sc_stream_state *state, uint32_t stream_id) { + if (!ctx || !state) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: invalid args"); return SC_ERR_INVALID_ARG; } + if (!ctx->session_ready) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: session not ready"); return SC_ERR_NOT_INITIALIZED; } + uint8_t nonce[SC_STREAM_NONCE_SIZE]; + sc_stream_derive_nonce(ctx->session_key, stream_id, nonce); + memcpy(state->ctr_block, nonce, SC_STREAM_NONCE_SIZE); + memset(state->ctr_block + SC_STREAM_NONCE_SIZE, 0, 4); + if (tc_aes128_set_encrypt_key((TCAesKeySched_t)state->sched_buf, ctx->session_key) != TC_CRYPTO_SUCCESS) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: tc_aes128_set_encrypt_key failed"); + return SC_ERR_CRYPTO; + } + state->initialized = 1; + DEBUG_INFO(DEBUG_CATEGORY_CRYPTO, "sc_stream_init: stream_id=%u nonce=%02x%02x%02x%02x...", + stream_id, nonce[0], nonce[1], nonce[2], nonce[3]); + return SC_OK; +} + +sc_status_t sc_stream_xor(struct sc_stream_state *state, uint8_t *data, size_t data_len) { + if (data_len == 0) return SC_OK; + if (!state || !data) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: invalid args"); return SC_ERR_INVALID_ARG; } + if (!state->initialized) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: not initialized"); return SC_ERR_NOT_INITIALIZED; } + if (tc_ctr_mode(data, (unsigned int)data_len, data, (unsigned int)data_len, + state->ctr_block, (TCAesKeySched_t)state->sched_buf) != TC_CRYPTO_SUCCESS) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "sc_stream_xor: tc_ctr_mode failed len=%zu", data_len); + return SC_ERR_CRYPTO; + } + return SC_OK; +} + +void sc_stream_cleanup(struct sc_stream_state *state) { + if (!state) return; + memset(state->sched_buf, 0, sizeof(state->sched_buf)); + memset(state->ctr_block, 0, sizeof(state->ctr_block)); + state->initialized = 0; +} + #endif sc_status_t sc_sha_transcode(const uint8_t *key, size_t key_len, uint8_t *data, size_t data_len) { diff --git a/src/secure_channel.h b/src/secure_channel.h index 165a705d..eaac7c46 100644 --- a/src/secure_channel.h +++ b/src/secure_channel.h @@ -75,4 +75,23 @@ sc_status_t sc_sha_transcode(const uint8_t *key, size_t key_len, uint8_t *data, // Obfuscate pubkey при передаче: XOR с SHA256(salt+peer_pubkey) || SHA256(peer_pubkey+salt) sc_status_t sc_obfuscate_pubkey(const uint8_t *salt, const uint8_t *peer_pubkey, const uint8_t *pubkey, uint8_t *output); +// --- Streaming cipher (AES-128-CTR, confidentiality only) --- + +#define SC_STREAM_NONCE_SIZE 12 +#define SC_STREAM_AES_SCHED_SIZE 176 // sizeof(struct tc_aes_key_sched_struct) = Nb*(Nr+1)*4 = 4*11*4 + +struct sc_stream_state { +#ifdef USE_OPENSSL + void *ectx; // EVP_CIPHER_CTX* (opaque, created/destroyed in .c) +#else + uint8_t sched_buf[SC_STREAM_AES_SCHED_SIZE]; // AES-128 key schedule + uint8_t ctr_block[16]; // nonce(12) + counter(4, BE) +#endif + uint8_t initialized; +}; + +sc_status_t sc_stream_init(sc_context_t *ctx, struct sc_stream_state *state, uint32_t stream_id); +sc_status_t sc_stream_xor(struct sc_stream_state *state, uint8_t *data, size_t data_len); +void sc_stream_cleanup(struct sc_stream_state *state); + #endif // SECURE_CHANNEL_H diff --git a/src/stcp.c b/src/stcp.c new file mode 100644 index 00000000..3ccb1a36 --- /dev/null +++ b/src/stcp.c @@ -0,0 +1,35 @@ +// stcp.c — shared stcp_conn lifecycle +#include "stcp.h" +#include "../lib/ll_queue.h" +#include "../lib/mem.h" +#include "../lib/debug_config.h" +#include +#include + +void stcp_conn_set_tx_queue(struct stcp_conn *c, struct ll_queue *q) { + if (!c || !q) return; + c->tx_queue = q; + if (c->tx_cb) queue_set_callback(q, c->tx_cb, c); +} + +void stcp_conn_set_rx_queue(struct stcp_conn *c, struct ll_queue *q) { + if (!c || !q) return; + c->rx_queue = q; +} + +void stcp_conn_set_on_close(struct stcp_conn *c, void (*cb)(struct stcp_conn *conn, int err, void *arg), void *arg) { + if (!c || !cb) return; + c->on_close = cb; + c->close_arg = arg; +} + +void stcp_conn_free(struct stcp_conn *c) { + if (!c) return; + 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; } + if (c->send_buf) { u_free(c->send_buf); c->send_buf = NULL; } + sc_stream_cleanup(&c->stream_send); + sc_stream_cleanup(&c->stream_recv); + if (c->allocated) u_free(c); +} diff --git a/src/stcp.h b/src/stcp.h new file mode 100644 index 00000000..caa8e11d --- /dev/null +++ b/src/stcp.h @@ -0,0 +1,80 @@ +// stcp.h — Streaming TCP: shared structures, constants, connection lifecycle +#ifndef STCP_H +#define STCP_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/debug_config.h" +#include + +#define STCP_MAX_MSG_SIZE 65535 +#define STCP_RECV_BUF_INIT 8192 +#define STCP_RECV_BUF_MAX 131072 +#define STCP_HS_TIMEOUT 5000 +#define STCP_CONNECT_TIMEOUT 10000 + +#define STCP_HS_CLIENT_MIN (SC_PUBKEY_ENC_SIZE + 6) // 72 + 2(padding_size enc) + 4(crc) = 78 +#define STCP_HS_SERVER_MIN (SC_PUBKEY_ENC_SIZE + 7) // 72 + 1(status) + 2(padding_size) + 4(crc) = 79 +#define STCP_HS_ENC_CLIENT 6 // padding_size(2) + CRC32(4) +#define STCP_HS_ENC_SERVER 7 // status(1) + padding_size(2) + CRC32(4) + +#define STCP_STREAM_CLIENT_SEND 0 +#define STCP_STREAM_SERVER_SEND 1 + +enum stcp_state { + STCP_STATE_INIT, + STCP_STATE_HS_CLIENT_SENT, + STCP_STATE_HS_SERVER_WAIT, + STCP_STATE_DATA, + STCP_STATE_CLOSED, + STCP_STATE_ERROR +}; + +struct stcp_conn { + socket_t sock; + struct UASYNC *ua; + void *socket_id; + + enum stcp_state state; + uint8_t is_server; + uint8_t allocated; // 1 = allocated by u_calloc, free in stcp_conn_free + + uint8_t session_key[SC_SESSION_KEY_SIZE]; + struct sc_stream_state stream_send; + struct sc_stream_state stream_recv; + + struct SC_MYKEYS my_keys; + uint8_t peer_pubkey[SC_PUBKEY_SIZE]; + uint8_t peer_pubkey_set; + + struct ll_queue *rx_queue; + struct ll_queue *tx_queue; + void (*tx_cb)(struct ll_queue *q, void *arg); + + uint8_t *recv_buf; + size_t recv_buf_len; + size_t recv_buf_cap; + + uint8_t *send_buf; + size_t send_len; + size_t send_offset; + + size_t hs_expected_len; + uint8_t hs_key_processed; + + void (*on_ready)(struct stcp_conn *conn, void *arg); + void *ready_arg; + + void (*on_close)(struct stcp_conn *conn, int err, void *arg); + void *close_arg; +}; + +void stcp_conn_free(struct stcp_conn *c); +void stcp_conn_set_tx_queue(struct stcp_conn *c, struct ll_queue *q); +void stcp_conn_set_rx_queue(struct stcp_conn *c, struct ll_queue *q); +void stcp_conn_set_on_close(struct stcp_conn *c, void (*cb)(struct stcp_conn *conn, int err, void *arg), void *arg); + +#endif diff --git a/src/stcp_client.c b/src/stcp_client.c new file mode 100644 index 00000000..90bff59a --- /dev/null +++ b/src/stcp_client.c @@ -0,0 +1,282 @@ +// 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 +#include +#include +#ifndef _WIN32 +#include +#include +#include +#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, "encrypt_and_crc: 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, "decrypt_and_check: 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, "decrypt_and_check: 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, "client_try_send: 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_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) { + if (!ua || !addr || !keys || !peer_pubkey || !ready_cb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_client_connect: invalid args"); return NULL; } + struct stcp_client *cli = u_calloc(1, sizeof(struct stcp_client)); + if (!cli) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_client_connect: 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->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; +} diff --git a/src/stcp_client.h b/src/stcp_client.h new file mode 100644 index 00000000..a47abb48 --- /dev/null +++ b/src/stcp_client.h @@ -0,0 +1,15 @@ +// stcp_client.h — STCP client: connect, handshake +#ifndef STCP_CLIENT_H +#define STCP_CLIENT_H + +#include "stcp.h" + +typedef void (*stcp_ready_cb)(struct stcp_conn *conn, void *arg); + +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); +void stcp_client_destroy(struct stcp_client *cli); +struct stcp_conn *stcp_client_get_conn(struct stcp_client *cli); + +#endif diff --git a/src/stcp_server.c b/src/stcp_server.c new file mode 100644 index 00000000..348bf81e --- /dev/null +++ b/src/stcp_server.c @@ -0,0 +1,375 @@ +// stcp_server.c — STCP server implementation +#include "stcp_server.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/memory_pool.h" +#include "../lib/mem.h" +#include "../lib/debug_config.h" +#include "../lib/platform_compat.h" +#include +#include +#include +#ifndef _WIN32 +#include +#include +#include +#endif + +struct stcp_server { + struct UASYNC *ua; + socket_t listen_sock; + void *listen_id; + struct SC_MYKEYS my_keys; + stcp_connect_cb connect_cb; + void *cb_arg; +}; + +static void server_accept_cb(socket_t sock, void *arg); +static void server_conn_read_cb(socket_t sock, void *arg); +static void server_conn_write_cb(socket_t sock, void *arg); +static void tx_queue_cb(struct ll_queue *q, void *arg); +static void stcp_conn_process_recv(struct stcp_conn *c); +static int stcp_conn_try_send(struct stcp_conn *c, const uint8_t *data, size_t len); +static void stcp_conn_do_close(struct stcp_conn *c, int err); + +static int stcp_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] = (data_len >> 0) & 0xFF; + output[1] = (data_len >> 8) & 0xFF; + 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); + size_t enc_len = 2 + data_len + 4; + if (sc_stream_xor(&c->stream_send, output + 2, data_len + 4) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_encrypt_and_crc: stream_xor failed len=%zu", data_len); + return -1; + } + *output_len = enc_len; + return 0; +} + +static int stcp_decrypt_and_check(uint8_t *plaintext, size_t plaintext_len, + struct sc_stream_state *stream, size_t *out_len) { + if (sc_stream_xor(stream, plaintext, plaintext_len) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_decrypt_and_check: stream_xor failed len=%zu", plaintext_len); + return -1; + } + size_t data_len = plaintext_len - 4; + uint32_t recv_crc = ((uint32_t)plaintext[data_len]) | ((uint32_t)plaintext[data_len + 1] << 8) | + ((uint32_t)plaintext[data_len + 2] << 16) | ((uint32_t)plaintext[data_len + 3] << 24); + uint32_t calc_crc = crc32_calc(plaintext, data_len); + if (recv_crc != calc_crc) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_decrypt_and_check: CRC mismatch recv=%08x calc=%08x", recv_crc, calc_crc); + return -1; + } + *out_len = data_len; + return 0; +} + +static void stcp_conn_do_close(struct stcp_conn *c, int err) { + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + int prev = c->state; + 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; } + DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_conn closed is_server=%d prev_state=%d err=%d", c->is_server, prev, err); + if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; cb(c, err, c->close_arg); } + else if (prev == STCP_STATE_HS_SERVER_WAIT) { sc_stream_cleanup(&c->stream_send); sc_stream_cleanup(&c->stream_recv); if (c->allocated) u_free(c); } +} + +static void stcp_conn_send_message(struct stcp_conn *c, const uint8_t *data, size_t len) { + if (c->state != STCP_STATE_DATA) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: not in DATA state"); return; } + size_t max_enc = STCP_MAX_MSG_SIZE; + if (len > max_enc) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: data too long %zu", len); return; } + size_t need = 2 + len + 4; + uint8_t *buf = u_malloc(need); + if (!buf) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: malloc(%zu) failed", need); return; } + if (stcp_encrypt_and_crc(c, data, len, buf, &need)) { u_free(buf); return; } + stcp_conn_try_send(c, buf, need); +} + +static int stcp_conn_try_send(struct stcp_conn *c, const uint8_t *data, size_t len) { + if (c->sock == SOCKET_INVALID) return -1; + if (c->send_buf) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_try_send: send_buf already busy"); + return -1; + } + ssize_t sent = send(c->sock, data, len, 0); + if (sent < 0) { + int err = socket_get_error(); + if (err == ERR_AGAIN || err == 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, "stcp_conn_try_send: send failed err=%d", err); + u_free((uint8_t *)data); + stcp_conn_do_close(c, err); + 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 server_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 err = socket_get_error(); + if (err == ERR_AGAIN || err == ERR_WOULDBLOCK) return; + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_write_cb: send failed err=%d", err); + stcp_conn_do_close(c, err); + 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); + return; + } +} + +static void 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 (e->dgram) stcp_conn_send_message(c, e->dgram, e->len); + queue_entry_free(e); + queue_resume_callback(q); +} + +static int derive_session_and_streams(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, "derive_session: 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_SERVER_SEND) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "derive_session: stream_send init failed"); + return -1; + } + if (sc_stream_init(&sc, &c->stream_recv, STCP_STREAM_CLIENT_SEND) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "derive_session: stream_recv init failed"); + return -1; + } + memcpy(c->peer_pubkey, peer_pubkey, SC_PUBKEY_SIZE); + c->peer_pubkey_set = 1; + return 0; +} + +static void process_client_handshake(struct stcp_conn *c) { + const uint8_t *salt = c->recv_buf; + const uint8_t *enc_pubkey = salt + SC_PUBKEY_ENC_SALT_SIZE; + uint8_t client_pubkey[SC_PUBKEY_SIZE]; + sc_obfuscate_pubkey(salt, c->my_keys.public_key, enc_pubkey, client_pubkey); + if (derive_session_and_streams(c, client_pubkey)) { stcp_conn_do_close(c, 1); return; } + uint8_t enc_hs[STCP_HS_ENC_CLIENT]; + memcpy(enc_hs, c->recv_buf + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_CLIENT); + size_t hs_data_len; + if (stcp_decrypt_and_check(enc_hs, STCP_HS_ENC_CLIENT, &c->stream_recv, &hs_data_len)) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "process_client_handshake: decrypt/CRC failed"); + stcp_conn_do_close(c, 2); return; + } + uint16_t padding_size = (uint16_t)enc_hs[0] | ((uint16_t)enc_hs[1] << 8); + c->hs_expected_len = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT + padding_size; + c->hs_key_processed = 1; +} + +static void finish_client_handshake(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; + + uint8_t salt2[SC_PUBKEY_ENC_SALT_SIZE]; + if (random_bytes(salt2, SC_PUBKEY_ENC_SALT_SIZE) != 0) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: random_bytes failed"); + stcp_conn_do_close(c, 3); return; + } + uint16_t padding = 8; + size_t total_resp = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + padding; + uint8_t *resp = u_malloc(total_resp); + if (!resp) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: malloc failed"); stcp_conn_do_close(c, 3); return; } + memcpy(resp, salt2, SC_PUBKEY_ENC_SALT_SIZE); + sc_obfuscate_pubkey(salt2, c->peer_pubkey, c->my_keys.public_key, resp + SC_PUBKEY_ENC_SALT_SIZE); + + uint8_t plain_hs[3] = {0, (uint8_t)padding, (uint8_t)(padding >> 8)}; // status=OK + padding_size + uint32_t crc = crc32_calc(plain_hs, 3); + uint8_t *enc_dst = resp + SC_PUBKEY_ENC_SIZE; + enc_dst[0] = plain_hs[0]; enc_dst[1] = plain_hs[1]; enc_dst[2] = plain_hs[2]; + enc_dst[3] = (uint8_t)(crc >> 0); enc_dst[4] = (uint8_t)(crc >> 8); + enc_dst[5] = (uint8_t)(crc >> 16); enc_dst[6] = (uint8_t)(crc >> 24); + if (sc_stream_xor(&c->stream_send, enc_dst, STCP_HS_ENC_SERVER) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: encrypt failed"); + u_free(resp); stcp_conn_do_close(c, 3); return; + } + for (int i = 0; i < padding; i++) resp[SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + i] = (uint8_t)(salt2[0] ^ i); + + c->state = STCP_STATE_DATA; + DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: handshake OK, entering DATA state"); + if (c->on_ready) c->on_ready(c, c->ready_arg); + stcp_conn_try_send(c, resp, total_resp); +} + +static void server_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_cap == 0) { + c->recv_buf_cap = STCP_RECV_BUF_INIT; + c->recv_buf = u_malloc(c->recv_buf_cap); + if (!c->recv_buf) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: malloc failed"); stcp_conn_do_close(c, ENOMEM); return; } + } + if (c->recv_buf_len + 4096 > c->recv_buf_cap) { + size_t new_cap = c->recv_buf_cap * 2; + if (new_cap > STCP_RECV_BUF_MAX) new_cap = STCP_RECV_BUF_MAX; + if (new_cap <= c->recv_buf_cap) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: recv buf full"); stcp_conn_do_close(c, ENOBUFS); return; } + uint8_t *nb = u_realloc(c->recv_buf, new_cap); + if (!nb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: realloc failed"); stcp_conn_do_close(c, ENOMEM); return; } + c->recv_buf = nb; c->recv_buf_cap = new_cap; + } + 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 err = socket_get_error(); + if (err == ERR_AGAIN || err == ERR_WOULDBLOCK) return; + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: recv failed err=%d", err); + stcp_conn_do_close(c, err); return; + } + if (n == 0) { DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: EOF"); stcp_conn_do_close(c, 0); return; } + c->recv_buf_len += (size_t)n; + stcp_conn_process_recv(c); +} + +static void stcp_conn_process_recv(struct stcp_conn *c) { + while (c->recv_buf_len > 0) { + switch (c->state) { + case STCP_STATE_HS_SERVER_WAIT: + if (!c->hs_key_processed) { + if (c->recv_buf_len < SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT) return; + process_client_handshake(c); + if (c->state != STCP_STATE_HS_SERVER_WAIT) return; + } + if (c->hs_key_processed && c->recv_buf_len >= c->hs_expected_len) { + finish_client_handshake(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, "msg too large: %u", msg_size); stcp_conn_do_close(c, 4); return; } + size_t total = 2 + msg_size + 4; + if (c->recv_buf_len < total) return; + uint8_t *enc_data = c->recv_buf + 2; + size_t data_len; + if (stcp_decrypt_and_check(enc_data, msg_size + 4, &c->stream_recv, &data_len)) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "data decrypt/CRC failed"); + stcp_conn_do_close(c, 5); 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); DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "rx malloc(%zu) failed", data_len); } + } else { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "queue_entry_new failed"); } + } + memmove(c->recv_buf, c->recv_buf + total, c->recv_buf_len - total); + c->recv_buf_len -= total; + break; + } + default: return; + } + } +} + +static void server_accept_cb(socket_t listen_sock, void *arg) { + struct stcp_server *srv = (struct stcp_server *)arg; + struct sockaddr_storage cli_addr; + socklen_t addr_len = sizeof(cli_addr); + socket_t cli_sock = accept(listen_sock, (struct sockaddr *)&cli_addr, &addr_len); + if (cli_sock == SOCKET_INVALID) { + int err = socket_get_error(); + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server accept failed err=%d", err); + return; + } + socket_set_nonblocking(cli_sock); + int opt = 1; + setsockopt(cli_sock, IPPROTO_TCP, TCP_NODELAY, (const char *)&opt, sizeof(opt)); + + struct stcp_conn *c = u_calloc(1, sizeof(struct stcp_conn)); + if (!c) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server: calloc conn failed"); socket_close_wrapper(cli_sock); return; } + c->sock = cli_sock; + c->ua = srv->ua; + c->state = STCP_STATE_HS_SERVER_WAIT; + c->is_server = 1; + c->allocated = 1; + c->tx_cb = tx_queue_cb; + c->my_keys = srv->my_keys; + c->on_ready = srv->connect_cb; + c->ready_arg = srv->cb_arg; + c->socket_id = uasync_add_socket_t(srv->ua, cli_sock, server_conn_read_cb, server_conn_write_cb, NULL, c); + DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: accepted connection fd=%d", (int)cli_sock); +} + +struct stcp_server *stcp_server_create(struct UASYNC *ua, uint16_t port, + struct SC_MYKEYS *keys, + stcp_connect_cb connect_cb, void *arg) { + if (!ua || !keys || !connect_cb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: invalid args"); return NULL; } + struct stcp_server *srv = u_calloc(1, sizeof(struct stcp_server)); + if (!srv) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: calloc failed"); return NULL; } + srv->ua = ua; + srv->my_keys = *keys; + srv->connect_cb = connect_cb; + srv->cb_arg = arg; + + srv->listen_sock = socket(AF_INET, SOCK_STREAM, 0); + if (srv->listen_sock == SOCKET_INVALID) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: socket failed"); u_free(srv); return NULL; } + socket_set_nonblocking(srv->listen_sock); + int reuse = 1; + setsockopt(srv->listen_sock, SOL_SOCKET, SO_REUSEADDR, (const char *)&reuse, sizeof(reuse)); + struct sockaddr_in addr; + memset(&addr, 0, sizeof(addr)); + addr.sin_family = AF_INET; + addr.sin_addr.s_addr = INADDR_ANY; + addr.sin_port = htons(port); + if (bind(srv->listen_sock, (struct sockaddr *)&addr, sizeof(addr)) < 0) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: bind port=%u failed err=%d", port, socket_get_error()); + socket_close_wrapper(srv->listen_sock); u_free(srv); return NULL; + } + if (listen(srv->listen_sock, 16) < 0) { + DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: listen failed err=%d", socket_get_error()); + socket_close_wrapper(srv->listen_sock); u_free(srv); return NULL; + } + srv->listen_id = uasync_add_socket_t(ua, srv->listen_sock, server_accept_cb, NULL, NULL, srv); + DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: listening on port %u", port); + return srv; +} + +void stcp_server_destroy(struct stcp_server *srv) { + if (!srv) return; + if (srv->listen_id) uasync_remove_socket_t(srv->ua, srv->listen_sock); + if (srv->listen_sock != SOCKET_INVALID) socket_close_wrapper(srv->listen_sock); + u_free(srv); + DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server destroyed"); +} \ No newline at end of file diff --git a/src/stcp_server.h b/src/stcp_server.h new file mode 100644 index 00000000..b510f0db --- /dev/null +++ b/src/stcp_server.h @@ -0,0 +1,14 @@ +// stcp_server.h — STCP server: listen, accept, handshake +#ifndef STCP_SERVER_H +#define STCP_SERVER_H + +#include "stcp.h" + +typedef void (*stcp_connect_cb)(struct stcp_conn *conn, void *arg); + +struct stcp_server *stcp_server_create(struct UASYNC *ua, uint16_t port, + struct SC_MYKEYS *keys, + stcp_connect_cb connect_cb, void *arg); +void stcp_server_destroy(struct stcp_server *srv); + +#endif diff --git a/tests/Makefile.am b/tests/Makefile.am index 5e72c999..032f524f 100644 --- a/tests/Makefile.am +++ b/tests/Makefile.am @@ -4,6 +4,8 @@ check_PROGRAMS = \ test_ll_queue \ test_serialize \ + test_stream_cipher \ + test_stcp \ test_packet_dump \ test_debug_categories \ test_config_debug \ @@ -176,6 +178,14 @@ test_etcp_crypto_SOURCES = test_etcp_crypto.c test_etcp_crypto_CFLAGS = -I$(top_srcdir)/src -I$(top_srcdir)/lib -I$(top_srcdir)/tinycrypt/lib/include -I$(top_srcdir)/tinycrypt/lib/source test_etcp_crypto_LDADD = $(SECURE_CHANNEL_OBJS) $(CRYPTO_LIBS) $(COMMON_LIBS) +test_stream_cipher_SOURCES = test_stream_cipher.c +test_stream_cipher_CFLAGS = -I$(top_srcdir)/src -I$(top_srcdir)/lib -I$(top_srcdir)/tinycrypt/lib/include -I$(top_srcdir)/tinycrypt/lib/source +test_stream_cipher_LDADD = $(SECURE_CHANNEL_OBJS) $(CRYPTO_LIBS) $(COMMON_LIBS) + +test_stcp_SOURCES = test_stcp.c +test_stcp_CFLAGS = -I$(top_srcdir)/src -I$(top_srcdir)/lib -I$(top_srcdir)/tinycrypt/lib/include -I$(top_srcdir)/tinycrypt/lib/source +test_stcp_LDADD = $(top_builddir)/src/utun-stcp.o $(top_builddir)/src/utun-stcp_server.o $(top_builddir)/src/utun-stcp_client.o $(SECURE_CHANNEL_OBJS) $(CRYPTO_LIBS) $(COMMON_LIBS) + if USE_OPENSSL else test_crypto_SOURCES = test_crypto.c diff --git a/tests/test_stcp.c b/tests/test_stcp.c new file mode 100644 index 00000000..cd41a84a --- /dev/null +++ b/tests/test_stcp.c @@ -0,0 +1,359 @@ +// test_stcp.c — comprehensive STCP integration tests +#include "../src/stcp.h" +#include "../src/stcp_server.h" +#include "../src/stcp_client.h" +#include "../src/secure_channel.h" +#include "../lib/u_async.h" +#include "../lib/ll_queue.h" +#include "../lib/debug_config.h" +#include "../lib/mem.h" +#include +#include +#include + +static int tests_passed = 0, tests_total = 0; +static struct SC_MYKEYS s_keys, c_keys; + +#define BASE_PORT 23456 + +#define TASSERT(cond) do { \ + if (!(cond)) { DEBUG_ERROR(DEBUG_CATEGORY_GENERAL, " FAIL: %s", #cond); return 1; } \ +} while(0) + +#define TRUN(name) do { \ + tests_total++; \ + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "--- %s ---", name); \ + int _r = name(); \ + if (_r == 0) { tests_passed++; DEBUG_INFO(DEBUG_CATEGORY_GENERAL, " PASS"); } \ +} while(0) + +// ======================= peer helper ======================= + +struct test_peer { + struct stcp_conn *conn; + struct ll_queue *rx, *tx; + int ready, closed, close_err, msg_count; + uint8_t *accum; + size_t accum_len, accum_cap; +}; + +static void peer_rx_cb(struct ll_queue *q, void *arg) { + struct test_peer *p = (struct test_peer *)arg; + struct ll_entry *e = queue_data_get(q); + if (!e) { queue_resume_callback(q); return; } + p->msg_count++; + size_t need = p->accum_len + e->len; + if (need > p->accum_cap) { p->accum_cap = need + 4096; p->accum = u_realloc(p->accum, p->accum_cap); } + memcpy(p->accum + p->accum_len, e->dgram, e->len); + p->accum_len += e->len; + queue_entry_free(e); + queue_resume_callback(q); +} + +static void peer_close_cb(struct stcp_conn *conn, int err, void *arg) { + (void)conn; + struct test_peer *p = (struct test_peer *)arg; + p->closed = 1; p->close_err = err; +} + +static void setup_peer(struct test_peer *p, struct stcp_conn *conn) { + p->conn = conn; + p->rx = queue_new(conn->ua, 0, 0, 0, "rx"); + p->tx = queue_new(conn->ua, 0, 0, 0, "tx"); + queue_set_callback(p->rx, peer_rx_cb, p); + stcp_conn_set_rx_queue(conn, p->rx); + stcp_conn_set_tx_queue(conn, p->tx); + stcp_conn_set_on_close(conn, peer_close_cb, p); + p->ready = 1; +} + +static void server_connect_cb(struct stcp_conn *conn, void *arg) { + struct test_peer *p = (struct test_peer *)arg; + setup_peer(p, conn); +} + +static void client_ready_cb(struct stcp_conn *conn, void *arg) { + struct test_peer *p = (struct test_peer *)arg; + setup_peer(p, conn); +} + +static void peer_cleanup(struct test_peer *p) { + if (p->rx) { queue_free(p->rx); p->rx = NULL; } + if (p->tx) { queue_free(p->tx); p->tx = NULL; } + if (p->accum) { u_free(p->accum); p->accum = NULL; } +} + +static int peer_send(struct test_peer *p, const uint8_t *data, size_t len) { + struct ll_entry *e = queue_entry_new(0); + if (!e) return -1; + e->dgram = u_malloc(len ? len : 1); + if (!e->dgram) { queue_entry_free(e); return -1; } + if (len) memcpy(e->dgram, data, len); + e->len = (uint16_t)len; + queue_data_put(p->tx, e); + return 0; +} + +// ======================= test 1: handshake + all message sizes ======================= + +static int test1_sizes(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 1; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, server_connect_cb, &srv); TASSERT(ss); + struct stcp_client *sc = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, s_keys.public_key, client_ready_cb, &cli); TASSERT(sc); + + size_t sizes[] = {0, 1, 16, 17, 255, 256, 1000, 65535}; + int n_sizes = 8; + size_t total = 0; for (int i = 0; i < n_sizes; i++) total += sizes[i]; + uint8_t *payload = u_malloc(65536); + for (int i = 0; i < 65536; i++) payload[i] = (uint8_t)(i * 7 + 13); + + int sent = 0, ticks = 0; + while (srv.msg_count < n_sizes && ticks < 5000) { + uasync_poll(ua, 10); + if (srv.ready && cli.ready && !sent) { + for (int i = 0; i < n_sizes; i++) TASSERT(peer_send(&cli, payload, sizes[i]) == 0); + sent = 1; + } + ticks++; + } + TASSERT(srv.msg_count == n_sizes); + TASSERT(srv.accum_len == total); + size_t off = 0; + for (int i = 0; i < n_sizes; i++) { + TASSERT(memcmp(srv.accum + off, payload, sizes[i]) == 0); + off += sizes[i]; + } + u_free(payload); + peer_cleanup(&srv); peer_cleanup(&cli); + if (srv.conn) stcp_conn_free(srv.conn); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= test 2: many sequential messages ======================= + +static int test2_many(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 2; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, server_connect_cb, &srv); TASSERT(ss); + struct stcp_client *sc = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, s_keys.public_key, client_ready_cb, &cli); TASSERT(sc); + + int sent = 0, ticks = 0; + while (srv.msg_count < 200 && ticks < 5000) { + uasync_poll(ua, 10); + if (srv.ready && cli.ready && !sent) { + for (int i = 0; i < 200; i++) { + uint8_t buf[8]; + buf[0] = (uint8_t)(i >> 0); buf[1] = (uint8_t)(i >> 8); + buf[2] = (uint8_t)(i >> 16); buf[3] = (uint8_t)(i >> 24); + buf[4] = (uint8_t)(i * 3); + TASSERT(peer_send(&cli, buf, 5) == 0); + } + sent = 1; + } + ticks++; + } + TASSERT(srv.msg_count == 200); + TASSERT(srv.accum_len == 200 * 5); + for (int i = 0; i < 200; i++) { + uint32_t v; memcpy(&v, srv.accum + i * 5, 4); TASSERT(v == (uint32_t)i); + TASSERT(srv.accum[i * 5 + 4] == (uint8_t)(i * 3)); + } + peer_cleanup(&srv); peer_cleanup(&cli); + if (srv.conn) stcp_conn_free(srv.conn); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= test 3: wrong peer pubkey ======================= + +static int test3_wrong_key(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 3; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, server_connect_cb, &srv); TASSERT(ss); + + struct SC_MYKEYS rogue; + TASSERT(sc_generate_keypair(&rogue) == SC_OK); + struct stcp_client *sc = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, rogue.public_key, client_ready_cb, &cli); TASSERT(sc); + + int ticks = 0; + while (ticks < 3000) { + uasync_poll(ua, 10); + if (cli.closed || srv.closed) break; + ticks++; + } + TASSERT(!cli.ready); + TASSERT(!srv.ready); + + peer_cleanup(&srv); peer_cleanup(&cli); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= test 4: close detection ======================= + +static int test4_close(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 4; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, server_connect_cb, &srv); TASSERT(ss); + struct stcp_client *sc = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, s_keys.public_key, client_ready_cb, &cli); TASSERT(sc); + + int closed = 0, ticks = 0; + while (!srv.closed && ticks < 3000) { + uasync_poll(ua, 10); + if (srv.ready && cli.ready && !closed) { + uint8_t m = 0xAB; peer_send(&cli, &m, 1); + closed = 1; + } + if (closed && srv.msg_count >= 1 && cli.conn) { + stcp_conn_free(cli.conn); cli.conn = NULL; + cli.closed = 1; + } + ticks++; + } + TASSERT(srv.closed); + + peer_cleanup(&srv); peer_cleanup(&cli); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= test 5: multiple concurrent clients ======================= + +static struct test_peer *g_multi_peers; +static int g_multi_idx, g_multi_max; + +static void multi_connect_cb(struct stcp_conn *conn, void *arg) { + (void)arg; + int i = g_multi_idx++; + if (i >= g_multi_max) return; + struct test_peer *p = &g_multi_peers[i]; + setup_peer(p, conn); +} + +static int test5_multi(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + uint16_t port = BASE_PORT + 5; + + #define NCLI 3 + struct test_peer srvp[NCLI]; + struct test_peer clip[NCLI]; + memset(srvp, 0, sizeof(srvp)); memset(clip, 0, sizeof(clip)); + + g_multi_peers = srvp; g_multi_idx = 0; g_multi_max = NCLI; + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, multi_connect_cb, NULL); TASSERT(ss); + + struct stcp_client *clients[NCLI] = {0}; + for (int i = 0; i < NCLI; i++) { + clients[i] = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, s_keys.public_key, client_ready_cb, &clip[i]); + TASSERT(clients[i]); + } + + int sent = 0, ticks = 0; + while (ticks < 5000) { + uasync_poll(ua, 10); + if (!sent) { + int all_ready = 1; + for (int i = 0; i < NCLI; i++) if (!clip[i].ready || !srvp[i].ready) all_ready = 0; + if (all_ready) { + for (int i = 0; i < NCLI; i++) { + uint8_t buf[4]; buf[0] = (uint8_t)i; buf[1] = (uint8_t)(i * 17 + 42); + TASSERT(peer_send(&clip[i], buf, 2) == 0); + } + sent = 1; + } + } + if (sent) { + int all_got = 1; + for (int i = 0; i < NCLI; i++) if (srvp[i].msg_count < 1) all_got = 0; + if (all_got) break; + } + ticks++; + } + for (int i = 0; i < NCLI; i++) { + TASSERT(srvp[i].msg_count >= 1); + TASSERT(srvp[i].accum_len == 2); + TASSERT(srvp[i].accum[0] == (uint8_t)i); + TASSERT(srvp[i].accum[1] == (uint8_t)(i * 17 + 42)); + } + for (int i = 0; i < NCLI; i++) { + peer_cleanup(&srvp[i]); peer_cleanup(&clip[i]); + if (srvp[i].conn) stcp_conn_free(srvp[i].conn); + stcp_client_destroy(clients[i]); + } + stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= test 6: interleaved send/recv ======================= + +static int test6_interleaved(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 6; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, server_connect_cb, &srv); TASSERT(ss); + struct stcp_client *sc = stcp_client_connect(ua, "127.0.0.1", port, &c_keys, s_keys.public_key, client_ready_cb, &cli); TASSERT(sc); + + int round = 0, ticks = 0; + while (srv.msg_count < 50 || cli.msg_count < 50) { + uasync_poll(ua, 10); + if (srv.ready && cli.ready && round < 50) { + uint8_t cb = (uint8_t)(round + 100); + uint8_t sb = (uint8_t)(round + 200); + if (peer_send(&cli, &cb, 1) == 0 && peer_send(&srv, &sb, 1) == 0) round++; + } + if (++ticks > 5000) break; + } + TASSERT(srv.msg_count >= 50); + TASSERT(cli.msg_count >= 50); + for (int i = 0; i < 50; i++) { + TASSERT(srv.accum[i] == (uint8_t)(i + 100)); + TASSERT(cli.accum[i] == (uint8_t)(i + 200)); + } + + peer_cleanup(&srv); peer_cleanup(&cli); + if (srv.conn) stcp_conn_free(srv.conn); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); + return 0; +} + +// ======================= main ======================= + +int main(void) { + debug_config_init(); + debug_set_level(DEBUG_LEVEL_INFO); + debug_set_categories(DEBUG_CATEGORY_GENERAL | DEBUG_CATEGORY_SOCKET | DEBUG_CATEGORY_CRYPTO); + + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "============================================"); + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "=== STCP Integration Tests ==="); + + TASSERT(sc_generate_keypair(&s_keys) == SC_OK); + TASSERT(sc_generate_keypair(&c_keys) == SC_OK); + + TRUN(test1_sizes); + TRUN(test2_many); + TRUN(test3_wrong_key); + TRUN(test4_close); + TRUN(test5_multi); + TRUN(test6_interleaved); + + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "============================================"); + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "Results: %d/%d passed", tests_passed, tests_total); + return tests_passed == tests_total ? 0 : 1; +} diff --git a/tests/test_stream_cipher.c b/tests/test_stream_cipher.c new file mode 100644 index 00000000..b215f5a4 --- /dev/null +++ b/tests/test_stream_cipher.c @@ -0,0 +1,141 @@ +// test_stream_cipher.c — tests for streaming cipher (AES-128-CTR) +#include "../src/secure_channel.h" +#include "../lib/debug_config.h" +#include +#include +#include + +static int test_failed = 0; + +#define CHECK(expr, msg) do { \ + if (!(expr)) { \ + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "FAIL: %s", msg); \ + test_failed = 1; \ + } else { \ + DEBUG_INFO(DEBUG_CATEGORY_CRYPTO, "PASS: %s", msg); \ + } \ +} while(0) + +int main(void) { + debug_config_init(); + debug_set_level(DEBUG_LEVEL_INFO); + debug_set_categories(DEBUG_CATEGORY_CRYPTO); + + DEBUG_INFO(DEBUG_CATEGORY_CRYPTO, "=== Stream Cipher Test ==="); + + struct SC_MYKEYS keys_a, keys_b; + CHECK(sc_generate_keypair(&keys_a) == SC_OK, "generate keypair A"); + CHECK(sc_generate_keypair(&keys_b) == SC_OK, "generate keypair B"); + + sc_context_t ctx_a, ctx_b; + CHECK(sc_init_ctx(&ctx_a, &keys_a) == SC_OK, "init ctx A"); + CHECK(sc_init_ctx(&ctx_b, &keys_b) == SC_OK, "init ctx B"); + + CHECK(sc_set_peer_public_key(&ctx_a, keys_b.public_key, SC_PEER_PUBKEY_BIN) == SC_OK, "set peer key A<-B"); + CHECK(sc_set_peer_public_key(&ctx_b, keys_a.public_key, SC_PEER_PUBKEY_BIN) == SC_OK, "set peer key B<-A"); + + struct sc_stream_state send_a, recv_a, send_b, recv_b; + CHECK(sc_stream_init(&ctx_a, &send_a, 0) == SC_OK, "init send_A (id=0)"); + CHECK(sc_stream_init(&ctx_a, &recv_a, 1) == SC_OK, "init recv_A (id=1)"); + CHECK(sc_stream_init(&ctx_b, &recv_b, 0) == SC_OK, "init recv_B (id=0)"); + CHECK(sc_stream_init(&ctx_b, &send_b, 1) == SC_OK, "init send_B (id=1)"); + + uint8_t original[2048], data[2048]; + for (int i = 0; i < (int)sizeof(original); i++) original[i] = (uint8_t)(i * 3 + 17); + + // Test 1: A->B, single chunk, 1024 bytes + memcpy(data, original, 1024); + CHECK(sc_stream_xor(&send_a, data, 1024) == SC_OK, "encrypt 1024B A->B"); + CHECK(sc_stream_xor(&recv_b, data, 1024) == SC_OK, "decrypt 1024B A->B"); + CHECK(memcmp(data, original, 1024) == 0, "round-trip A->B matches"); + + // Test 2: B->A, single chunk + memcpy(data, original, 1024); + CHECK(sc_stream_xor(&send_b, data, 1024) == SC_OK, "encrypt 1024B B->A"); + CHECK(sc_stream_xor(&recv_a, data, 1024) == SC_OK, "decrypt 1024B B->A"); + CHECK(memcmp(data, original, 1024) == 0, "round-trip B->A matches"); + + // Test 3: incremental feeding — encrypt in 4 chunks, decrypt in 2 + memcpy(data, original, 1024); + CHECK(sc_stream_xor(&send_a, data, 7) == SC_OK, "inc_A enc 7B"); + CHECK(sc_stream_xor(&send_a, data + 7, 33) == SC_OK, "inc_A enc 33B"); + CHECK(sc_stream_xor(&send_a, data + 40, 1) == SC_OK, "inc_A enc 1B"); + CHECK(sc_stream_xor(&send_a, data + 41, 1024 - 41) == SC_OK, "inc_A enc 983B"); + CHECK(sc_stream_xor(&recv_b, data, 513) == SC_OK, "inc_B dec 513B"); + CHECK(sc_stream_xor(&recv_b, data + 513, 1024 - 513) == SC_OK, "inc_B dec 511B"); + CHECK(memcmp(data, original, 1024) == 0, "inc round-trip matches"); + + // Test 4: all sizes 1..255 — catches alignment/keystream issues + for (size_t sz = 1; sz <= 255; sz++) { + uint8_t buf[256], cpy[256]; + for (size_t i = 0; i < sz; i++) buf[i] = (uint8_t)(sz ^ i); + memcpy(cpy, buf, sz); + CHECK(sc_stream_xor(&send_a, buf, sz) == SC_OK, "various-size enc"); + CHECK(sc_stream_xor(&recv_b, buf, sz) == SC_OK, "various-size dec"); + CHECK(memcmp(buf, cpy, sz) == 0, "various-size match"); + } + + // Test 5: 2048 bytes (= 128 AES blocks, whole number) + memcpy(data, original, 2048); + CHECK(sc_stream_xor(&send_a, data, 2048) == SC_OK, "enc 2048B"); + CHECK(sc_stream_xor(&recv_b, data, 2048) == SC_OK, "dec 2048B"); + CHECK(memcmp(data, original, 2048) == 0, "round-trip 2048 matches"); + + // Test 6: A and B streams must produce DIFFERENT keystreams + { + uint8_t test_data[16]; + memset(test_data, 0xAA, 16); + uint8_t a_out[16], b_out[16]; + memcpy(a_out, test_data, 16); + memcpy(b_out, test_data, 16); + CHECK(sc_stream_xor(&send_a, a_out, 16) == SC_OK, "xor on send_A"); + CHECK(sc_stream_xor(&send_b, b_out, 16) == SC_OK, "xor on send_B"); + CHECK(memcmp(a_out, b_out, 16) != 0, "stream_id=0 != stream_id=1 keystream"); + // Decrypt back + CHECK(sc_stream_xor(&recv_b, a_out, 16) == SC_OK, "recv_B decrypt send_A data"); + CHECK(sc_stream_xor(&recv_a, b_out, 16) == SC_OK, "recv_A decrypt send_B data"); + CHECK(memcmp(a_out, test_data, 16) == 0, "recv_B restores original"); + CHECK(memcmp(b_out, test_data, 16) == 0, "recv_A restores original"); + } + + // Test 7: encrypt 1 byte — extreme alignment edge + for (int j = 0; j < 50; j++) { + uint8_t one_byte = (uint8_t)j; + CHECK(sc_stream_xor(&send_a, &one_byte, 1) == SC_OK, "enc 1B"); + CHECK(sc_stream_xor(&recv_b, &one_byte, 1) == SC_OK, "dec 1B"); + CHECK(one_byte == (uint8_t)j, "1-byte round-trip"); + } + + // Test 8: error handling + CHECK(sc_stream_xor(&send_a, NULL, 0) == SC_OK, "xor len=0 OK"); + CHECK(sc_stream_xor(&send_a, data, 0) == SC_OK, "xor data_len=0 OK"); + CHECK(sc_stream_init(&ctx_a, NULL, 0) == SC_ERR_INVALID_ARG, "init NULL state"); + CHECK(sc_stream_init(NULL, &send_a, 0) == SC_ERR_INVALID_ARG, "init NULL ctx"); + + struct sc_stream_state uninit; + memset(&uninit, 0, sizeof(uninit)); + CHECK(sc_stream_xor(&uninit, data, 1) == SC_ERR_NOT_INITIALIZED, "xor uninitialized"); + + // У ctx без session_ready + sc_context_t no_session_ctx; + struct SC_MYKEYS dummy_keys; + memset(&dummy_keys, 0, sizeof(dummy_keys)); + sc_init_ctx(&no_session_ctx, &dummy_keys); + struct sc_stream_state bad_stream; + CHECK(sc_stream_init(&no_session_ctx, &bad_stream, 0) == SC_ERR_NOT_INITIALIZED, "init without session"); + + // Cleanup + sc_stream_cleanup(&send_a); + sc_stream_cleanup(&recv_a); + sc_stream_cleanup(&send_b); + sc_stream_cleanup(&recv_b); + sc_stream_cleanup(NULL); + sc_stream_cleanup(&uninit); + + if (test_failed) { + DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "=== Stream Cipher Test: FAILED ==="); + return 1; + } + DEBUG_INFO(DEBUG_CATEGORY_CRYPTO, "=== Stream Cipher Test: PASSED ==="); + return 0; +}