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.
 
 
 
 
 
 

306 lines
13 KiB

#include "eim_nat.h"
#include "config_parser.h"
#include "../lib/debug_config.h"
#include "../lib/mem.h"
#include <string.h>
#include <stdlib.h>
// IP header offsets
#define IP_IHL_OFFSET 0
#define IP_PROTO_OFFSET 9
#define IP_CHECKSUM_OFFSET 10
#define IP_SRC_ADDR_OFFSET 12
#define IP_DST_ADDR_OFFSET 16
#define IP_HDR_MIN_SIZE 20
#define IP_FRAG_MF_MASK 0x2000
#define IP_FRAG_OFF_MASK 0x1FFF
#define TCP_SRC_PORT_OFFSET 0
#define TCP_DST_PORT_OFFSET 2
#define TCP_CHECKSUM_OFFSET 16
#define UDP_SRC_PORT_OFFSET 0
#define UDP_DST_PORT_OFFSET 2
#define UDP_CHECKSUM_OFFSET 6
#define ICMP_TYPE_OFFSET 0
#define ICMP_ID_OFFSET 4
#define ICMP_CHECKSUM_OFFSET 2
#define ICMP_ECHO_REQUEST 8
#define ICMP_ECHO_REPLY 0
static ip_str_t ip_host_to_str(uint32_t ip_host) {
struct in_addr a; a.s_addr = htonl(ip_host);
return ip_to_str(&a, AF_INET);
}
static uint16_t csum_update_n(uint16_t old_csum, const uint16_t* old_vals,
const uint16_t* new_vals, int n) {
uint32_t sum = (uint32_t)(uint16_t)(~old_csum);
for (int i = 0; i < n; i++) { sum -= old_vals[i]; sum += new_vals[i]; }
sum = (sum & 0xFFFF) + (sum >> 16); sum = (sum & 0xFFFF) + (sum >> 16);
return (uint16_t)(~sum);
}
static void csum_update_ip(uint8_t* ip_data, uint32_t old_ip_host, uint32_t new_ip_host) {
uint16_t old_csum; memcpy(&old_csum, ip_data + IP_CHECKSUM_OFFSET, 2);
uint32_t old_ip_net = htonl(old_ip_host), new_ip_net = htonl(new_ip_host);
uint16_t old_w[2] = {(uint16_t)(old_ip_net & 0xFFFF), (uint16_t)(old_ip_net >> 16)};
uint16_t new_w[2] = {(uint16_t)(new_ip_net & 0xFFFF), (uint16_t)(new_ip_net >> 16)};
uint16_t new_csum = csum_update_n(old_csum, old_w, new_w, 2);
memcpy(ip_data + IP_CHECKSUM_OFFSET, &new_csum, 2);
}
static void csum_update_transport(uint8_t* transport, int csum_off,
uint32_t old_ip_host, uint32_t new_ip_host,
uint16_t old_port_net, uint16_t new_port_net) {
uint16_t old_csum; memcpy(&old_csum, transport + csum_off, 2);
uint32_t old_ip_net = htonl(old_ip_host), new_ip_net = htonl(new_ip_host);
uint16_t old_w[3] = {(uint16_t)(old_ip_net & 0xFFFF), (uint16_t)(old_ip_net >> 16), old_port_net};
uint16_t new_w[3] = {(uint16_t)(new_ip_net & 0xFFFF), (uint16_t)(new_ip_net >> 16), new_port_net};
uint16_t new_csum = csum_update_n(old_csum, old_w, new_w, 3);
memcpy(transport + csum_off, &new_csum, 2);
}
static void csum_update_icmp(uint8_t* icmp, uint16_t old_word, uint16_t new_word) {
uint16_t old_csum; memcpy(&old_csum, icmp + ICMP_CHECKSUM_OFFSET, 2);
uint16_t new_csum = csum_update_n(old_csum, &old_word, &new_word, 1);
memcpy(icmp + ICMP_CHECKSUM_OFFSET, &new_csum, 2);
}
static struct eim_nat_entry* eim_nat_find_egress(struct eim_nat_ctx* ctx,
uint8_t proto, uint32_t ip_host, uint16_t port_net) {
for (uint16_t i = ctx->port_start; i <= ctx->port_end; i++) {
struct eim_nat_entry* e = &ctx->table[i];
if (e->state != EIM_NAT_ENTRY_FREE && e->proto == proto &&
e->internal_ip == ip_host && e->internal_port == port_net) return e;
}
return NULL;
}
static uint16_t eim_nat_alloc_port(struct eim_nat_ctx* ctx) {
uint16_t start = ctx->next_port;
do {
if (ctx->table[ctx->next_port].state == EIM_NAT_ENTRY_FREE) {
uint16_t port = ctx->next_port;
ctx->next_port++;
if (ctx->next_port > ctx->port_end || ctx->next_port < ctx->port_start) ctx->next_port = ctx->port_start;
return port;
}
ctx->next_port++;
if (ctx->next_port > ctx->port_end || ctx->next_port < ctx->port_start) ctx->next_port = ctx->port_start;
} while (ctx->next_port != start);
return 0;
}
// ==================== Public API ====================
int eim_nat_egress(struct eim_nat_ctx* ctx, uint8_t* ip_data, size_t ip_len,
uint64_t src_node_id, struct ETCP_CONN* src_conn) {
if (ip_len < IP_HDR_MIN_SIZE) return -1;
uint8_t ihl = ip_data[IP_IHL_OFFSET] & 0x0F;
if (ihl < 5) return -1;
uint16_t ip_hdr_len = ihl * 4;
if (ip_len < ip_hdr_len) return -1;
uint16_t frag_off = ntohs(*(uint16_t*)(ip_data + 6));
if ((frag_off & (IP_FRAG_MF_MASK | IP_FRAG_OFF_MASK)) != 0) return 0;
uint8_t proto = ip_data[IP_PROTO_OFFSET];
uint32_t src_ip_net, src_ip_host;
memcpy(&src_ip_net, ip_data + IP_SRC_ADDR_OFFSET, 4);
src_ip_host = ntohl(src_ip_net);
uint8_t* transport = ip_data + ip_hdr_len;
size_t tlen = ip_len - ip_hdr_len;
uint16_t src_port_net = 0;
int is_tcp = (proto == IPPROTO_TCP_UINT8 && tlen >= 20);
int is_udp = (proto == IPPROTO_UDP_UINT8 && tlen >= 8);
int is_icmp_echo = 0;
if (is_tcp) memcpy(&src_port_net, transport + TCP_SRC_PORT_OFFSET, 2);
else if (is_udp) memcpy(&src_port_net, transport + UDP_SRC_PORT_OFFSET, 2);
else if (proto == IPPROTO_ICMP_UINT8 && tlen >= 8) {
uint8_t icmp_type = transport[ICMP_TYPE_OFFSET];
if (icmp_type == ICMP_ECHO_REQUEST || icmp_type == ICMP_ECHO_REPLY) {
is_icmp_echo = 1; memcpy(&src_port_net, transport + ICMP_ID_OFFSET, 2);
}
}
if (!is_tcp && !is_udp && !is_icmp_echo) return 0;
struct eim_nat_entry* entry = eim_nat_find_egress(ctx, proto, src_ip_host, src_port_net);
uint16_t ext_port_host;
if (entry) {
ext_port_host = (uint16_t)(entry - ctx->table);
} else {
ext_port_host = eim_nat_alloc_port(ctx);
if (ext_port_host == 0) {
DEBUG_ERROR(DEBUG_CATEGORY_NAT, "NAT port range exhausted (%s:%u proto=%u)",
ip_host_to_str(src_ip_host).str, ntohs(src_port_net), proto);
return -1;
}
entry = &ctx->table[ext_port_host];
entry->internal_ip = src_ip_host;
entry->internal_port = src_port_net;
entry->proto = proto;
entry->state = EIM_NAT_ENTRY_ACTIVE;
entry->src_node_id = src_node_id;
entry->src_conn = src_conn;
entry->last_seen = 0;
DEBUG_DEBUG(DEBUG_CATEGORY_NAT, "NAT map: %s:%u -> %s:%u (proto=%u) from node %016llx",
ip_host_to_str(src_ip_host).str, ntohs(src_port_net),
ip_host_to_str(ctx->gateway_ip).str, ext_port_host, proto,
(unsigned long long)src_node_id);
}
uint32_t gw_ip_net = htonl(ctx->gateway_ip);
memcpy(ip_data + IP_SRC_ADDR_OFFSET, &gw_ip_net, 4);
csum_update_ip(ip_data, src_ip_host, ctx->gateway_ip);
uint16_t new_port_net = htons(ext_port_host);
if (is_tcp) {
memcpy(transport + TCP_SRC_PORT_OFFSET, &new_port_net, 2);
csum_update_transport(transport, TCP_CHECKSUM_OFFSET, src_ip_host, ctx->gateway_ip, src_port_net, new_port_net);
} else if (is_udp) {
memcpy(transport + UDP_SRC_PORT_OFFSET, &new_port_net, 2);
csum_update_transport(transport, UDP_CHECKSUM_OFFSET, src_ip_host, ctx->gateway_ip, src_port_net, new_port_net);
} else if (is_icmp_echo) {
memcpy(transport + ICMP_ID_OFFSET, &new_port_net, 2);
csum_update_icmp(transport, src_port_net, new_port_net);
}
return 0;
}
int eim_nat_ingress(struct eim_nat_ctx* ctx, uint8_t* ip_data, size_t ip_len,
struct eim_nat_entry** out_entry) {
if (ip_len < IP_HDR_MIN_SIZE) return -1;
uint8_t ihl = ip_data[IP_IHL_OFFSET] & 0x0F;
if (ihl < 5) return -1;
uint16_t ip_hdr_len = ihl * 4;
if (ip_len < ip_hdr_len) return -1;
uint16_t frag_off = ntohs(*(uint16_t*)(ip_data + 6));
if ((frag_off & (IP_FRAG_MF_MASK | IP_FRAG_OFF_MASK)) != 0) return 0;
uint8_t proto = ip_data[IP_PROTO_OFFSET];
uint32_t dst_ip_net;
memcpy(&dst_ip_net, ip_data + IP_DST_ADDR_OFFSET, 4);
if (dst_ip_net != htonl(ctx->gateway_ip)) return 0;
uint8_t* transport = ip_data + ip_hdr_len;
size_t tlen = ip_len - ip_hdr_len;
uint16_t dst_port_net = 0;
int is_tcp = (proto == IPPROTO_TCP_UINT8 && tlen >= 20);
int is_udp = (proto == IPPROTO_UDP_UINT8 && tlen >= 8);
int is_icmp_echo = 0;
if (is_tcp) memcpy(&dst_port_net, transport + TCP_DST_PORT_OFFSET, 2);
else if (is_udp) memcpy(&dst_port_net, transport + UDP_DST_PORT_OFFSET, 2);
else if (proto == IPPROTO_ICMP_UINT8 && tlen >= 8) {
uint8_t icmp_type = transport[ICMP_TYPE_OFFSET];
if (icmp_type == ICMP_ECHO_REQUEST || icmp_type == ICMP_ECHO_REPLY) {
is_icmp_echo = 1; memcpy(&dst_port_net, transport + ICMP_ID_OFFSET, 2);
}
}
if (!is_tcp && !is_udp && !is_icmp_echo) return 0;
uint16_t ext_port_host = ntohs(dst_port_net);
if (ext_port_host < ctx->port_start || ext_port_host > ctx->port_end) return 0;
struct eim_nat_entry* entry = &ctx->table[ext_port_host];
if (entry->state == EIM_NAT_ENTRY_FREE) return 0;
if (entry->proto != proto) return 0;
uint32_t dst_ip_host = ntohl(dst_ip_net);
uint32_t internal_ip_net = htonl(entry->internal_ip);
memcpy(ip_data + IP_DST_ADDR_OFFSET, &internal_ip_net, 4);
csum_update_ip(ip_data, dst_ip_host, entry->internal_ip);
uint16_t internal_port_net = entry->internal_port;
if (is_tcp) {
memcpy(transport + TCP_DST_PORT_OFFSET, &internal_port_net, 2);
csum_update_transport(transport, TCP_CHECKSUM_OFFSET, dst_ip_host, entry->internal_ip, dst_port_net, internal_port_net);
} else if (is_udp) {
memcpy(transport + UDP_DST_PORT_OFFSET, &internal_port_net, 2);
csum_update_transport(transport, UDP_CHECKSUM_OFFSET, dst_ip_host, entry->internal_ip, dst_port_net, internal_port_net);
} else if (is_icmp_echo) {
memcpy(transport + ICMP_ID_OFFSET, &internal_port_net, 2);
csum_update_icmp(transport, dst_port_net, internal_port_net);
}
DEBUG_DEBUG(DEBUG_CATEGORY_NAT, "NAT unmap: %s:%u <- %s:%u (proto=%u) back to node %016llx",
ip_host_to_str(entry->internal_ip).str, ntohs(entry->internal_port),
ip_host_to_str(ctx->gateway_ip).str, ext_port_host, proto,
(unsigned long long)entry->src_node_id);
if (out_entry) *out_entry = entry;
return 0;
}
// ==================== Init / Destroy ====================
int eim_nat_init_ctx(struct eim_nat_ctx* ctx, const struct global_config* g) {
if (!ctx || !g) return -1;
memset(ctx, 0, sizeof(*ctx));
if (!g->nat_enabled) return 0;
if (g->nat_tun_ip.family == AF_INET) {
ctx->gateway_ip = ntohl(g->nat_tun_ip.addr.v4.s_addr);
} else if (g->tun_ip.family == AF_INET) {
ctx->gateway_ip = ntohl(g->tun_ip.addr.v4.s_addr);
} else {
DEBUG_ERROR(DEBUG_CATEGORY_NAT, "No valid IP for NAT gateway");
return -1;
}
ctx->port_start = g->nat_port_start;
ctx->port_end = g->nat_port_end;
if (ctx->port_start == 0 || ctx->port_end == 0 || ctx->port_start >= ctx->port_end) {
DEBUG_ERROR(DEBUG_CATEGORY_NAT, "Invalid NAT port range: %u-%u", ctx->port_start, ctx->port_end);
return -1;
}
ctx->next_port = ctx->port_start;
ctx->table_size = EIM_NAT_TABLE_SIZE;
ctx->table = u_calloc(ctx->table_size, sizeof(struct eim_nat_entry));
if (!ctx->table) { DEBUG_ERROR(DEBUG_CATEGORY_NAT, "Failed to allocate NAT table"); return -1; }
for (int i = 0; i < g->nat_forward_count; i++) {
uint16_t ext = g->nat_forwards[i].external_port;
if (ext >= ctx->table_size) continue;
struct eim_nat_entry* e = &ctx->table[ext];
e->internal_ip = g->nat_forwards[i].internal_ip_host;
e->internal_port = g->nat_forwards[i].internal_port_net;
e->proto = g->nat_forwards[i].proto;
e->state = EIM_NAT_ENTRY_STATIC;
DEBUG_INFO(DEBUG_CATEGORY_NAT, "Port forward: %s:%u <- :%u (proto=%u)",
ip_host_to_str(e->internal_ip).str, ntohs(e->internal_port), ext, e->proto);
}
ctx->initialized = 1;
DEBUG_INFO(DEBUG_CATEGORY_NAT, "NAT engine initialized: gw=%s ports=%u-%u",
ip_host_to_str(ctx->gateway_ip).str, ctx->port_start, ctx->port_end);
return 0;
}
void eim_nat_destroy_ctx(struct eim_nat_ctx* ctx) {
if (!ctx || !ctx->initialized) return;
u_free(ctx->table);
memset(ctx, 0, sizeof(*ctx));
DEBUG_INFO(DEBUG_CATEGORY_NAT, "NAT engine destroyed");
}
int eim_nat_add_forward(struct eim_nat_ctx* ctx, uint8_t proto,
uint32_t internal_ip_host, uint16_t internal_port_net,
uint16_t external_port) {
if (!ctx || !ctx->initialized || external_port >= ctx->table_size) return -1;
if (ctx->table[external_port].state != EIM_NAT_ENTRY_FREE) return -1;
struct eim_nat_entry* e = &ctx->table[external_port];
e->internal_ip = internal_ip_host;
e->internal_port = internal_port_net;
e->proto = proto;
e->state = EIM_NAT_ENTRY_STATIC;
DEBUG_INFO(DEBUG_CATEGORY_NAT, "Port forward: %s:%u <- :%u (proto=%u)",
ip_host_to_str(internal_ip_host).str, ntohs(internal_port_net), external_port, proto);
return 0;
}