#include "invite_link_c.h" #include #include #include #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); }