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.
 
 
 
 
 
 

358 lines
12 KiB

// test_stcp.c — comprehensive STCP integration tests
#include "stcp.h"
#include "stcp_server.h"
#include "stcp_client.h"
#include "secure_channel.h"
#include "../lib/u_async.h"
#include "../lib/ll_queue.h"
#include "../lib/debug_config.h"
#include "../lib/mem.h"
#include <stdio.h>
#include <string.h>
#include <stdlib.h>
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);
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, peer_close_cb, &srv, AF_INET); 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, peer_close_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 < 200) {
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, peer_close_cb, &srv, AF_INET); 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, peer_close_cb, &cli); TASSERT(sc);
int sent = 0, ticks = 0;
while (srv.msg_count < 200 && ticks < 200) {
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, peer_close_cb, &srv, AF_INET); 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, peer_close_cb, &cli); TASSERT(sc);
int ticks = 0;
while (ticks < 200) {
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, peer_close_cb, &srv, AF_INET); 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, peer_close_cb, &cli); TASSERT(sc);
int closed = 0, ticks = 0;
while (!srv.closed && ticks < 200) {
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, NULL, NULL, AF_INET); 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], peer_close_cb, &clip[i]);
TASSERT(clients[i]);
}
int sent = 0, ticks = 0;
while (ticks < 200) {
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, peer_close_cb, &srv, AF_INET); 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, peer_close_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 > 200) 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;
}