diff --git a/src/payload.c b/src/payload.c new file mode 100644 index 0000000..27f07e2 --- /dev/null +++ b/src/payload.c @@ -0,0 +1,263 @@ +#include "rpc/protocol.h" + +#include +#include +#include + +static void put_u32(uint8_t *out, uint32_t value) { + out[0] = (uint8_t)(value >> 24u); + out[1] = (uint8_t)(value >> 16u); + out[2] = (uint8_t)(value >> 8u); + out[3] = (uint8_t)value; +} + +static void put_u64(uint8_t *out, uint64_t value) { + for (int i = 7; i >= 0; --i) { + out[7 - i] = (uint8_t)(value >> (unsigned)(i * 8)); + } +} + +static uint32_t get_u32(const uint8_t *in) { + return ((uint32_t)in[0] << 24u) | ((uint32_t)in[1] << 16u) | + ((uint32_t)in[2] << 8u) | (uint32_t)in[3]; +} + +static uint64_t get_u64(const uint8_t *in) { + uint64_t value = 0; + for (size_t i = 0; i < 8; ++i) { + value = (value << 8u) | (uint64_t)in[i]; + } + return value; +} + +static int writer_reserve(rpc_writer *writer, size_t extra) { + if (!writer || extra > RPC_MAX_PAYLOAD_SIZE || + writer->len > RPC_MAX_PAYLOAD_SIZE - extra) { + return -1; + } + size_t need = writer->len + extra; + if (need <= writer->cap) { + return 0; + } + + size_t cap = writer->cap ? writer->cap : 64u; + while (cap < need) { + if (cap > RPC_MAX_PAYLOAD_SIZE / 2u) { + cap = RPC_MAX_PAYLOAD_SIZE; + break; + } + cap *= 2u; + } + uint8_t *next = realloc(writer->data, cap); + if (!next) { + return -1; + } + writer->data = next; + writer->cap = cap; + return 0; +} + +static int writer_push(rpc_writer *writer, const void *data, size_t len) { + if (writer_reserve(writer, len) != 0) { + return -1; + } + if (len > 0) { + memcpy(writer->data + writer->len, data, len); + } + writer->len += len; + return 0; +} + +void rpc_writer_init(rpc_writer *writer) { + if (writer) { + memset(writer, 0, sizeof(*writer)); + } +} + +void rpc_writer_reset(rpc_writer *writer) { + if (writer) { + writer->len = 0; + } +} + +void rpc_writer_free(rpc_writer *writer) { + if (writer) { + free(writer->data); + memset(writer, 0, sizeof(*writer)); + } +} + +int rpc_writer_null(rpc_writer *writer) { + uint8_t type = RPC_TYPE_NULL; + return writer_push(writer, &type, 1); +} + +int rpc_writer_bool(rpc_writer *writer, bool value) { + uint8_t data[2] = {RPC_TYPE_BOOL, value ? 1u : 0u}; + return writer_push(writer, data, sizeof(data)); +} + +int rpc_writer_i64(rpc_writer *writer, int64_t value) { + uint8_t data[9]; + data[0] = RPC_TYPE_I64; + put_u64(data + 1, (uint64_t)value); + return writer_push(writer, data, sizeof(data)); +} + +int rpc_writer_u64(rpc_writer *writer, uint64_t value) { + uint8_t data[9]; + data[0] = RPC_TYPE_U64; + put_u64(data + 1, value); + return writer_push(writer, data, sizeof(data)); +} + +int rpc_writer_f64(rpc_writer *writer, double value) { + uint8_t data[9]; + uint64_t bits = 0; + memcpy(&bits, &value, sizeof(bits)); + data[0] = RPC_TYPE_F64; + put_u64(data + 1, bits); + return writer_push(writer, data, sizeof(data)); +} + +int rpc_writer_bytes(rpc_writer *writer, const void *data, uint32_t len) { + uint8_t prefix[5]; + if (len > 0 && !data) { + return -1; + } + prefix[0] = RPC_TYPE_BYTES; + put_u32(prefix + 1, len); + if (writer_push(writer, prefix, sizeof(prefix)) != 0) { + return -1; + } + return writer_push(writer, data, len); +} + +int rpc_writer_string(rpc_writer *writer, const char *data, uint32_t len) { + uint8_t prefix[5]; + if (len > 0 && !data) { + return -1; + } + prefix[0] = RPC_TYPE_STRING; + put_u32(prefix + 1, len); + if (writer_push(writer, prefix, sizeof(prefix)) != 0) { + return -1; + } + return writer_push(writer, data, len); +} + +int rpc_payload_decode(const uint8_t *data, size_t len, rpc_value **out_values, + size_t *out_count) { + static const void *dispatch[] = { + [RPC_TYPE_NULL] = &&type_null, + [RPC_TYPE_BOOL] = &&type_bool, + [RPC_TYPE_I64] = &&type_i64, + [RPC_TYPE_U64] = &&type_u64, + [RPC_TYPE_F64] = &&type_f64, + [RPC_TYPE_BYTES] = &&type_bytes, + [RPC_TYPE_STRING] = &&type_string, + }; + + if ((!data && len > 0) || !out_values || !out_count) { + return -1; + } + + rpc_value *values = NULL; + size_t count = 0; + size_t cap = 0; + size_t off = 0; + + while (off < len) { + if (count == cap) { + size_t next_cap = cap ? cap * 2u : 4u; + rpc_value *next = realloc(values, next_cap * sizeof(*values)); + if (!next) { + free(values); + return -1; + } + values = next; + cap = next_cap; + } + + rpc_value value; + memset(&value, 0, sizeof(value)); + value.type = (rpc_type)data[off++]; + + if ((size_t)value.type >= sizeof(dispatch) / sizeof(*dispatch) || + !dispatch[value.type]) { + goto malformed; + } + goto *dispatch[value.type]; + +type_null: + goto store; + +type_bool: + if (off + 1u > len || (data[off] != 0u && data[off] != 1u)) { + goto malformed; + } + value.as.boolean = data[off++] != 0u; + goto store; + +type_i64: + if (off + 8u > len) { + goto malformed; + } + value.as.i64 = (int64_t)get_u64(data + off); + off += 8u; + goto store; + +type_u64: + if (off + 8u > len) { + goto malformed; + } + value.as.u64 = get_u64(data + off); + off += 8u; + goto store; + +type_f64: { + if (off + 8u > len) { + goto malformed; + } + uint64_t bits = get_u64(data + off); + memcpy(&value.as.f64, &bits, sizeof(bits)); + off += 8u; + goto store; + } + +type_bytes: +type_string: { + if (off + 4u > len) { + goto malformed; + } + uint32_t value_len = get_u32(data + off); + off += 4u; + if (off + value_len > len) { + goto malformed; + } + if (value.type == RPC_TYPE_BYTES) { + value.as.bytes.data = data + off; + value.as.bytes.len = value_len; + } else { + value.as.string.data = (const char *)(data + off); + value.as.string.len = value_len; + } + off += value_len; + goto store; + } + +store: + values[count++] = value; + continue; + +malformed: + free(values); + return -1; + } + + *out_values = values; + *out_count = count; + return 0; +} + +void rpc_values_free(rpc_value *values) { free(values); } diff --git a/src/protocol.c b/src/protocol.c new file mode 100644 index 0000000..29787d5 --- /dev/null +++ b/src/protocol.c @@ -0,0 +1,65 @@ +#include "rpc/protocol.h" + +#include + +static void put_u32(uint8_t *out, uint32_t value) { + out[0] = (uint8_t)(value >> 24u); + out[1] = (uint8_t)(value >> 16u); + out[2] = (uint8_t)(value >> 8u); + out[3] = (uint8_t)value; +} + +static void put_u64(uint8_t *out, uint64_t value) { + for (int i = 7; i >= 0; --i) { + out[7 - i] = (uint8_t)(value >> (unsigned)(i * 8)); + } +} + +static uint32_t get_u32(const uint8_t *in) { + return ((uint32_t)in[0] << 24u) | ((uint32_t)in[1] << 16u) | + ((uint32_t)in[2] << 8u) | (uint32_t)in[3]; +} + +static uint64_t get_u64(const uint8_t *in) { + uint64_t value = 0; + for (size_t i = 0; i < 8; ++i) { + value = (value << 8u) | (uint64_t)in[i]; + } + return value; +} + +static int valid_op(uint8_t op) { + return op == RPC_OP_RPC || op == RPC_OP_PING || op == RPC_OP_DISCONNECT || + op == RPC_OP_RESPONSE || op == RPC_OP_ERROR; +} + +int rpc_header_encode(const rpc_header *header, uint8_t out[RPC_HEADER_SIZE]) { + if (!header || !out || !valid_op((uint8_t)header->op) || + header->size > RPC_MAX_PAYLOAD_SIZE) { + return -1; + } + + out[0] = (uint8_t)header->op; + out[1] = header->flags; + put_u32(out + 2, header->proc_id); + put_u32(out + 6, header->size); + put_u64(out + 10, header->call_id); + return 0; +} + +int rpc_header_decode(const uint8_t in[RPC_HEADER_SIZE], rpc_header *out) { + if (!in || !out || !valid_op(in[0])) { + return -1; + } + + memset(out, 0, sizeof(*out)); + out->op = (rpc_op)in[0]; + out->flags = in[1]; + out->proc_id = get_u32(in + 2); + out->size = get_u32(in + 6); + out->call_id = get_u64(in + 10); + if (out->size > RPC_MAX_PAYLOAD_SIZE) { + return -1; + } + return 0; +} diff --git a/tests/test_protocol.c b/tests/test_protocol.c new file mode 100644 index 0000000..2aaf0a3 --- /dev/null +++ b/tests/test_protocol.c @@ -0,0 +1,57 @@ +#include "rpc/protocol.h" + +#include +#include +#include + +int main(void) { + uint8_t buf[RPC_HEADER_SIZE]; + rpc_header h = { + .op = RPC_OP_RPC, + .flags = RPC_FLAG_MORE, + .proc_id = 0x11223344u, + .size = 9, + .call_id = 0x0102030405060708ull, + }; + assert(rpc_header_encode(&h, buf) == 0); + rpc_header out; + assert(rpc_header_decode(buf, &out) == 0); + assert(out.op == h.op); + assert(out.flags == h.flags); + assert(out.proc_id == h.proc_id); + assert(out.size == h.size); + assert(out.call_id == h.call_id); + + buf[0] = 99; + assert(rpc_header_decode(buf, &out) != 0); + + rpc_writer w; + rpc_writer_init(&w); + assert(rpc_writer_null(&w) == 0); + assert(rpc_writer_bool(&w, true) == 0); + assert(rpc_writer_i64(&w, -42) == 0); + assert(rpc_writer_u64(&w, 42) == 0); + assert(rpc_writer_f64(&w, 3.5) == 0); + assert(rpc_writer_bytes(&w, "abc", 3) == 0); + assert(rpc_writer_string(&w, "hello", 5) == 0); + + rpc_value *values = NULL; + size_t count = 0; + assert(rpc_payload_decode(w.data, w.len, &values, &count) == 0); + assert(count == 7); + assert(values[0].type == RPC_TYPE_NULL); + assert(values[1].as.boolean); + assert(values[2].as.i64 == -42); + assert(values[3].as.u64 == 42); + assert(fabs(values[4].as.f64 - 3.5) < 0.001); + assert(values[5].as.bytes.len == 3); + assert(memcmp(values[5].as.bytes.data, "abc", 3) == 0); + assert(values[6].as.string.len == 5); + assert(memcmp(values[6].as.string.data, "hello", 5) == 0); + rpc_values_free(values); + + uint8_t malformed[] = {RPC_TYPE_STRING, 0, 0, 0, 10, 'x'}; + assert(rpc_payload_decode(malformed, sizeof(malformed), &values, &count) != 0); + rpc_writer_free(&w); + return 0; +}