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
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); |
|
} |
|
}
|
|
|