diff --git a/src/transport_layer/auto_socket.c b/src/transport_layer/auto_socket.c index 429e2157..a31c0ef5 100644 --- a/src/transport_layer/auto_socket.c +++ b/src/transport_layer/auto_socket.c @@ -28,13 +28,6 @@ #include "node_conn_direct.h" #include "socket_monitor.h" #include "stcp_link.h" - -/* struct stcp_server layout from stcp_link.c — opaque in public header, - need this to safely remove from inst->stcp_servers linked list */ -struct stcp_link_server_local { - struct stcp_link_server_local *next; - struct stcp_server *srv; -}; #include "utun_instance.h" #include "config_parser.h" #include "../chat/chat_event.h" @@ -326,7 +319,7 @@ static struct stcp_server* create_iface_tcp_socket(struct AUTO_SOCKET* as, uint3 (ia->ss_family == AF_INET6 && IN6_IS_ADDR_UNSPECIFIED(&((struct sockaddr_in6*)ia)->sin6_addr))) { DEBUG_INFO(DEBUG_CATEGORY_AS, "[as] TCP socket %s has null addr — removing", server.name); if (ts) tcp_socket_remove(ts); - { struct stcp_link_server_local** pp = (struct stcp_link_server_local**)&inst->stcp_servers; while (*pp) { if ((struct stcp_server*)*pp == tsrv) { *pp = (*pp)->next; break; } pp = &(*pp)->next; } } + stcp_server_list_remove(inst, tsrv); stcp_link_server_destroy(tsrv); continue; } } @@ -429,6 +422,7 @@ static void remove_iface_sockets(struct AUTO_SOCKET* as, uint32_t ifindex) { tsp = &(*tsp)->next; } } + stcp_server_list_remove(as->instance, ifa->v4_tcp); stcp_link_server_destroy(ifa->v4_tcp); delete_port_from_db(as, ifname, AF_INET, AS_PROTO_TCP); } @@ -444,6 +438,7 @@ static void remove_iface_sockets(struct AUTO_SOCKET* as, uint32_t ifindex) { tsp = &(*tsp)->next; } } + stcp_server_list_remove(as->instance, ifa->v6_tcp); stcp_link_server_destroy(ifa->v6_tcp); delete_port_from_db(as, ifname, AF_INET6, AS_PROTO_TCP); } @@ -721,6 +716,7 @@ static void reconcile_iface(struct AUTO_SOCKET* as, uint32_t ifindex, const char if (a->ss_family == AF_INET) tp = ntohs(((struct sockaddr_in*)a)->sin_port); if (tp == p && p > 0) { struct TCP_SOCKET* rm = *tsp; *tsp = rm->next; u_free(rm); break; } tsp = &(*tsp)->next; } } + stcp_server_list_remove(as->instance, srv); stcp_link_server_destroy(srv); delete_port_from_db(as, ifname, AF_INET, AS_PROTO_TCP); ifa->v4_tcp = NULL; changed = 1; @@ -743,6 +739,7 @@ static void reconcile_iface(struct AUTO_SOCKET* as, uint32_t ifindex, const char if (a->ss_family == AF_INET6) tp = ntohs(((struct sockaddr_in6*)a)->sin6_port); if (tp == p && p > 0) { struct TCP_SOCKET* rm = *tsp; *tsp = rm->next; u_free(rm); break; } tsp = &(*tsp)->next; } } + stcp_server_list_remove(as->instance, srv); stcp_link_server_destroy(srv); delete_port_from_db(as, ifname, AF_INET6, AS_PROTO_TCP); ifa->v6_tcp = NULL; changed = 1; diff --git a/src/transport_layer/stcp.c b/src/transport_layer/stcp.c index 21a5e90c..f907e886 100644 --- a/src/transport_layer/stcp.c +++ b/src/transport_layer/stcp.c @@ -1,4 +1,4 @@ -// stcp.c — shared stcp_conn lifecycle + frame encrypt/decrypt + pending queue + send +// stcp.c — shared stcp_conn lifecycle + frame encrypt/decrypt + pending queue + send + recv FSM #include "stcp.h" #include "../lib/ll_queue.h" #include "../lib/mem.h" @@ -48,12 +48,12 @@ int stcp_frame_encrypt(struct stcp_conn *c, const uint8_t *data, size_t data_len 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) { + size_t total = 2 + data_len + SC_CRC32_SIZE; + if (sc_stream_xor(&c->stream_send, output, total) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_frame_encrypt stream_xor failed len=%zu", data_len); return -1; } - *output_len = enc_len; + *output_len = total; return 0; } @@ -99,17 +99,14 @@ void stcp_pending_clear(struct stcp_conn *c) { c->pending_tail = NULL; } -int stcp_try_send(struct stcp_conn *c, const uint8_t *data, size_t len) { +int stcp_try_send(struct stcp_conn *c, 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; - } + 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; + c->send_buf = data; c->send_len = len; c->send_offset = 0; uasync_set_socket_write(c->ua, c->socket_id, 1); return 1; } @@ -117,11 +114,11 @@ int stcp_try_send(struct stcp_conn *c, const uint8_t *data, size_t len) { return -1; } if ((size_t)sent < len) { - c->send_buf = (uint8_t *)data; c->send_len = len; c->send_offset = (size_t)sent; + c->send_buf = 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); + u_free(data); return 0; } @@ -163,11 +160,192 @@ void stcp_flush_pending(struct stcp_conn *c) { 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) { u_free(enc); if (c->on_write_error) c->on_write_error(c, ECANCELED); return; } if (r > 0) return; } } + +// ====== unified recv ====== + +int stcp_conn_read(struct stcp_conn *c) { + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return -1; + 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) { stcp_conn_do_close(c, ENOMEM); return -1; } + } + 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) { stcp_conn_do_close(c, ENOBUFS); return -1; } + uint8_t *nb = u_realloc(c->recv_buf, new_cap); + if (!nb) { stcp_conn_do_close(c, ENOMEM); return -1; } + c->recv_buf = nb; c->recv_buf_cap = new_cap; + } + ssize_t n = recv(c->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 0; + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp recv failed err=%d", err); + stcp_conn_do_close(c, err); return -1; + } + if (n == 0) { DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp recv EOF"); stcp_conn_do_close(c, 0); return -1; } + c->recv_buf_len += (size_t)n; + return 1; +} + +void stcp_recv_set(struct stcp_conn *c, size_t need, int streaming, + void (*on_chunk)(struct stcp_conn *c, uint8_t *data, size_t len)) { + if (!c) return; + c->recv_need = need; + c->recv_streaming = streaming; + c->recv_on_chunk = on_chunk; + if (streaming) { c->recv_in_meta = 1; c->recv_need = 2; } +} + +void stcp_recv_try(struct stcp_conn *c) { + if (!c || c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + + while (c->recv_buf_len > 0) { + if (!c->recv_streaming) { + if (c->recv_buf_len < c->recv_need) return; + void (*cb)(struct stcp_conn*, uint8_t*, size_t) = c->recv_on_chunk; + if (!cb) return; + size_t consumed = c->recv_need; + cb(c, c->recv_buf, consumed); + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + memmove(c->recv_buf, c->recv_buf + consumed, c->recv_buf_len - consumed); + c->recv_buf_len -= consumed; + if (c->recv_streaming) continue; + if (c->recv_need == 0) return; + continue; + } + + // stream mode + if (c->recv_in_meta) { + if (c->recv_buf_len < 2) return; + if (sc_stream_xor(&c->stream_recv, c->recv_buf, 2) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_recv_try stream_xor header failed"); + stcp_conn_do_close(c, ECANCELED); 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_ETCP, "stcp_recv_try msg too large: %u", msg_size); + stcp_conn_do_close(c, 4); return; + } + c->recv_msg_size = msg_size; + c->recv_need = 2 + msg_size + SC_CRC32_SIZE; + c->recv_in_meta = 0; + } + + if (c->recv_buf_len < c->recv_need) return; + + if (c->recv_need > 2) { + if (sc_stream_xor(&c->stream_recv, c->recv_buf + 2, c->recv_need - 2) != SC_OK) { + DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_recv_try stream_xor payload failed"); + stcp_conn_do_close(c, ECANCELED); return; + } + } + + 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 (c->recv_on_chunk) c->recv_on_chunk(c, plain, data_len); + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + + memmove(c->recv_buf, c->recv_buf + c->recv_need, c->recv_buf_len - c->recv_need); + c->recv_buf_len -= c->recv_need; + c->recv_in_meta = 1; + c->recv_need = 2; + } +} + +// ====== unified close ====== + +static void stcp_conn_deferred_free(void *arg) { + stcp_conn_free((struct stcp_conn*)arg); +} + +void hs_timeout_cb(void *arg) { + struct stcp_conn *c = (struct stcp_conn *)arg; + c->hs_timer = NULL; + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + DEBUG_WARN(DEBUG_CATEGORY_ETCP, "stcp handshake timeout is_server=%d prev_state=%d", c->is_server, c->state); + stcp_conn_do_close(c, ETIMEDOUT); +} + +void stcp_conn_do_close(struct stcp_conn *c, int err) { + if (!c) return; + if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; + int prev = c->state; + c->state = STCP_STATE_CLOSED; + + if (c->hs_timer) { uasync_cancel_timeout(c->ua, c->hs_timer); c->hs_timer = NULL; } + + DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_conn close is_server=%d prev_state=%d err=%d sock=%d", c->is_server, prev, err, (int)c->sock); + + 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; } + stcp_pending_clear(c); + + if (c->on_close) { + void (*cb)(struct stcp_conn*, int, void*) = c->on_close; + c->on_close = NULL; + cb(c, err, c->close_arg); + } else if (c->allocated) { + uasync_call_soon(c->ua, c, stcp_conn_deferred_free); + } +} + +// ====== unified tx queue callback ====== + +void stcp_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); +} + +// ====== rx backpressure ====== + +void stcp_rx_push(struct stcp_conn *c, uint8_t *data, size_t len) { + if (!c || !c->rx_queue) return; + struct ll_entry *e = queue_entry_new(0); + if (!e) return; + e->dgram = u_malloc(len ? len : 1); + if (!e->dgram) { queue_entry_free(e); return; } + if (len) memcpy(e->dgram, data, len); + e->len = (uint16_t)len; + queue_data_put(c->rx_queue, e); + + if (!c->rx_paused && c->socket_id && + (c->rx_queue->count > STCP_RX_QUEUE_MAX_PACKETS || c->rx_queue->total_bytes > STCP_RX_QUEUE_MAX_BYTES)) { + c->rx_paused = 1; + uasync_set_socket_read(c->ua, c->socket_id, 0); + DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp rx paused count=%d bytes=%zu", c->rx_queue->count, c->rx_queue->total_bytes); + } +} + +void stcp_rx_resume_if_needed(struct stcp_conn *c) { + if (!c || !c->rx_paused || !c->rx_queue) return; + if (c->state != STCP_STATE_DATA) return; + if (c->rx_queue->count < STCP_RX_QUEUE_MAX_PACKETS && c->rx_queue->total_bytes < STCP_RX_QUEUE_MAX_BYTES) { + c->rx_paused = 0; + if (c->socket_id) uasync_set_socket_read(c->ua, c->socket_id, 1); + DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp rx resumed count=%d bytes=%zu", c->rx_queue->count, c->rx_queue->total_bytes); + } +} diff --git a/src/transport_layer/stcp.h b/src/transport_layer/stcp.h index e606521c..23336c0b 100644 --- a/src/transport_layer/stcp.h +++ b/src/transport_layer/stcp.h @@ -18,8 +18,8 @@ extern "C" { #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_TIMEOUT 50000 // 5s in 0.1ms timebase units +#define STCP_CONNECT_TIMEOUT 100000 // 10s in 0.1ms timebase units #define STCP_HS_CLIENT_MIN (SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT) // 72+38=110 #define STCP_HS_SERVER_MIN (SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER) // 72+39=111 @@ -29,6 +29,9 @@ extern "C" { #define STCP_STREAM_CLIENT_SEND 0 #define STCP_STREAM_SERVER_SEND 1 +#define STCP_RX_QUEUE_MAX_PACKETS 64 +#define STCP_RX_QUEUE_MAX_BYTES (256 * 1024) + enum stcp_state { STCP_STATE_INIT, STCP_STATE_HS_CLIENT_SENT, @@ -77,9 +80,6 @@ struct stcp_conn { size_t send_len; size_t send_offset; - 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); @@ -89,6 +89,16 @@ struct stcp_conn { void (*on_close)(struct stcp_conn *conn, int err, void *arg); void *close_arg; + + void *hs_timer; // handshake timeout handle, NULL when not active + + // recv state machine (stcp_recv_try) + size_t recv_need; // expected next chunk size (0 = disabled) + uint8_t recv_streaming; // 0=raw (handshake), 1=stream (DATA, auto-XOR+CRC) + uint8_t recv_in_meta; // stream: waiting for 2-byte header + uint16_t recv_msg_size; // stream: size from decoded header + void (*recv_on_chunk)(struct stcp_conn *c, uint8_t *data, size_t len); + uint8_t rx_paused; // reading paused due to rx_queue overflow }; void stcp_conn_free(struct stcp_conn *c); @@ -96,14 +106,42 @@ 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); +// tx: encrypts frame, XOR covers len+data+CRC (continuous stream) int stcp_frame_encrypt(struct stcp_conn *c, const uint8_t *data, size_t data_len, uint8_t *output, size_t *output_len); +// decrypt: XOR-in-place data+CRC blob (handshake only; DATA uses stcp_recv_try stream mode) 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); + +// takes ownership of data (must be u_malloc'd, will be u_free'd) +int stcp_try_send(struct stcp_conn *c, uint8_t *data, size_t len); void stcp_write_cb(socket_t sock, void *arg); void stcp_flush_pending(struct stcp_conn *c); +// unified recv: read from TCP, grow recv_buf, calls stcp_recv_try. returns 1=data, 0=wait, -1=error/closed +int stcp_conn_read(struct stcp_conn *c); + +// set next expected chunk: need=0+streaming=1 enters DATA stream mode +void stcp_recv_set(struct stcp_conn *c, size_t need, int streaming, + void (*on_chunk)(struct stcp_conn *c, uint8_t *data, size_t len)); +// try to process buffered data according to recv_need/streaming/on_chunk +void stcp_recv_try(struct stcp_conn *c); + +// unified close: closes socket, frees buffers, calls on_close. does NOT clean streams (stcp_conn_free does) +void stcp_conn_do_close(struct stcp_conn *c, int err); + +// handshake timeout callback (for use with uasync_set_timeout) +void hs_timeout_cb(void *arg); + +// unified tx queue callback (for use as tx_cb) +void stcp_tx_queue_cb(struct ll_queue *q, void *arg); + +// push decoded data to rx_queue; pauses reading if queue over threshold +void stcp_rx_push(struct stcp_conn *c, uint8_t *data, size_t len); +// call after consumer dequeues from rx_queue; resumes reading when below threshold +void stcp_rx_resume_if_needed(struct stcp_conn *c); + #ifdef __cplusplus } diff --git a/src/transport_layer/stcp_client.c b/src/transport_layer/stcp_client.c index 54073ed1..c39eca78 100644 --- a/src/transport_layer/stcp_client.c +++ b/src/transport_layer/stcp_client.c @@ -28,9 +28,9 @@ 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_tx_queue_cb(struct ll_queue *q, void *arg); -static void stcp_conn_process_recv(struct stcp_conn *c); -static void client_do_close(struct stcp_conn *c, int err); +static void client_hs_cb(struct stcp_conn *c, uint8_t *data, size_t len); +static void client_hs_padding_cb(struct stcp_conn *c, uint8_t *data, size_t len); +static void client_data_cb(struct stcp_conn *c, uint8_t *plain_data, size_t data_len); static int client_derive_session(struct stcp_conn *c, const uint8_t *peer_pubkey) { struct secure_channel sc; @@ -46,11 +46,11 @@ static int client_derive_session(struct stcp_conn *c, const uint8_t *peer_pubkey static void client_send_handshake(struct stcp_conn *c, const uint8_t *server_pubkey, const uint8_t *my_ed25519) { uint8_t salt[SC_PUBKEY_ENC_SALT_SIZE]; - if (random_bytes(salt, SC_PUBKEY_ENC_SALT_SIZE) != 0) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "client random_bytes failed"); client_do_close(c, 1); return; } + if (random_bytes(salt, SC_PUBKEY_ENC_SALT_SIZE) != 0) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "client random_bytes failed"); stcp_conn_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_ETCP, "client hs malloc failed"); client_do_close(c, 1); return; } + if (!hs) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "client hs malloc failed"); stcp_conn_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); @@ -59,7 +59,7 @@ static void client_send_handshake(struct stcp_conn *c, const uint8_t *server_pub uint8_t *enc_dst = hs + SC_PUBKEY_ENC_SIZE; memcpy(enc_dst, plain, 34); enc_dst[34] = (uint8_t)(crc >> 0); enc_dst[35] = (uint8_t)(crc >> 8); enc_dst[36] = (uint8_t)(crc >> 16); enc_dst[37] = (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; } + if (sc_stream_xor(&c->stream_send, enc_dst, STCP_HS_ENC_CLIENT) != SC_OK) { u_free(hs); stcp_conn_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; @@ -68,107 +68,53 @@ static void client_send_handshake(struct stcp_conn *c, const uint8_t *server_pub uasync_set_socket_write(c->ua, c->socket_id, 1); } -static void process_server_response(struct stcp_conn *c) { - log_dump(DEBUG_LEVEL_DEBUG, DEBUG_CATEGORY_ETCP, "stcp_client process_srv_resp recv_buf", c->recv_buf, c->recv_buf_len); - 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); - 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; +static void client_hs_cb(struct stcp_conn *c, uint8_t *data, size_t len) { + (void)len; + log_dump(DEBUG_LEVEL_DEBUG, DEBUG_CATEGORY_ETCP, "stcp_client process_srv_resp recv_buf", data, len); + + const uint8_t *salt = data; + 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"); + stcp_conn_do_close(c, 4); return; } + uint8_t enc_hs[STCP_HS_ENC_SERVER]; - memcpy(enc_hs, c->recv_buf + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_SERVER); + memcpy(enc_hs, data + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_SERVER); log_dump(DEBUG_LEVEL_DEBUG, DEBUG_CATEGORY_CRYPTO, "stcp_client enc_hs BEFORE xor", enc_hs, STCP_HS_ENC_SERVER); size_t hs_data_len; - if (stcp_frame_decrypt(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)) { stcp_conn_do_close(c, 3); return; } log_dump(DEBUG_LEVEL_DEBUG, DEBUG_CATEGORY_CRYPTO, "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]; uint16_t padding_size = (uint16_t)enc_hs[33] | ((uint16_t)enc_hs[34] << 8); - c->hs_expected_len = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + padding_size; DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_client: server response OK status=%d padding=%u", status, padding_size); - if (status != 0) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "server handshake status=%d", status); client_do_close(c, 4); return; } + if (status != 0) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "server handshake status=%d", status); stcp_conn_do_close(c, 4); return; } + stcp_recv_set(c, padding_size, 0, client_hs_padding_cb); } -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; +static void client_hs_padding_cb(struct stcp_conn *c, uint8_t *data, size_t len) { + (void)data; (void)len; + if (c->hs_timer) { uasync_cancel_timeout(c->ua, c->hs_timer); c->hs_timer = NULL; } c->state = STCP_STATE_DATA; DEBUG_INFO(DEBUG_CATEGORY_ETCP, "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); } + stcp_recv_set(c, 0, 1, client_data_cb); } -static void stcp_conn_process_recv(struct stcp_conn *c) { - DEBUG_TRACE(DEBUG_CATEGORY_ETCP, "stcp_client: process_recv entry state=%d sock=%d recv_buf=%p len=%zu cap=%zu", - c->state, (int)c->sock, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap); - if (!c->recv_buf && c->recv_buf_len > 0) { - DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_client: process_recv NULL recv_buf with len=%zu — closing", c->recv_buf_len); - client_do_close(c, EFAULT); return; - } - 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_ETCP, "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 (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); } } - if (c->state != STCP_STATE_DATA) return; - } - if ((uintptr_t)c->recv_buf < 4096) { - DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_client: suspicious recv_buf=%p (addr < 4K) — closing", (void*)c->recv_buf); - client_do_close(c, EFAULT); return; - } - memmove(c->recv_buf, c->recv_buf + total, c->recv_buf_len - total); - c->recv_buf_len -= total; - break; - } - default: return; - } - } +static void client_data_cb(struct stcp_conn *c, uint8_t *plain_data, size_t data_len) { + stcp_rx_push(c, plain_data, data_len); } 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) { - DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp_client: read_cb SKIP state=%d sock=%d c=%p recv_buf=%p", c->state, (int)sock, (void*)c, (void*)c->recv_buf); - return; - } - DEBUG_TRACE(DEBUG_CATEGORY_ETCP, "stcp_client: read_cb state=%d sock=%d c=%p recv_buf=%p len=%zu cap=%zu", - c->state, (int)sock, (void*)c, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap); - 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); + (void)sock; + int r = stcp_conn_read(c); + if (r == -1) return; + if (r == 0) return; + stcp_recv_try(c); } static void client_connect_write_cb(socket_t sock, void *arg) { @@ -178,42 +124,17 @@ static void client_connect_write_cb(socket_t sock, void *arg) { socklen_t len = sizeof(err); if (getsockopt(sock, SOL_SOCKET, SO_ERROR, (char *)&err, &len) < 0 || err != 0) { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_client connect failed err=%d", err); - client_do_close(c, err); return; + stcp_conn_do_close(c, err); return; } 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, stcp_write_cb, NULL, c); - if (!c->socket_id) { client_do_close(c, ENOMEM); return; } + if (!c->socket_id) { stcp_conn_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; } + if (client_derive_session(c, cli->peer_pubkey)) { stcp_conn_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_ETCP, "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_ETCP, "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_ETCP, "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_ETCP, "stcp_client: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp_client: do_close on_close returned"); } + // after handshake sent, wait for server response + stcp_recv_set(c, SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER, 0, client_hs_cb); } struct stcp_client *stcp_client_connect(struct UASYNC *ua, const char *addr, uint16_t port, @@ -231,8 +152,8 @@ 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; + c->on_write_error = stcp_conn_do_close; + c->tx_cb = stcp_tx_queue_cb; struct addrinfo hints = {0}; hints.ai_family = AF_UNSPEC; @@ -256,19 +177,22 @@ 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; } + c->hs_timer = uasync_set_timeout(ua, STCP_CONNECT_TIMEOUT, c, hs_timeout_cb, "stcp_hs"); } else { 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, 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; } + 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; } 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); } return cli; } void stcp_client_destroy(struct stcp_client *cli) { if (!cli) return; - client_do_close(&cli->conn, 0); + stcp_conn_do_close(&cli->conn, 0); stcp_conn_free(&cli->conn); u_free(cli); } diff --git a/src/transport_layer/stcp_link.c b/src/transport_layer/stcp_link.c index edc1dd05..4a3be4a1 100644 --- a/src/transport_layer/stcp_link.c +++ b/src/transport_layer/stcp_link.c @@ -76,12 +76,11 @@ static void link_rx_cb(struct ll_queue *q, void *arg) { else { queue_dgram_free(e); queue_entry_free(e); } } queue_resume_callback(q); + if (link->conn) stcp_rx_resume_if_needed(link->conn); } // ====== STCP close → link down ====== -#include "etcp_connections.h" - 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; @@ -104,6 +103,7 @@ static void server_accept_cb(struct stcp_conn *conn, void *arg) { struct ll_queue *rx = queue_new(conn->ua, 0, 0, 0, "srx"); queue_set_callback(rx, link_rx_cb, link); + queue_set_threshold(rx, STCP_RX_QUEUE_MAX_PACKETS, STCP_RX_QUEUE_MAX_BYTES); stcp_conn_set_rx_queue(conn, rx); link->tx_queue = queue_new(conn->ua, 0, 0, 0, "stx"); queue_set_threshold(link->tx_queue, 0, 0); @@ -126,6 +126,7 @@ static void client_ready_cb(struct stcp_conn *conn, void *arg) { struct ll_queue *rx = queue_new(conn->ua, 0, 0, 0, "crx"); queue_set_callback(rx, link_rx_cb, link); + queue_set_threshold(rx, STCP_RX_QUEUE_MAX_PACKETS, STCP_RX_QUEUE_MAX_BYTES); stcp_conn_set_rx_queue(conn, rx); link->tx_queue = queue_new(conn->ua, 0, 0, 0, "ctx"); queue_set_threshold(link->tx_queue, 0, 0); @@ -209,13 +210,14 @@ 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->on_close_cb) link->on_close_cb(link, 0, link->close_arg); - if (link->conn) { + if (link->cli) { + 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); } - if (link->cli) stcp_client_destroy(link->cli); u_free(link); } @@ -289,6 +291,16 @@ void stcp_server_list_add(struct UTUN_INSTANCE *inst, struct stcp_server *srv) { DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server_list_add: %p total=%d", (void*)srv, stcp_server_list_count(inst)); } +void stcp_server_list_remove(struct UTUN_INSTANCE *inst, struct stcp_server *srv) { + if (!inst || !srv) return; + struct stcp_server **pp = &inst->stcp_servers; + while (*pp) { + if (*pp == srv) { *pp = (*pp)->next; DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server_list_remove: %p total=%d", (void*)srv, stcp_server_list_count(inst)); return; } + pp = &(*pp)->next; + } + DEBUG_WARN(DEBUG_CATEGORY_ETCP, "stcp_server_list_remove: not found %p", (void*)srv); +} + void stcp_server_list_destroy_all(struct UTUN_INSTANCE *inst) { if (!inst) return; int count = 0; diff --git a/src/transport_layer/stcp_link.h b/src/transport_layer/stcp_link.h index 95e67609..64f3cb70 100644 --- a/src/transport_layer/stcp_link.h +++ b/src/transport_layer/stcp_link.h @@ -41,6 +41,7 @@ void stcp_link_server_destroy(struct stcp_server *srv); /* Server list management (multiple servers per instance) */ void stcp_server_list_add(struct UTUN_INSTANCE *inst, struct stcp_server *srv); +void stcp_server_list_remove(struct UTUN_INSTANCE *inst, struct stcp_server *srv); void stcp_server_list_destroy_all(struct UTUN_INSTANCE *inst); int stcp_server_list_count(struct UTUN_INSTANCE *inst); diff --git a/src/transport_layer/stcp_server.c b/src/transport_layer/stcp_server.c index a5a9d44f..5d9caea8 100644 --- a/src/transport_layer/stcp_server.c +++ b/src/transport_layer/stcp_server.c @@ -5,7 +5,6 @@ #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" @@ -32,9 +31,9 @@ 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 tx_queue_cb(struct ll_queue *q, void *arg); -static void stcp_conn_process_recv(struct stcp_conn *c); -static void stcp_conn_do_close(struct stcp_conn *c, int err); +static void server_hs_phase1_cb(struct stcp_conn *c, uint8_t *data, size_t len); +static void server_hs_phase2_cb(struct stcp_conn *c, uint8_t *data, size_t len); +static void server_data_cb(struct stcp_conn *c, uint8_t *plain_data, size_t data_len); static int derive_session_and_streams(struct stcp_conn *c, const uint8_t *peer_pubkey) { struct secure_channel sc; @@ -58,14 +57,15 @@ static int derive_session_and_streams(struct stcp_conn *c, const uint8_t *peer_p return 0; } -static void process_client_handshake(struct stcp_conn *c) { - const uint8_t *salt = c->recv_buf; +static void server_hs_phase1_cb(struct stcp_conn *c, uint8_t *data, size_t len) { + (void)len; + const uint8_t *salt = data; 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); + memcpy(enc_hs, data + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_CLIENT); size_t 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"); @@ -73,16 +73,14 @@ static void process_client_handshake(struct stcp_conn *c) { } memcpy(c->peer_ed25519_pubkey, enc_hs, SC_PUBKEY_SIZE); c->peer_ed25519_set = 1; uint16_t padding_size = (uint16_t)enc_hs[32] | ((uint16_t)enc_hs[33] << 8); - c->hs_expected_len = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT + padding_size; - c->hs_key_processed = 1; - DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server: client handshake key processed, expecting %zu bytes padding", c->hs_expected_len - SC_PUBKEY_ENC_SIZE - STCP_HS_ENC_CLIENT); + DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server: client handshake key processed, expecting %u bytes padding", padding_size); + stcp_recv_set(c, padding_size, 0, server_hs_phase2_cb); } -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; +static void server_hs_phase2_cb(struct stcp_conn *c, uint8_t *data, size_t len) { + (void)data; (void)len; + + if (c->hs_timer) { uasync_cancel_timeout(c->ua, c->hs_timer); c->hs_timer = NULL; } uint8_t salt2[SC_PUBKEY_ENC_SALT_SIZE]; if (random_bytes(salt2, SC_PUBKEY_ENC_SALT_SIZE) != 0) { @@ -115,125 +113,22 @@ static void finish_client_handshake(struct stcp_conn *c) { 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); + stcp_recv_set(c, 0, 1, server_data_cb); } -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) { - DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp_server: read_cb SKIP state=%d sock=%d c=%p recv_buf=%p", c->state, (int)sock, (void*)c, (void*)c->recv_buf); - return; - } - DEBUG_TRACE(DEBUG_CATEGORY_ETCP, "stcp_server: read_cb state=%d sock=%d c=%p recv_buf=%p len=%zu cap=%zu", - c->state, (int)sock, (void*)c, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap); - - 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_ETCP, "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_ETCP, "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_ETCP, "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_ETCP, "recv failed err=%d", err); - stcp_conn_do_close(c, err); return; - } - if (n == 0) { DEBUG_INFO(DEBUG_CATEGORY_ETCP, "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) { - DEBUG_TRACE(DEBUG_CATEGORY_ETCP, "stcp_server: process_recv entry state=%d sock=%d recv_buf=%p len=%zu cap=%zu", - c->state, (int)c->sock, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap); - if (!c->recv_buf && c->recv_buf_len > 0) { - DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_server: process_recv NULL recv_buf with len=%zu — closing", c->recv_buf_len); - stcp_conn_do_close(c, EFAULT); return; - } - 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_ETCP, "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_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; - } - 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_ETCP, "rx malloc(%zu) failed", data_len); } - } else { DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "queue_entry_new failed"); } - } - if ((uintptr_t)c->recv_buf < 4096) { - DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp_server: suspicious recv_buf=%p (addr < 4K) — closing", (void*)c->recv_buf); - stcp_conn_do_close(c, EFAULT); return; - } - memmove(c->recv_buf, c->recv_buf + total, c->recv_buf_len - total); - c->recv_buf_len -= total; - break; - } - default: return; - } - } +static void server_data_cb(struct stcp_conn *c, uint8_t *plain_data, size_t data_len) { + stcp_rx_push(c, plain_data, data_len); } -static void tx_queue_cb(struct ll_queue *q, void *arg) { +static void server_conn_read_cb(socket_t sock, 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_ETCP, "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_ETCP, "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_ETCP, "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_ETCP, "stcp_server: do_close calling on_close=%p", (void*)cb); cb(c, err, c->close_arg); DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "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); } + (void)sock; + DEBUG_TRACE(DEBUG_CATEGORY_ETCP, "stcp_server: read_cb state=%d sock=%d c=%p recv_buf=%p len=%zu cap=%zu", + c->state, (int)c->sock, (void*)c, (void*)c->recv_buf, c->recv_buf_len, c->recv_buf_cap); + int r = stcp_conn_read(c); + if (r == -1) return; + if (r == 0) return; + stcp_recv_try(c); } static void server_accept_cb(socket_t listen_sock, void *arg) { @@ -241,11 +136,7 @@ static void server_accept_cb(socket_t listen_sock, void *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_ETCP, "stcp_server accept failed err=%d", err); - return; - } + if (cli_sock == SOCKET_INVALID) { int err = socket_get_error(); DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "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)); @@ -258,7 +149,7 @@ static void server_accept_cb(socket_t listen_sock, void *arg) { c->is_server = 1; c->allocated = 1; c->on_write_error = stcp_conn_do_close; - c->tx_cb = tx_queue_cb; + c->tx_cb = stcp_tx_queue_cb; c->my_keys = srv->my_keys; memcpy(c->my_ed25519_pubkey, srv->my_ed25519_pubkey, SC_PUBKEY_SIZE); c->on_ready = srv->connect_cb; @@ -266,6 +157,8 @@ static void server_accept_cb(socket_t listen_sock, void *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, stcp_write_cb, NULL, c); + stcp_recv_set(c, SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT, 0, server_hs_phase1_cb); + c->hs_timer = uasync_set_timeout(c->ua, STCP_HS_TIMEOUT, c, hs_timeout_cb, "stcp_hs"); DEBUG_INFO(DEBUG_CATEGORY_ETCP, "stcp_server: accepted connection fd=%d", (int)cli_sock); } diff --git a/tests/test_stcp.c b/tests/test_stcp.c index 723416b3..1b0fa159 100644 --- a/tests/test_stcp.c +++ b/tests/test_stcp.c @@ -48,6 +48,7 @@ static void peer_rx_cb(struct ll_queue *q, void *arg) { p->accum_len += e->len; queue_entry_free(e); queue_resume_callback(q); + if (p->conn) stcp_rx_resume_if_needed(p->conn); } static void peer_close_cb(struct stcp_conn *conn, int err, void *arg) { @@ -60,6 +61,7 @@ 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_threshold(p->rx, STCP_RX_QUEUE_MAX_PACKETS, STCP_RX_QUEUE_MAX_BYTES); 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);