diff --git a/src/transport_layer/stcp.c b/src/transport_layer/stcp.c index df9cfca8..52a77da8 100644 --- a/src/transport_layer/stcp.c +++ b/src/transport_layer/stcp.c @@ -1,10 +1,12 @@ -// stcp.c — shared stcp_conn lifecycle +// stcp.c — shared stcp_conn lifecycle + frame encrypt/decrypt + pending queue + send #include "stcp.h" #include "../lib/ll_queue.h" #include "../lib/mem.h" #include "../lib/debug_config.h" +#include "../lib/platform_compat.h" #include #include +#include void stcp_conn_set_tx_queue(struct stcp_conn *c, struct ll_queue *q) { if (!c || !q) return; @@ -31,7 +33,141 @@ void stcp_conn_free(struct stcp_conn *c) { 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; } + stcp_pending_clear(c); sc_stream_cleanup(&c->stream_send); sc_stream_cleanup(&c->stream_recv); if (c->allocated) u_free(c); } + +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); + output[1] = (uint8_t)(data_len >> 8); + if (data_len) 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_ETCP, "stcp_frame_encrypt stream_xor failed len=%zu", data_len); + return -1; + } + *output_len = enc_len; + return 0; +} + +int stcp_frame_decrypt(uint8_t *data, size_t len, struct sc_stream_state *stream, size_t *out_len) { + if (len < SC_CRC32_SIZE) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_frame_decrypt too short len=%zu", len); return -1; } + if (sc_stream_xor(stream, data, len) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_frame_decrypt stream_xor failed len=%zu", len); + 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; + } + *out_len = data_len; + return 0; +} + +void stcp_pending_queue(struct stcp_conn *c, const uint8_t *data, size_t len) { + if (!c) return; + struct pending_entry *pe = u_malloc(sizeof(*pe)); + if (!pe) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_pending_queue malloc entry failed"); return; } + if (len && data) { pe->data = u_malloc(len); if (!pe->data) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_pending_queue malloc data(%zu) failed", len); u_free(pe); return; } memcpy(pe->data, data, len); } + else { pe->data = NULL; } + pe->len = len; + pe->next = NULL; + if (c->pending_tail) c->pending_tail->next = pe; + else c->pending_head = pe; + c->pending_tail = pe; +} + +void stcp_pending_clear(struct stcp_conn *c) { + if (!c) return; + while (c->pending_head) { + struct pending_entry *pe = c->pending_head; + c->pending_head = pe->next; + if (pe->data) u_free(pe->data); + u_free(pe); + } + c->pending_tail = NULL; +} + +int stcp_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_ETCP, "stcp_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 1; + } + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_try_send send failed err=%d", 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 1; + } + u_free((uint8_t *)data); + return 0; +} + +void stcp_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_ETCP, "stcp_write_cb send failed err=%d", err); + if (c->on_write_error) c->on_write_error(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); + stcp_flush_pending(c); + } +} + +void stcp_flush_pending(struct stcp_conn *c) { + if (!c) return; + while (!c->send_buf && c->pending_head) { + struct pending_entry *pe = c->pending_head; + c->pending_head = pe->next; + if (!c->pending_head) c->pending_tail = NULL; + + size_t enc_len; + uint8_t *enc = u_malloc(2 + pe->len + SC_CRC32_SIZE); + if (!enc) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_flush_pending malloc enc failed len=%zu", pe->len); if (pe->data) u_free(pe->data); u_free(pe); continue; } + if (stcp_frame_encrypt(c, pe->data, pe->len, enc, &enc_len)) { + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_flush_pending encrypt failed"); u_free(enc); if (pe->data) u_free(pe->data); u_free(pe); + if (c->on_write_error) c->on_write_error(c, ECANCELED); + return; + } + if (pe->data) u_free(pe->data); + u_free(pe); + + int r = stcp_try_send(c, enc, enc_len); + if (r < 0) { + u_free(enc); + if (c->on_write_error) c->on_write_error(c, ECANCELED); + return; + } + if (r > 0) return; + } +} diff --git a/src/transport_layer/stcp.h b/src/transport_layer/stcp.h index 56263c44..e606521c 100644 --- a/src/transport_layer/stcp.h +++ b/src/transport_layer/stcp.h @@ -38,6 +38,12 @@ enum stcp_state { STCP_STATE_ERROR }; +struct pending_entry { + struct pending_entry *next; + uint8_t *data; + size_t len; +}; + struct stcp_conn { socket_t sock; struct UASYNC *ua; @@ -74,6 +80,10 @@ struct stcp_conn { size_t hs_expected_len; uint8_t hs_key_processed; + struct pending_entry *pending_head; + struct pending_entry *pending_tail; + void (*on_write_error)(struct stcp_conn *c, int err); + void (*on_ready)(struct stcp_conn *conn, void *arg); void *ready_arg; @@ -86,6 +96,14 @@ 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); +int stcp_frame_encrypt(struct stcp_conn *c, const uint8_t *data, size_t data_len, uint8_t *output, size_t *output_len); +int stcp_frame_decrypt(uint8_t *data, size_t len, struct sc_stream_state *stream, size_t *out_len); +void stcp_pending_queue(struct stcp_conn *c, const uint8_t *data, size_t len); +void stcp_pending_clear(struct stcp_conn *c); +int stcp_try_send(struct stcp_conn *c, const uint8_t *data, size_t len); +void stcp_write_cb(socket_t sock, void *arg); +void stcp_flush_pending(struct stcp_conn *c); + #ifdef __cplusplus } diff --git a/src/transport_layer/stcp_client.c b/src/transport_layer/stcp_client.c index d312ab13..bab942eb 100644 --- a/src/transport_layer/stcp_client.c +++ b/src/transport_layer/stcp_client.c @@ -28,104 +28,16 @@ struct stcp_client { 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_ETCP, "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_ETCP, "stream_xor failed"); return -1; } - log_dump(DEBUG_LEVEL_ERROR, DEBUG_CATEGORY_DEBUG, "decrypt_and_check AFTER xor", data, len); - 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, "CRC mismatch recv=%08x calc=%08x", recv_crc, calc_crc); 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; - int prev = c->state; - c->state = STCP_STATE_CLOSED; - DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_client: closed is_server=%d err=%d prev_state=%d", c->is_server, err, prev); - DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close sock=%d sock_id=%p recv_buf=%p len=%zu cap=%zu send_buf=%p on_close=%p", - (int)c->sock, (void*)c->socket_id, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap, (void*)c->send_buf, (void*)c->on_close); - 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) { void *p = c->recv_buf; u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close u_free recv_buf=%p", p); } - if (c->send_buf) { void *p = c->send_buf; u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close u_free send_buf=%p", p); } - if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close on_close returned"); } -} - -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_ETCP, "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_ETCP, "client ECDH failed"); return -1; } memcpy(c->session_key, sc.session_key, SC_SESSION_KEY_SIZE); + memcpy(c->peer_pubkey, peer_pubkey, SC_PUBKEY_SIZE); c->peer_pubkey_set = 1; log_dump(DEBUG_LEVEL_ERROR, DEBUG_CATEGORY_DEBUG, "stcp_client session_key", c->session_key, 16); if (sc_stream_init(&sc, &c->stream_send, STCP_STREAM_CLIENT_SEND) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "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_ETCP, "client stream_recv init failed"); return -1; } @@ -152,7 +64,8 @@ static void client_send_handshake(struct stcp_conn *c, const uint8_t *server_pub c->state = STCP_STATE_HS_CLIENT_SENT; DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_client: handshake sent (%zu bytes), entering HS_CLIENT_SENT", total); - client_try_send(c, hs, total); + c->send_buf = hs; c->send_len = total; c->send_offset = 0; + uasync_set_socket_write(c->ua, c->socket_id, 1); } static void process_server_response(struct stcp_conn *c) { @@ -162,13 +75,17 @@ static void process_server_response(struct stcp_conn *c) { 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); + if (c->peer_pubkey_set && memcmp(server_pubkey, c->peer_pubkey, SC_PUBKEY_SIZE) != 0) { + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_client: server pubkey mismatch — possible MITM"); + client_do_close(c, 4); return; + } 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); log_dump(DEBUG_LEVEL_ERROR, DEBUG_CATEGORY_DEBUG, "stcp_client enc_hs BEFORE xor", enc_hs, 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; } + if (stcp_frame_decrypt(enc_hs, STCP_HS_ENC_SERVER, &c->stream_recv, &hs_data_len)) { client_do_close(c, 3); return; } log_dump(DEBUG_LEVEL_ERROR, DEBUG_CATEGORY_DEBUG, "stcp_client enc_hs AFTER xor", enc_hs, STCP_HS_ENC_SERVER); memcpy(c->peer_ed25519_pubkey, enc_hs, SC_PUBKEY_SIZE); c->peer_ed25519_set = 1; uint8_t status = enc_hs[32]; @@ -213,7 +130,7 @@ static void stcp_conn_process_recv(struct stcp_conn *c) { 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 (stcp_frame_decrypt(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); } } @@ -267,13 +184,40 @@ static void client_connect_write_cb(socket_t sock, void *arg) { } DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_client: TCP connected, starting handshake"); 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); + c->socket_id = uasync_add_socket_t(cli->ua, sock, client_conn_read_cb, stcp_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, cli->my_ed25519_pubkey); } +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) + stcp_pending_queue(c, e->dgram, e->len); + queue_dgram_free(e); + queue_entry_free(e); + queue_resume_callback(q); + stcp_flush_pending(c); +} + +static void client_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; + DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_client: closed is_server=%d err=%d prev_state=%d", c->is_server, err, prev); + DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close sock=%d sock_id=%p recv_buf=%p len=%zu cap=%zu send_buf=%p on_close=%p", + (int)c->sock, (void*)c->socket_id, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap, (void*)c->send_buf, (void*)c->on_close); + 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) { void *p = c->recv_buf; u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close u_free recv_buf=%p", p); } + if (c->send_buf) { void *p = c->send_buf; u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close u_free send_buf=%p", p); } + stcp_pending_clear(c); + if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_client: do_close on_close returned"); } +} + struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uint16_t port, struct SC_MYKEYS *keys, const uint8_t *peer_pubkey, const uint8_t *my_ed25519_pubkey, @@ -289,6 +233,7 @@ struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uin 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->on_write_error = client_do_close; c->tx_cb = client_tx_queue_cb; struct addrinfo hints = {0}; @@ -314,9 +259,8 @@ struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uin 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; } } else { - /* immediate connect (localhost) — запускаем handshake сразу */ int opt = 1; setsockopt(c->sock, IPPROTO_TCP, TCP_NODELAY, (const char *)&opt, sizeof(opt)); - c->socket_id = uasync_add_socket_t(ua, c->sock, client_conn_read_cb, client_conn_write_cb, NULL, c); + 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; } if (client_derive_session(c, cli->peer_pubkey)) { client_do_close(c, 1); return cli; } client_send_handshake(c, cli->peer_pubkey, cli->my_ed25519_pubkey); diff --git a/src/transport_layer/stcp_server.c b/src/transport_layer/stcp_server.c index 7635f4f2..a797b5ca 100644 --- a/src/transport_layer/stcp_server.c +++ b/src/transport_layer/stcp_server.c @@ -32,140 +32,10 @@ struct stcp_server { 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_ETCP, "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_ETCP, "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_ETCP, "CRC mismatch recv=%08x calc=%08x", recv_crc, calc_crc); - return -1; - } - *out_len = data_len; - return 0; -} - -static void stcp_conn_deferred_free(void* arg) { - u_free((struct stcp_conn*)arg); -} - -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) { void *p = c->recv_buf; u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close u_free recv_buf=%p", p); } - if (c->send_buf) { void *p = c->send_buf; u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close u_free send_buf=%p", p); } - DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_conn closed is_server=%d prev_state=%d err=%d sock=%d", c->is_server, prev, err, (int)c->sock); - DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close sock=%d sock_id=%p recv_buf was=%p on_close=%p", - (int)c->sock, (void*)c->socket_id, (void*)c->recv_buf, (void*)c->on_close); - if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close on_close returned"); } - else if (prev == STCP_STATE_HS_SERVER_WAIT) { sc_stream_cleanup(&c->stream_send); sc_stream_cleanup(&c->stream_recv); if (c->allocated) uasync_call_soon(c->ua, c, stcp_conn_deferred_free); } -} - -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_ETCP, "not in DATA state"); return; } - size_t max_enc = STCP_MAX_MSG_SIZE; - if (len > max_enc) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "data too long %zu", len); return; } - size_t need = 2 + len + 4; - uint8_t *buf = u_malloc(need); - if (!buf) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "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_ETCP, "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_ETCP, "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_ETCP, "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_dgram_free(e); - 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); @@ -197,7 +67,7 @@ static void process_client_handshake(struct stcp_conn *c) { 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)) { + if (stcp_frame_decrypt(enc_hs, STCP_HS_ENC_CLIENT, &c->stream_recv, &hs_data_len)) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "decrypt/CRC failed"); stcp_conn_do_close(c, 2); return; } @@ -242,7 +112,8 @@ static void finish_client_handshake(struct stcp_conn *c) { c->state = STCP_STATE_DATA; DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server: handshake OK, entering DATA state"); log_dump(DEBUG_LEVEL_ERROR, DEBUG_CATEGORY_DEBUG, "stcp_server FULL RESP", resp, total_resp); - stcp_conn_try_send(c, resp, total_resp); // send handshake response BEFORE on_ready — otherwise ETCP sends INIT data which mixes into client's recv buffer + int r = stcp_try_send(c, resp, total_resp); + if (r < 0) { u_free(resp); stcp_conn_do_close(c, 3); return; } if (c->on_ready) c->on_ready(c, c->ready_arg); } @@ -308,7 +179,7 @@ static void stcp_conn_process_recv(struct stcp_conn *c) { 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)) { + if (stcp_frame_decrypt(enc_data, msg_size + 4, &c->stream_recv, &data_len)) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "data decrypt/CRC failed"); stcp_conn_do_close(c, 5); return; } @@ -335,6 +206,38 @@ static void stcp_conn_process_recv(struct stcp_conn *c) { } } +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 (c->state == STCP_STATE_DATA && e->dgram) + stcp_pending_queue(c, e->dgram, e->len); + queue_dgram_free(e); + queue_entry_free(e); + queue_resume_callback(q); + stcp_flush_pending(c); +} + +static void stcp_conn_deferred_free(void* arg) { + stcp_conn_free((struct stcp_conn*)arg); +} + +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) { void *p = c->recv_buf; u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close u_free recv_buf=%p", p); } + if (c->send_buf) { void *p = c->send_buf; u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close u_free send_buf=%p", p); } + stcp_pending_clear(c); + DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_conn closed is_server=%d prev_state=%d err=%d sock=%d", c->is_server, prev, err, (int)c->sock); + DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close sock=%d sock_id=%p recv_buf was=%p on_close=%p", + (int)c->sock, (void*)c->socket_id, (void*)c->recv_buf, (void*)c->on_close); + if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; c->on_close = NULL; DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_DEBUG, "stcp_server: do_close on_close returned"); } + else if (prev == STCP_STATE_HS_SERVER_WAIT) { sc_stream_cleanup(&c->stream_send); sc_stream_cleanup(&c->stream_recv); if (c->allocated) uasync_call_soon(c->ua, c, stcp_conn_deferred_free); } +} + static void server_accept_cb(socket_t listen_sock, void *arg) { struct stcp_server *srv = (struct stcp_server *)arg; struct sockaddr_storage cli_addr; @@ -356,6 +259,7 @@ static void server_accept_cb(socket_t listen_sock, void *arg) { c->state = STCP_STATE_HS_SERVER_WAIT; c->is_server = 1; c->allocated = 1; + c->on_write_error = stcp_conn_do_close; c->tx_cb = tx_queue_cb; c->my_keys = srv->my_keys; memcpy(c->my_ed25519_pubkey, srv->my_ed25519_pubkey, SC_PUBKEY_SIZE); @@ -363,7 +267,7 @@ static void server_accept_cb(socket_t listen_sock, void *arg) { c->ready_arg = srv->cb_arg; c->on_close = srv->close_cb; c->close_arg = srv->close_arg; - c->socket_id = uasync_add_socket_t(srv->ua, cli_sock, server_conn_read_cb, server_conn_write_cb, NULL, c); + c->socket_id = uasync_add_socket_t(srv->ua, cli_sock, server_conn_read_cb, stcp_write_cb, NULL, c); DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server: accepted connection fd=%d", (int)cli_sock); } @@ -430,4 +334,4 @@ void stcp_server_destroy(struct stcp_server *srv) { if (srv->listen_sock != SOCKET_INVALID) socket_close_wrapper(srv->listen_sock); u_free(srv); DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server destroyed"); -} \ No newline at end of file +} diff --git a/tests/test_stcp.c b/tests/test_stcp.c index 22834c64..723416b3 100644 --- a/tests/test_stcp.c +++ b/tests/test_stcp.c @@ -332,6 +332,45 @@ static int test6_interleaved(void) { return 0; } +// ======================= test 7: 4MB bulk transfer (pending queue stress) ======================= + +static int test7_bulk_4mb(void) { + struct UASYNC *ua = uasync_create(); TASSERT(ua); + struct test_peer srv = {0}, cli = {0}; + uint16_t port = BASE_PORT + 7; + + 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); + + #define N_BULK 64 + #define SZ_BULK 65535 + size_t total = N_BULK * SZ_BULK; // 4,194,240 bytes + uint8_t *payload = u_malloc(total); TASSERT(payload); + for (size_t i = 0; i < total; i++) payload[i] = (uint8_t)(i * 7 + 13); + + int sent = 0, ticks = 0; + while (srv.msg_count < N_BULK && ticks < 10000) { + uasync_poll(ua, 10); + if (srv.ready && cli.ready && !sent) { + for (int i = 0; i < N_BULK; i++) + TASSERT(peer_send(&cli, payload + i * SZ_BULK, SZ_BULK) == 0); + sent = 1; + } + ticks++; + } + DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "bulk: server received %d msgs, %zu bytes in %d ticks", srv.msg_count, srv.accum_len, ticks); + TASSERT(srv.msg_count == N_BULK); + TASSERT(srv.accum_len == total); + TASSERT(memcmp(srv.accum, payload, total) == 0); + + 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; +} + // ======================= main ======================= int main(void) { @@ -351,6 +390,7 @@ int main(void) { TRUN(test4_close); TRUN(test5_multi); TRUN(test6_interleaved); + TRUN(test7_bulk_4mb); DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "============================================"); DEBUG_INFO(DEBUG_CATEGORY_GENERAL, "Results: %d/%d passed", tests_passed, tests_total);