diff --git a/src/transport_layer/stcp.c b/src/transport_layer/stcp.c index f907e886..9077642f 100644 --- a/src/transport_layer/stcp.c +++ b/src/transport_layer/stcp.c @@ -39,6 +39,18 @@ void stcp_conn_free(struct stcp_conn *c) { if (c->allocated) u_free(c); } +static int stcp_check_crc(const uint8_t *data, size_t data_len) { + 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_ETCP, "stcp CRC mismatch recv=%08x calc=%08x data_len=%zu", + recv_crc, calc_crc, data_len); + return -1; + } + return 0; +} + int stcp_frame_encrypt(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); @@ -64,13 +76,7 @@ int stcp_frame_decrypt(uint8_t *data, size_t len, struct sc_stream_state *stream 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_ETCP, "stcp_frame_decrypt CRC mismatch recv=%08x calc=%08x", recv_crc, calc_crc); - return -1; - } + if (stcp_check_crc(data, data_len) != 0) return -1; *out_len = data_len; return 0; } @@ -249,13 +255,7 @@ void stcp_recv_try(struct stcp_conn *c) { size_t data_len = (size_t)c->recv_msg_size; uint8_t *plain = c->recv_buf + 2; - uint32_t recv_crc = ((uint32_t)plain[data_len]) | ((uint32_t)plain[data_len + 1] << 8) | - ((uint32_t)plain[data_len + 2] << 16) | ((uint32_t)plain[data_len + 3] << 24); - uint32_t calc_crc = crc32_calc(plain, data_len); - if (recv_crc != calc_crc) { - DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_recv_try CRC mismatch recv=%08x calc=%08x data_len=%zu", recv_crc, calc_crc, data_len); - stcp_conn_do_close(c, 5); return; - } + if (stcp_check_crc(plain, data_len) != 0) { stcp_conn_do_close(c, 5); return; } if (c->recv_on_chunk) c->recv_on_chunk(c, plain, data_len); if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; @@ -301,7 +301,8 @@ void stcp_conn_do_close(struct stcp_conn *c, int err) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; cb(c, err, c->close_arg); - } else if (c->allocated) { + } + if (c->allocated) { uasync_call_soon(c->ua, c, stcp_conn_deferred_free); } } diff --git a/src/transport_layer/stcp_client.c b/src/transport_layer/stcp_client.c index c39eca78..c89e7cc9 100644 --- a/src/transport_layer/stcp_client.c +++ b/src/transport_layer/stcp_client.c @@ -183,7 +183,7 @@ struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uin c->socket_id = uasync_add_socket_t(ua, c->sock, client_conn_read_cb, stcp_write_cb, NULL, c); if (!c->socket_id) { socket_close_wrapper(c->sock); u_free(cli); return NULL; } c->hs_timer = uasync_set_timeout(ua, STCP_CONNECT_TIMEOUT, c, hs_timeout_cb, "stcp_hs"); - if (client_derive_session(c, cli->peer_pubkey)) { stcp_conn_do_close(c, 1); return cli; } + if (client_derive_session(c, cli->peer_pubkey)) { stcp_conn_do_close(c, 1); u_free(cli); return NULL; } client_send_handshake(c, cli->peer_pubkey, cli->my_ed25519_pubkey); stcp_recv_set(c, SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER, 0, client_hs_cb); } diff --git a/src/transport_layer/stcp_link.c b/src/transport_layer/stcp_link.c index 4a3be4a1..2b82f979 100644 --- a/src/transport_layer/stcp_link.c +++ b/src/transport_layer/stcp_link.c @@ -84,7 +84,7 @@ static void link_rx_cb(struct ll_queue *q, void *arg) { static void stcp_link_on_stcp_close(struct stcp_conn *conn, int err, void *arg) { struct stcp_link *link = (struct stcp_link *)arg; (void)conn; (void)err; - link->conn = NULL; link->ready = 0; + link->ready = 0; if (link->etcp_link) { link->etcp_link->recv_keepalive = 0; link->etcp_link->link_status = 0; } if (link->on_close_cb) link->on_close_cb(link, err, link->close_arg); } @@ -210,13 +210,18 @@ struct stcp_link *stcp_link_connect(struct stcp_link_config *cfg) { static void stcp_link_close_impl(void *arg) { struct stcp_link *link = (struct stcp_link *)arg; + if (link->cli) { + struct stcp_conn *c = stcp_client_get_conn(link->cli); + if (c && c->rx_queue) queue_free(c->rx_queue); + if (link->tx_queue) queue_free(link->tx_queue); stcp_client_destroy(link->cli); } else if (link->conn) { - if (link->conn->rx_queue) queue_free(link->conn->rx_queue); - if (link->tx_queue) queue_free(link->tx_queue); - stcp_conn_do_close(link->conn, 0); - stcp_conn_free(link->conn); + struct stcp_conn *conn = link->conn; + if (conn->rx_queue) { queue_free(conn->rx_queue); conn->rx_queue = NULL; } + if (link->tx_queue) { queue_free(link->tx_queue); link->tx_queue = NULL; } + stcp_conn_do_close(conn, 0); + link->conn = NULL; } u_free(link); } diff --git a/tests/test_stcp.c b/tests/test_stcp.c index 1b0fa159..d06024e9 100644 --- a/tests/test_stcp.c +++ b/tests/test_stcp.c @@ -7,6 +7,7 @@ #include "../lib/ll_queue.h" #include "../lib/debug_config.h" #include "../lib/mem.h" +#include "../lib/socket_compat.h" #include #include #include @@ -373,6 +374,36 @@ static int test7_bulk_4mb(void) { return 0; } +// ======================= test 8: server recv error → close chain (allocated=1 + on_close + deferred free) ======================= + +static int test8_srv_recv_close(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 8; + + struct stcp_server *ss = stcp_server_create(ua, port, &s_keys, NULL, 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, NULL, client_ready_cb, &cli, peer_close_cb, &cli); TASSERT(sc); + + int ticks = 0; + while ((!srv.ready || !cli.ready) && ticks < 200) { uasync_poll(ua, 10); ticks++; } + TASSERT(srv.ready && cli.ready); + + socket_close_wrapper(cli.conn->sock); cli.conn->sock = SOCKET_INVALID; + + ticks = 0; + while (!srv.closed && ticks < 200) { uasync_poll(ua, 10); ticks++; } + TASSERT(srv.closed); + + // server conn freed via deferred free (allocated=1 + on_close=peer_close_cb) + srv.conn = NULL; // prevent explicit stcp_conn_free double-free + srv.ready = 0; + + peer_cleanup(&srv); peer_cleanup(&cli); + stcp_client_destroy(sc); stcp_server_destroy(ss); + uasync_destroy(ua, 1); // flush deferred → stcp_conn_free(srv_conn) + return 0; +} + // ======================= main ======================= int main(void) { @@ -393,6 +424,7 @@ int main(void) { TRUN(test5_multi); TRUN(test6_interleaved); TRUN(test7_bulk_4mb); + TRUN(test8_srv_recv_close); DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "============================================"); DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "Results: %d/%d passed", tests_passed, tests_total);