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.
237 lines
8.8 KiB
237 lines
8.8 KiB
#include "invite_link_c.h" |
|
#include <string.h> |
|
#include <stdlib.h> |
|
#include <stdio.h> |
|
|
|
#include "../../../src/transport_layer/secure_channel.h" |
|
|
|
#define INVITE_PREFIX "utun://" |
|
#define INVITE_PREFIX_LEN 7 |
|
|
|
static const char base64_table[] = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"; |
|
|
|
static int base64_char_val(char c) { |
|
if (c >= 'A' && c <= 'Z') return c - 'A'; |
|
if (c >= 'a' && c <= 'z') return c - 'a' + 26; |
|
if (c >= '0' && c <= '9') return c - '0' + 52; |
|
if (c == '+') return 62; |
|
if (c == '/') return 63; |
|
return -1; |
|
} |
|
|
|
static int base64_decode(const char* in, size_t in_len, uint8_t* out, size_t out_cap) { |
|
size_t opos = 0; |
|
int group[4]; |
|
int gi = 0; |
|
for (size_t i = 0; i < in_len && in[i] != '='; i++) { |
|
int v = base64_char_val(in[i]); |
|
if (v < 0) return -1; |
|
group[gi++] = v; |
|
if (gi == 4) { |
|
out[opos++] = (uint8_t)((group[0] << 2) | (group[1] >> 4)); |
|
if (opos >= out_cap) return -1; |
|
out[opos++] = (uint8_t)(((group[1] & 0xF) << 4) | (group[2] >> 2)); |
|
if (opos >= out_cap) return -1; |
|
out[opos++] = (uint8_t)(((group[2] & 0x3) << 6) | group[3]); |
|
if (opos >= out_cap) return -1; |
|
gi = 0; |
|
} |
|
} |
|
if (gi == 2) { |
|
out[opos++] = (uint8_t)((group[0] << 2) | (group[1] >> 4)); |
|
} else if (gi == 3) { |
|
out[opos++] = (uint8_t)((group[0] << 2) | (group[1] >> 4)); |
|
out[opos++] = (uint8_t)(((group[1] & 0xF) << 4) | (group[2] >> 2)); |
|
} |
|
return (int)opos; |
|
} |
|
|
|
int invite_link_decode(const char* link, size_t link_len, struct InviteDataC* out, |
|
char* error_buf, size_t error_buf_size) { |
|
if (!link || !out) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "null argument"); |
|
return -1; |
|
} |
|
if (link_len < INVITE_PREFIX_LEN || strncmp(link, INVITE_PREFIX, INVITE_PREFIX_LEN) != 0) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "invalid prefix (expected utun://)"); |
|
return -1; |
|
} |
|
const char* b64 = link + INVITE_PREFIX_LEN; |
|
size_t b64_len = link_len - INVITE_PREFIX_LEN; |
|
|
|
uint8_t raw[4096]; |
|
int raw_len = base64_decode(b64, b64_len, raw, sizeof(raw)); |
|
if (raw_len < 0) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "base64 decode failed"); |
|
return -1; |
|
} |
|
if (raw_len < 11) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "blob too short (%d bytes)", raw_len); |
|
return -1; |
|
} |
|
|
|
int off = 0; |
|
uint8_t ver = raw[off++]; |
|
if (ver != INVITE_LINK_VERSION && ver != 0x02) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "unsupported version 0x%02x", ver); |
|
return -1; |
|
} |
|
out->password[0] = '\0'; out->password_len = 0; |
|
|
|
/* v2: password field before addresses */ |
|
if (ver == 0x02) { |
|
if (off + 1 > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at pass_len"); |
|
return -1; |
|
} |
|
uint8_t plen = raw[off++]; |
|
if (plen > 0) { |
|
if (plen > INVITE_PASS_MAX - 1) { if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "password too long %d", plen); return -1; } |
|
if (off + plen > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at password"); |
|
return -1; |
|
} |
|
memcpy(out->password, raw + off, plen); out->password[plen] = '\0'; out->password_len = plen; |
|
off += plen; |
|
} |
|
} |
|
|
|
if (off + 8 > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at channel_id"); |
|
return -1; |
|
} |
|
uint64_t chId = 0; |
|
for (int i = 0; i < 8; i++) |
|
chId = (chId << 8) | raw[off++]; |
|
out->channelId = chId; |
|
|
|
if (off + 8 > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at join_key"); |
|
return -1; |
|
} |
|
uint64_t jk = 0; |
|
for (int i = 0; i < 8; i++) |
|
jk = (jk << 8) | raw[off++]; |
|
out->join_key = jk; |
|
|
|
memset(out->pubkey, 0, INVITE_PUBKEY_SIZE); |
|
out->addrCount = 0; |
|
|
|
int first_block = 1; |
|
while (off + 1 <= raw_len) { |
|
uint8_t header = raw[off++]; |
|
int cnt = (header & 0x03) + 1; |
|
if (cnt < 1 || cnt > 4) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "invalid addr count %d", cnt); |
|
return -1; |
|
} |
|
if (off + 32 > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at pubkey"); |
|
return -1; |
|
} |
|
if (first_block) { |
|
memcpy(out->pubkey, raw + off, INVITE_PUBKEY_SIZE); |
|
out->nodeId = sc_derive_node_id_from_pubkey(out->pubkey); |
|
first_block = 0; |
|
} |
|
off += 32; |
|
for (int j = 0; j < cnt; j++) { |
|
int is_v6 = (header >> (2 + j)) & 1; |
|
int ip_len = is_v6 ? 16 : 4; |
|
if (off + 2 + ip_len + 2 > raw_len) { |
|
if (error_buf && error_buf_size) snprintf(error_buf, error_buf_size, "truncated at addr %d", j); |
|
return -1; |
|
} |
|
if (out->addrCount >= INVITE_ADDR_MAX) break; |
|
struct InviteAddrC* a = &out->addrs[out->addrCount++]; |
|
a->socketId = raw[off++]; |
|
a->proto = raw[off++]; |
|
a->family = is_v6 ? 6 : 4; |
|
memcpy(a->address, raw + off, (size_t)ip_len); |
|
off += ip_len; |
|
a->port = (uint16_t)((raw[off] << 8) | raw[off + 1]); |
|
off += 2; |
|
} |
|
} |
|
return 0; |
|
} |
|
|
|
int invite_serialize_addrs(const struct InviteDataC* data, uint8_t* buf, size_t buf_size) { |
|
if (!data || !buf) return -1; |
|
size_t needed = 0; |
|
for (int i = 0; i < (int)data->addrCount; i++) |
|
needed += (data->addrs[i].family == 6) ? 21U : 9U; |
|
if (needed > buf_size) return -1; |
|
size_t pos = 0; |
|
for (int i = 0; i < (int)data->addrCount; i++) { |
|
const struct InviteAddrC* a = &data->addrs[i]; |
|
int ip_len = (a->family == 6) ? 16 : 4; |
|
buf[pos++] = (uint8_t)a->family; |
|
buf[pos++] = a->socketId; |
|
buf[pos++] = a->proto; |
|
memcpy(buf + pos, a->address, (size_t)ip_len); pos += ip_len; |
|
buf[pos++] = (uint8_t)(a->port >> 8); |
|
buf[pos++] = (uint8_t)(a->port & 0xFF); |
|
} |
|
return (int)pos; |
|
} |
|
|
|
int invite_link_encode(const struct InviteDataC* data, const char* password, char* out, size_t out_size) { |
|
if (!data || !out || data->pubkey[0] == 0) return -1; |
|
if (data->addrCount == 0) return -1; |
|
|
|
int has_pass = (password && password[0]) ? 1 : 0; |
|
uint8_t raw[4096]; |
|
size_t pos = 0; |
|
raw[pos++] = has_pass ? 0x02 : INVITE_LINK_VERSION; |
|
if (has_pass) { |
|
size_t plen = strlen(password); |
|
if (plen > INVITE_PASS_MAX - 1) return -1; |
|
raw[pos++] = (uint8_t)plen; |
|
memcpy(raw + pos, password, plen); pos += plen; |
|
} |
|
for (int i = 7; i >= 0; i--) raw[pos++] = (uint8_t)((data->channelId >> (i * 8)) & 0xFF); |
|
for (int i = 7; i >= 0; i--) raw[pos++] = (uint8_t)((data->join_key >> (i * 8)) & 0xFF); |
|
uint8_t pk[32]; memcpy(pk, data->pubkey, 32); |
|
|
|
int ia = 0; |
|
while (ia < (int)data->addrCount) { |
|
int rem = (int)data->addrCount - ia; |
|
int cnt = (rem > 4) ? 4 : rem; |
|
uint8_t header = (uint8_t)((cnt - 1) & 0x03); |
|
for (int j = 0; j < cnt; j++) { |
|
if (data->addrs[ia + j].family == 6) header |= (uint8_t)(1 << (2 + j)); |
|
} |
|
raw[pos++] = header; |
|
memcpy(raw + pos, pk, 32); pos += 32; |
|
for (int j = 0; j < cnt; j++) { |
|
const struct InviteAddrC* a = &data->addrs[ia + j]; |
|
int ip_len = (a->family == 6) ? 16 : 4; |
|
raw[pos++] = a->socketId; |
|
raw[pos++] = a->proto; |
|
memcpy(raw + pos, a->address, (size_t)ip_len); pos += (size_t)ip_len; |
|
raw[pos++] = (uint8_t)(a->port >> 8); |
|
raw[pos++] = (uint8_t)(a->port & 0xFF); |
|
} |
|
ia += cnt; |
|
} |
|
|
|
/* base64 encode */ |
|
if (out_size < 12 + (pos * 4 + 2) / 3 + 1) return -1; |
|
memcpy(out, INVITE_PREFIX, INVITE_PREFIX_LEN); |
|
size_t opos = INVITE_PREFIX_LEN; |
|
size_t i = 0; |
|
while (i < pos) { |
|
uint32_t val = (uint32_t)raw[i] << 16; |
|
val |= (i + 1 < pos) ? (uint32_t)raw[i + 1] << 8 : 0; |
|
val |= (i + 2 < pos) ? (uint32_t)raw[i + 2] : 0; |
|
out[opos++] = base64_table[(val >> 18) & 0x3F]; |
|
out[opos++] = base64_table[(val >> 12) & 0x3F]; |
|
out[opos++] = (i + 1 < pos) ? base64_table[(val >> 6) & 0x3F] : '='; |
|
out[opos++] = (i + 2 < pos) ? base64_table[val & 0x3F] : '='; |
|
i += 3; |
|
} |
|
if (opos >= out_size) return -1; |
|
out[opos] = '\0'; |
|
return (int)(opos - INVITE_PREFIX_LEN); |
|
}
|
|
|