You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

363 lines
14 KiB

// 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"
#include "../lib/debug_config.h"
#include "../lib/platform_compat.h"
#include <stdlib.h>
#include <string.h>
#include <errno.h>
void stcp_conn_set_tx_queue(struct stcp_conn *c, struct ll_queue *q) {
if (!c || !q) return;
c->tx_queue = q;
if (c->tx_cb) queue_set_callback(q, c->tx_cb, c);
}
void stcp_conn_set_rx_queue(struct stcp_conn *c, struct ll_queue *q) {
if (!c || !q) return;
c->rx_queue = q;
}
void stcp_conn_set_on_close(struct stcp_conn *c, void (*cb)(struct stcp_conn *conn, int err, void *arg), void *arg) {
if (!c || !cb) return;
c->on_close = cb;
c->close_arg = arg;
}
void stcp_conn_free(struct stcp_conn *c) {
if (!c) return;
DEBUG_DEBUG(DEBUG_CATEGORY_ETCP, "stcp_conn_free: c=%p sock=%d sock_id=%p recv_buf=%p state=%d allocated=%d",
(void*)c, (int)c->sock, (void*)c->socket_id, (void*)c->recv_buf, (int)c->state, (int)c->allocated);
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; }
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);
}
static int stcp_check_crc(const uint8_t *data, size_t data_len) {
uint32_t recv_crc = ((uint32_t)data[data_len]) | ((uint32_t)data[data_len + 1] << 8) |
((uint32_t)data[data_len + 2] << 16) | ((uint32_t)data[data_len + 3] << 24);
uint32_t calc_crc = crc32_calc(data, data_len);
if (recv_crc != calc_crc) {
DEBUG_ERROR(DEBUG_CATEGORY_ETCP, "stcp CRC mismatch recv=%08x calc=%08x data_len=%zu",
recv_crc, calc_crc, data_len);
return -1;
}
return 0;
}
int stcp_frame_encrypt(struct stcp_conn *c, const uint8_t *data, size_t data_len, uint8_t *output, size_t *output_len) {
uint32_t crc = crc32_calc(data, data_len);
output[0] = (uint8_t)(data_len >> 0);
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 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 = total;
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;
if (stcp_check_crc(data, data_len) != 0) 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, 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 = 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 = 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(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);
if (c->close_after_send) { stcp_conn_do_close(c, 0); return; }
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;
}
}
// ====== 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;
if (stcp_check_crc(plain, data_len) != 0) { stcp_conn_do_close(c, 5); return; }
if (c->recv_on_chunk) c->recv_on_chunk(c, plain, data_len);
if (c->state == STCP_STATE_CLOSED || c->state == STCP_STATE_ERROR) return;
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);
}
static void stcp_free_on_close_cb(void *arg) {
u_free(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);
sc_stream_cleanup(&c->stream_send);
sc_stream_cleanup(&c->stream_recv);
if (c->on_close) {
void (*cb)(struct stcp_conn*, int, void*) = c->on_close;
c->on_close = NULL;
cb(c, err, c->close_arg);
}
if (c->free_on_close) {
uasync_call_soon(c->ua, c->free_on_close, stcp_free_on_close_cb);
c->free_on_close = NULL;
}
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);
}
}