#include "eim_nat.h" #include "config_parser.h" #include "../lib/debug_config.h" #include "../lib/mem.h" #include #include // 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; }