// stcp_server.c — STCP server implementation #include "stcp_server.h" #include "secure_channel.h" #include "crc32.h" #include "../lib/u_async.h" #include "../lib/socket_compat.h" #include "../lib/ll_queue.h" #include "../lib/memory_pool.h" #include "../lib/mem.h" #include "../lib/debug_config.h" #include "../lib/platform_compat.h" #include #include #include #ifndef _WIN32 #include #include #include #endif struct stcp_server { struct UASYNC *ua; socket_t listen_sock; void *listen_id; struct SC_MYKEYS my_keys; stcp_connect_cb connect_cb; void *cb_arg; }; static void server_accept_cb(socket_t sock, void *arg); static void server_conn_read_cb(socket_t sock, void *arg); static void server_conn_write_cb(socket_t sock, void *arg); static void tx_queue_cb(struct ll_queue *q, void *arg); static void stcp_conn_process_recv(struct stcp_conn *c); static int stcp_conn_try_send(struct stcp_conn *c, const uint8_t *data, size_t len); static void stcp_conn_do_close(struct stcp_conn *c, int err); static int stcp_encrypt_and_crc(struct stcp_conn *c, const uint8_t *data, size_t data_len, uint8_t *output, size_t *output_len) { uint32_t crc = crc32_calc(data, data_len); output[0] = (data_len >> 0) & 0xFF; output[1] = (data_len >> 8) & 0xFF; memcpy(output + 2, data, data_len); output[2 + data_len + 0] = (uint8_t)(crc >> 0); output[2 + data_len + 1] = (uint8_t)(crc >> 8); output[2 + data_len + 2] = (uint8_t)(crc >> 16); output[2 + data_len + 3] = (uint8_t)(crc >> 24); size_t enc_len = 2 + data_len + 4; if (sc_stream_xor(&c->stream_send, output + 2, data_len + 4) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_encrypt_and_crc: stream_xor failed len=%zu", data_len); return -1; } *output_len = enc_len; return 0; } static int stcp_decrypt_and_check(uint8_t *plaintext, size_t plaintext_len, struct sc_stream_state *stream, size_t *out_len) { if (sc_stream_xor(stream, plaintext, plaintext_len) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_decrypt_and_check: stream_xor failed len=%zu", plaintext_len); return -1; } size_t data_len = plaintext_len - 4; uint32_t recv_crc = ((uint32_t)plaintext[data_len]) | ((uint32_t)plaintext[data_len + 1] << 8) | ((uint32_t)plaintext[data_len + 2] << 16) | ((uint32_t)plaintext[data_len + 3] << 24); uint32_t calc_crc = crc32_calc(plaintext, data_len); if (recv_crc != calc_crc) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_decrypt_and_check: CRC mismatch recv=%08x calc=%08x", recv_crc, calc_crc); return -1; } *out_len = data_len; return 0; } static void stcp_conn_do_close(struct stcp_conn *c, int err) { if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; int prev = c->state; c->state = STCP_STATE_CLOSED; if (c->socket_id) { uasync_remove_socket_t(c->ua, c->sock); c->socket_id = NULL; } if (c->sock != SOCKET_INVALID) { socket_close_wrapper(c->sock); c->sock = SOCKET_INVALID; } if (c->recv_buf) { u_free(c->recv_buf); c->recv_buf = NULL; c->recv_buf_len = 0; c->recv_buf_cap = 0; } if (c->send_buf) { u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; } DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_conn closed is_server=%d prev_state=%d err=%d", c->is_server, prev, err); if (c->on_close) { void (*cb)(struct stcp_conn*, int, void*) = c->on_close; cb(c, err, c->close_arg); } else if (prev == STCP_STATE_HS_SERVER_WAIT) { sc_stream_cleanup(&c->stream_send); sc_stream_cleanup(&c->stream_recv); if (c->allocated) u_free(c); } } static void stcp_conn_send_message(struct stcp_conn *c, const uint8_t *data, size_t len) { if (c->state != STCP_STATE_DATA) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: not in DATA state"); return; } size_t max_enc = STCP_MAX_MSG_SIZE; if (len > max_enc) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: data too long %zu", len); return; } size_t need = 2 + len + 4; uint8_t *buf = u_malloc(need); if (!buf) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_send_message: malloc(%zu) failed", need); return; } if (stcp_encrypt_and_crc(c, data, len, buf, &need)) { u_free(buf); return; } stcp_conn_try_send(c, buf, need); } static int stcp_conn_try_send(struct stcp_conn *c, const uint8_t *data, size_t len) { if (c->sock == SOCKET_INVALID) return -1; if (c->send_buf) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_try_send: send_buf already busy"); return -1; } ssize_t sent = send(c->sock, data, len, 0); if (sent < 0) { int err = socket_get_error(); if (err == ERR_AGAIN || err == ERR_WOULDBLOCK) { c->send_buf = (uint8_t *)data; c->send_len = len; c->send_offset = 0; uasync_set_socket_write(c->ua, c->socket_id, 1); return 0; } DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_conn_try_send: send failed err=%d", err); u_free((uint8_t *)data); stcp_conn_do_close(c, err); return -1; } if ((size_t)sent < len) { c->send_buf = (uint8_t *)data; c->send_len = len; c->send_offset = (size_t)sent; uasync_set_socket_write(c->ua, c->socket_id, 1); return 0; } u_free((uint8_t *)data); return 0; } static void server_conn_write_cb(socket_t sock, void *arg) { struct stcp_conn *c = (struct stcp_conn *)arg; if (!c->send_buf) { uasync_set_socket_write(c->ua, c->socket_id, 0); return; } ssize_t sent = send(sock, c->send_buf + c->send_offset, c->send_len - c->send_offset, 0); if (sent < 0) { int err = socket_get_error(); if (err == ERR_AGAIN || err == ERR_WOULDBLOCK) return; DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_write_cb: send failed err=%d", err); stcp_conn_do_close(c, err); return; } c->send_offset += (size_t)sent; if (c->send_offset >= c->send_len) { u_free(c->send_buf); c->send_buf = NULL; c->send_len = 0; c->send_offset = 0; uasync_set_socket_write(c->ua, c->socket_id, 0); return; } } static void tx_queue_cb(struct ll_queue *q, void *arg) { struct stcp_conn *c = (struct stcp_conn *)arg; struct ll_entry *e = queue_data_get(q); if (!e) { queue_resume_callback(q); return; } if (e->dgram) stcp_conn_send_message(c, e->dgram, e->len); queue_entry_free(e); queue_resume_callback(q); } static int derive_session_and_streams(struct stcp_conn *c, const uint8_t *peer_pubkey) { struct secure_channel sc; sc_init_ctx(&sc, &c->my_keys); if (sc_set_peer_public_key(&sc, peer_pubkey, SC_PEER_PUBKEY_BIN) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "derive_session: ECDH failed"); return -1; } memcpy(c->session_key, sc.session_key, SC_SESSION_KEY_SIZE); if (sc_stream_init(&sc, &c->stream_send, STCP_STREAM_SERVER_SEND) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "derive_session: stream_send init failed"); return -1; } if (sc_stream_init(&sc, &c->stream_recv, STCP_STREAM_CLIENT_SEND) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_CRYPTO, "derive_session: stream_recv init failed"); return -1; } memcpy(c->peer_pubkey, peer_pubkey, SC_PUBKEY_SIZE); c->peer_pubkey_set = 1; return 0; } static void process_client_handshake(struct stcp_conn *c) { const uint8_t *salt = c->recv_buf; const uint8_t *enc_pubkey = salt + SC_PUBKEY_ENC_SALT_SIZE; uint8_t client_pubkey[SC_PUBKEY_SIZE]; sc_obfuscate_pubkey(salt, c->my_keys.public_key, enc_pubkey, client_pubkey); if (derive_session_and_streams(c, client_pubkey)) { stcp_conn_do_close(c, 1); return; } uint8_t enc_hs[STCP_HS_ENC_CLIENT]; memcpy(enc_hs, c->recv_buf + SC_PUBKEY_ENC_SIZE, STCP_HS_ENC_CLIENT); size_t hs_data_len; if (stcp_decrypt_and_check(enc_hs, STCP_HS_ENC_CLIENT, &c->stream_recv, &hs_data_len)) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "process_client_handshake: decrypt/CRC failed"); stcp_conn_do_close(c, 2); return; } uint16_t padding_size = (uint16_t)enc_hs[0] | ((uint16_t)enc_hs[1] << 8); c->hs_expected_len = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT + padding_size; c->hs_key_processed = 1; } static void finish_client_handshake(struct stcp_conn *c) { memmove(c->recv_buf, c->recv_buf + c->hs_expected_len, c->recv_buf_len - c->hs_expected_len); c->recv_buf_len -= c->hs_expected_len; c->hs_expected_len = 0; c->hs_key_processed = 0; uint8_t salt2[SC_PUBKEY_ENC_SALT_SIZE]; if (random_bytes(salt2, SC_PUBKEY_ENC_SALT_SIZE) != 0) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: random_bytes failed"); stcp_conn_do_close(c, 3); return; } uint16_t padding = 8; size_t total_resp = SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + padding; uint8_t *resp = u_malloc(total_resp); if (!resp) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: malloc failed"); stcp_conn_do_close(c, 3); return; } memcpy(resp, salt2, SC_PUBKEY_ENC_SALT_SIZE); sc_obfuscate_pubkey(salt2, c->peer_pubkey, c->my_keys.public_key, resp + SC_PUBKEY_ENC_SALT_SIZE); uint8_t plain_hs[3] = {0, (uint8_t)padding, (uint8_t)(padding >> 8)}; // status=OK + padding_size uint32_t crc = crc32_calc(plain_hs, 3); uint8_t *enc_dst = resp + SC_PUBKEY_ENC_SIZE; enc_dst[0] = plain_hs[0]; enc_dst[1] = plain_hs[1]; enc_dst[2] = plain_hs[2]; enc_dst[3] = (uint8_t)(crc >> 0); enc_dst[4] = (uint8_t)(crc >> 8); enc_dst[5] = (uint8_t)(crc >> 16); enc_dst[6] = (uint8_t)(crc >> 24); if (sc_stream_xor(&c->stream_send, enc_dst, STCP_HS_ENC_SERVER) != SC_OK) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "finish_client_handshake: encrypt failed"); u_free(resp); stcp_conn_do_close(c, 3); return; } for (int i = 0; i < padding; i++) resp[SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_SERVER + i] = (uint8_t)(salt2[0] ^ i); c->state = STCP_STATE_DATA; DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: handshake OK, entering DATA state"); if (c->on_ready) c->on_ready(c, c->ready_arg); stcp_conn_try_send(c, resp, total_resp); } static void server_conn_read_cb(socket_t sock, void *arg) { struct stcp_conn *c = (struct stcp_conn *)arg; if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return; if (c->recv_buf_cap == 0) { c->recv_buf_cap = STCP_RECV_BUF_INIT; c->recv_buf = u_malloc(c->recv_buf_cap); if (!c->recv_buf) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: malloc failed"); stcp_conn_do_close(c, ENOMEM); return; } } if (c->recv_buf_len + 4096 > c->recv_buf_cap) { size_t new_cap = c->recv_buf_cap * 2; if (new_cap > STCP_RECV_BUF_MAX) new_cap = STCP_RECV_BUF_MAX; if (new_cap <= c->recv_buf_cap) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: recv buf full"); stcp_conn_do_close(c, ENOBUFS); return; } uint8_t *nb = u_realloc(c->recv_buf, new_cap); if (!nb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: realloc failed"); stcp_conn_do_close(c, ENOMEM); return; } c->recv_buf = nb; c->recv_buf_cap = new_cap; } ssize_t n = recv(sock, c->recv_buf + c->recv_buf_len, c->recv_buf_cap - c->recv_buf_len, 0); if (n < 0) { int err = socket_get_error(); if (err == ERR_AGAIN || err == ERR_WOULDBLOCK) return; DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: recv failed err=%d", err); stcp_conn_do_close(c, err); return; } if (n == 0) { DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "server_conn_read_cb: EOF"); stcp_conn_do_close(c, 0); return; } c->recv_buf_len += (size_t)n; stcp_conn_process_recv(c); } static void stcp_conn_process_recv(struct stcp_conn *c) { while (c->recv_buf_len > 0) { switch (c->state) { case STCP_STATE_HS_SERVER_WAIT: if (!c->hs_key_processed) { if (c->recv_buf_len < SC_PUBKEY_ENC_SIZE + STCP_HS_ENC_CLIENT) return; process_client_handshake(c); if (c->state != STCP_STATE_HS_SERVER_WAIT) return; } if (c->hs_key_processed && c->recv_buf_len >= c->hs_expected_len) { finish_client_handshake(c); return; } return; case STCP_STATE_DATA: { if (c->recv_buf_len < 2) return; uint16_t msg_size = (uint16_t)c->recv_buf[0] | ((uint16_t)c->recv_buf[1] << 8); if (msg_size > STCP_MAX_MSG_SIZE) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "msg too large: %u", msg_size); stcp_conn_do_close(c, 4); return; } size_t total = 2 + msg_size + 4; if (c->recv_buf_len < total) return; uint8_t *enc_data = c->recv_buf + 2; size_t data_len; if (stcp_decrypt_and_check(enc_data, msg_size + 4, &c->stream_recv, &data_len)) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "data decrypt/CRC failed"); stcp_conn_do_close(c, 5); return; } if (c->rx_queue) { struct ll_entry *e = queue_entry_new(0); if (e) { e->dgram = u_malloc(data_len); if (e->dgram) { if (data_len) memcpy(e->dgram, enc_data, data_len); e->len = (uint16_t)data_len; queue_data_put(c->rx_queue, e); } else { queue_entry_free(e); DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "rx malloc(%zu) failed", data_len); } } else { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "queue_entry_new failed"); } } memmove(c->recv_buf, c->recv_buf + total, c->recv_buf_len - total); c->recv_buf_len -= total; break; } default: return; } } } static void server_accept_cb(socket_t listen_sock, void *arg) { struct stcp_server *srv = (struct stcp_server *)arg; struct sockaddr_storage cli_addr; socklen_t addr_len = sizeof(cli_addr); socket_t cli_sock = accept(listen_sock, (struct sockaddr *)&cli_addr, &addr_len); if (cli_sock == SOCKET_INVALID) { int err = socket_get_error(); DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server accept failed err=%d", err); return; } socket_set_nonblocking(cli_sock); int opt = 1; setsockopt(cli_sock, IPPROTO_TCP, TCP_NODELAY, (const char *)&opt, sizeof(opt)); struct stcp_conn *c = u_calloc(1, sizeof(struct stcp_conn)); if (!c) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server: calloc conn failed"); socket_close_wrapper(cli_sock); return; } c->sock = cli_sock; c->ua = srv->ua; c->state = STCP_STATE_HS_SERVER_WAIT; c->is_server = 1; c->allocated = 1; c->tx_cb = tx_queue_cb; c->my_keys = srv->my_keys; c->on_ready = srv->connect_cb; c->ready_arg = srv->cb_arg; c->socket_id = uasync_add_socket_t(srv->ua, cli_sock, server_conn_read_cb, server_conn_write_cb, NULL, c); DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: accepted connection fd=%d", (int)cli_sock); } struct stcp_server *stcp_server_create(struct UASYNC *ua, uint16_t port, struct SC_MYKEYS *keys, stcp_connect_cb connect_cb, void *arg) { if (!ua || !keys || !connect_cb) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: invalid args"); return NULL; } struct stcp_server *srv = u_calloc(1, sizeof(struct stcp_server)); if (!srv) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: calloc failed"); return NULL; } srv->ua = ua; srv->my_keys = *keys; srv->connect_cb = connect_cb; srv->cb_arg = arg; srv->listen_sock = socket(AF_INET, SOCK_STREAM, 0); if (srv->listen_sock == SOCKET_INVALID) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: socket failed"); u_free(srv); return NULL; } socket_set_nonblocking(srv->listen_sock); int reuse = 1; setsockopt(srv->listen_sock, SOL_SOCKET, SO_REUSEADDR, (const char *)&reuse, sizeof(reuse)); struct sockaddr_in addr; memset(&addr, 0, sizeof(addr)); addr.sin_family = AF_INET; addr.sin_addr.s_addr = INADDR_ANY; addr.sin_port = htons(port); if (bind(srv->listen_sock, (struct sockaddr *)&addr, sizeof(addr)) < 0) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: bind port=%u failed err=%d", port, socket_get_error()); socket_close_wrapper(srv->listen_sock); u_free(srv); return NULL; } if (listen(srv->listen_sock, 16) < 0) { DEBUG_ERROR(DEBUG_CATEGORY_SOCKET, "stcp_server_create: listen failed err=%d", socket_get_error()); socket_close_wrapper(srv->listen_sock); u_free(srv); return NULL; } srv->listen_id = uasync_add_socket_t(ua, srv->listen_sock, server_accept_cb, NULL, NULL, srv); DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server: listening on port %u", port); return srv; } void stcp_server_destroy(struct stcp_server *srv) { if (!srv) return; if (srv->listen_id) uasync_remove_socket_t(srv->ua, srv->listen_sock); if (srv->listen_sock != SOCKET_INVALID) socket_close_wrapper(srv->listen_sock); u_free(srv); DEBUG_INFO(DEBUG_CATEGORY_SOCKET, "stcp_server destroyed"); }