const std = @import("std"); const posix = std.posix; const Allocator = std.mem.Allocator; const packet = @import("../dns/packet.zig"); const types = @import("../dns/types.zig"); /// UDP query task for worker pool const QueryTask = struct { data: [types.EDNS_DEFAULT_SIZE]u8, len: usize, src_addr: posix.sockaddr, addr_len: posix.socklen_t, }; /// Simple bounded work queue for UDP query tasks const WorkQueue = struct { items: [QUEUE_SIZE]?QueryTask, head: usize, tail: usize, count: usize, mutex: std.Thread.Mutex, not_empty: std.Thread.Condition, const QUEUE_SIZE = 64; fn init() WorkQueue { return .{ .items = [_]?QueryTask{null} ** QUEUE_SIZE, .head = 0, .tail = 0, .count = 0, .mutex = .{}, .not_empty = .{}, }; } fn tryPush(self: *WorkQueue, task: QueryTask) bool { self.mutex.lock(); defer self.mutex.unlock(); if (self.count >= QUEUE_SIZE) { return false; } self.items[self.tail] = task; self.tail = (self.tail + 1) % QUEUE_SIZE; self.count += 1; self.not_empty.signal(); return true; } fn pop(self: *WorkQueue, running: *std.atomic.Value(bool)) ?QueryTask { self.mutex.lock(); defer self.mutex.unlock(); while (self.count == 0) { if (!running.load(.acquire)) { return null; } self.not_empty.timedWait(&self.mutex, 100 * std.time.ns_per_ms) catch {}; } if (self.count == 0) return null; const task = self.items[self.head]; self.items[self.head] = null; self.head = (self.head + 1) % QUEUE_SIZE; self.count -= 1; return task; } fn wakeAll(self: *WorkQueue) void { self.mutex.lock(); defer self.mutex.unlock(); self.not_empty.broadcast(); } }; pub const UdpServer = struct { socket: posix.socket_t, allocator: Allocator, handler: *Handler, running: std.atomic.Value(bool), work_queue: WorkQueue, workers: []std.Thread, num_workers: u32, dropped_queries: std.atomic.Value(u64), /// Default number of worker threads pub const DEFAULT_NUM_WORKERS: u32 = 8; pub const Handler = struct { context: *anyopaque, handleFn: *const fn (*anyopaque, []const u8, std.net.Address, Allocator) ?[]const u8, pub fn handle(self: Handler, query: []const u8, client_addr: std.net.Address, allocator: Allocator) ?[]const u8 { return self.handleFn(self.context, query, client_addr, allocator); } }; pub const InitError = error{ SocketCreationFailed, SetSockOptFailed, BindFailed, } || posix.SocketError || posix.SetSockOptError; pub const Config = struct { num_workers: u32 = DEFAULT_NUM_WORKERS, }; /// Initialize the UDP server pub fn init(bind_addr: std.net.Address, handler: *Handler, allocator: Allocator) InitError!UdpServer { return initWithConfig(bind_addr, handler, allocator, .{}); } /// Initialize the UDP server with custom configuration pub fn initWithConfig(bind_addr: std.net.Address, handler: *Handler, allocator: Allocator, config: Config) InitError!UdpServer { // Create UDP socket const socket = try posix.socket( bind_addr.any.family, posix.SOCK.DGRAM, 0, ); errdefer posix.close(socket); // Bind to address (no SO_REUSEADDR - we want bind to fail if another instance is running) posix.bind(socket, &bind_addr.any, bind_addr.getOsSockLen()) catch { return error.BindFailed; }; return UdpServer{ .socket = socket, .allocator = allocator, .handler = handler, .running = std.atomic.Value(bool).init(false), .work_queue = WorkQueue.init(), .workers = &[_]std.Thread{}, .num_workers = config.num_workers, .dropped_queries = std.atomic.Value(u64).init(0), }; } /// Get the count of dropped queries (for monitoring) pub fn getDroppedQueries(self: *UdpServer) u64 { return self.dropped_queries.load(.monotonic); } /// Start the server loop pub fn run(self: *UdpServer) !void { self.running.store(true, .release); // Start worker threads self.workers = self.allocator.alloc(std.Thread, self.num_workers) catch |err| { std.log.err("UDP: failed to allocate worker threads: {}", .{err}); return error.OutOfMemory; }; errdefer self.allocator.free(self.workers); var started: u32 = 0; errdefer { self.running.store(false, .release); self.work_queue.wakeAll(); for (self.workers[0..started]) |w| w.join(); } for (self.workers) |*worker| { worker.* = std.Thread.spawn(.{}, workerLoop, .{self}) catch |err| { std.log.err("UDP: failed to start worker thread: {}", .{err}); return error.ThreadSpawnFailed; }; started += 1; } std.log.info("UDP: started {} worker threads", .{self.num_workers}); var buffer: [types.EDNS_DEFAULT_SIZE]u8 = undefined; while (self.running.load(.acquire)) { // Use poll with timeout to allow checking running flag var fds = [1]posix.pollfd{ .{ .fd = self.socket, .events = posix.POLL.IN, .revents = 0, }, }; const poll_result = posix.poll(&fds, 100) catch |err| { std.log.warn("UDP poll error: {}", .{err}); continue; }; // Timeout - check running flag and continue if (poll_result == 0) continue; // No data available if (fds[0].revents & posix.POLL.IN == 0) continue; var src_addr: posix.sockaddr = undefined; var addr_len: posix.socklen_t = @sizeOf(posix.sockaddr); // Receive query const recv_len = posix.recvfrom( self.socket, &buffer, 0, &src_addr, &addr_len, ) catch |err| { std.log.warn("UDP receive error: {}", .{err}); continue; }; if (recv_len < types.DNS_HEADER_SIZE) { continue; // Too small to be valid DNS } // Create task and submit to worker pool var task = QueryTask{ .data = undefined, .len = recv_len, .src_addr = src_addr, .addr_len = addr_len, }; @memcpy(task.data[0..recv_len], buffer[0..recv_len]); if (!self.work_queue.tryPush(task)) { // Backpressure: send SERVFAIL instead of silent drop _ = self.dropped_queries.fetchAdd(1, .monotonic); self.sendServfail(task.data[0..task.len], &task.src_addr, task.addr_len); } } // Shutdown: wait for workers self.work_queue.wakeAll(); for (self.workers) |w| w.join(); self.allocator.free(self.workers); self.workers = &[_]std.Thread{}; } /// Worker thread loop fn workerLoop(self: *UdpServer) void { while (self.running.load(.acquire)) { if (self.work_queue.pop(&self.running)) |task| { self.processQuery(task); } } } /// Process a single query fn processQuery(self: *UdpServer, task: QueryTask) void { const client_addr = std.net.Address{ .any = task.src_addr }; // Handle the query const response = self.handler.handle( task.data[0..task.len], client_addr, self.allocator, ) orelse return; defer self.allocator.free(response); // Determine max response size based on query EDNS support const max_response_size = getMaxResponseSize(task.data[0..task.len]); // Send response (truncate if needed) if (response.len > max_response_size) { // Set TC (truncation) bit in response header var truncated_response: [types.EDNS_DEFAULT_SIZE]u8 = undefined; const safe_max = @min(max_response_size, types.EDNS_DEFAULT_SIZE); const truncated_len = @min(response.len, safe_max); @memcpy(truncated_response[0..truncated_len], response[0..truncated_len]); truncated_response[2] |= 0x02; _ = posix.sendto( self.socket, truncated_response[0..truncated_len], 0, &task.src_addr, task.addr_len, ) catch |err| { std.log.warn("UDP send error: {}", .{err}); }; } else { _ = posix.sendto( self.socket, response, 0, &task.src_addr, task.addr_len, ) catch |err| { std.log.warn("UDP send error: {}", .{err}); }; } } /// Determine maximum response size based on EDNS in query /// Returns 512 (RFC 1035 default) if no EDNS, otherwise client's advertised size fn getMaxResponseSize(query: []const u8) usize { // Need at least header + minimal question if (query.len < types.DNS_HEADER_SIZE) { return types.DNS_UDP_SIZE; } // Check ARCOUNT (additional record count) - bytes 10-11 const arcount = std.mem.readInt(u16, query[10..12], .big); if (arcount == 0) { return types.DNS_UDP_SIZE; } // Quick scan for OPT record (type 41) // OPT records have root name (0x00), type 0x0029 // This is a simplified scan - look for the pattern in additional section var i: usize = types.DNS_HEADER_SIZE; // Skip questions const qdcount = std.mem.readInt(u16, query[4..6], .big); var q: u16 = 0; while (q < qdcount and i < query.len) : (q += 1) { // Skip name while (i < query.len) { const len = query[i]; if (len == 0) { i += 1; break; } else if ((len & 0xC0) == 0xC0) { i += 2; break; } else { i += 1 + len; } } i += 4; // Skip QTYPE and QCLASS } // Skip answers const ancount = std.mem.readInt(u16, query[6..8], .big); var a: u16 = 0; while (a < ancount and i < query.len) : (a += 1) { i = skipResourceRecord(query, i); } // Skip authority const nscount = std.mem.readInt(u16, query[8..10], .big); var n: u16 = 0; while (n < nscount and i < query.len) : (n += 1) { i = skipResourceRecord(query, i); } // Look for OPT in additional var ar: u16 = 0; while (ar < arcount and i + 11 <= query.len) : (ar += 1) { const name_start = i; // Skip name while (i < query.len) { const len = query[i]; if (len == 0) { i += 1; break; } else if ((len & 0xC0) == 0xC0) { i += 2; break; } else { i += 1 + len; } } if (i + 10 > query.len) break; const rtype = std.mem.readInt(u16, query[i..][0..2], .big); if (rtype == 41 and query[name_start] == 0) { // Found OPT record - CLASS field contains UDP payload size const udp_size = std.mem.readInt(u16, query[i + 2 ..][0..2], .big); // Return client's size, capped at our max return @min(udp_size, types.EDNS_DEFAULT_SIZE); } // Skip to next record const rdlength = std.mem.readInt(u16, query[i + 8 ..][0..2], .big); i += 10 + rdlength; } return types.DNS_UDP_SIZE; } /// Skip a resource record and return new position fn skipResourceRecord(data: []const u8, start: usize) usize { var i = start; // Skip name while (i < data.len) { const len = data[i]; if (len == 0) { i += 1; break; } else if ((len & 0xC0) == 0xC0) { i += 2; break; } else { i += 1 + len; } } // Need TYPE(2) + CLASS(2) + TTL(4) + RDLENGTH(2) if (i + 10 > data.len) return data.len; const rdlength = std.mem.readInt(u16, data[i + 8 ..][0..2], .big); const new_pos = i + 10 + rdlength; // Clamp to data.len to ensure callers don't need to handle overflow return @min(new_pos, data.len); } /// Stop the server pub fn stop(self: *UdpServer) void { self.running.store(false, .release); self.work_queue.wakeAll(); } /// Close the server socket pub fn deinit(self: *UdpServer) void { posix.close(self.socket); } /// Send a SERVFAIL response for backpressure fn sendServfail(self: *UdpServer, query: []const u8, addr: *const posix.sockaddr, addr_len: posix.socklen_t) void { if (query.len < types.DNS_HEADER_SIZE) return; // Build minimal SERVFAIL response (12 bytes - header only) var response: [12]u8 = undefined; // Copy transaction ID (bytes 0-1) response[0] = query[0]; response[1] = query[1]; // Flags: QR=1 (response), OPCODE=copy, AA=0, TC=0, RD=copy, RA=1, Z=0, RCODE=2 (SERVFAIL) const opcode = query[2] & 0x78; // Extract OPCODE bits const rd = query[2] & 0x01; // Extract RD bit response[2] = 0x80 | opcode | rd; // QR=1, copy OPCODE and RD response[3] = 0x82; // RA=1, RCODE=2 (SERVFAIL) // Counts: all zeros (no questions/answers in minimal response) response[4] = 0; response[5] = 0; response[6] = 0; response[7] = 0; response[8] = 0; response[9] = 0; response[10] = 0; response[11] = 0; _ = posix.sendto(self.socket, &response, 0, addr, addr_len) catch {}; } }; /// Create a simple echo handler for testing /// Caller must call destroyEchoHandler when done to free allocated context pub fn createEchoHandler(allocator: Allocator) !UdpServer.Handler { const EchoContext = struct { allocator: Allocator, fn handle(ctx: *anyopaque, query: []const u8, _: std.net.Address, alloc: Allocator) ?[]const u8 { _ = ctx; // Parse query and create response var pkt = packet.Packet.parse(query, alloc) catch return null; defer pkt.deinit(); // Create simple response echoing the query var response = packet.Packet.createDeniedResponse(&pkt, alloc) catch return null; defer response.deinit(); var response_buffer: [types.EDNS_DEFAULT_SIZE]u8 = undefined; const response_len = response.encode(&response_buffer) catch return null; return alloc.dupe(u8, response_buffer[0..response_len]) catch return null; } }; const ctx = try allocator.create(EchoContext); ctx.* = EchoContext{ .allocator = allocator }; return UdpServer.Handler{ .context = ctx, .handleFn = EchoContext.handle, }; } /// Free the echo handler context allocated by createEchoHandler pub fn destroyEchoHandler(handler: *UdpServer.Handler, allocator: Allocator) void { const EchoContext = struct { allocator: Allocator }; const ctx: *EchoContext = @ptrCast(@alignCast(handler.context)); allocator.destroy(ctx); handler.context = undefined; } test "UDP server creation" { const testing = std.testing; const allocator = testing.allocator; var handler = try createEchoHandler(allocator); defer destroyEchoHandler(&handler, allocator); // Try to create server on a high port to avoid permission issues const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 15353); var server = UdpServer.init(addr, &handler, allocator) catch |err| { // Skip test if we can't bind (e.g., in CI) std.log.warn("Could not create UDP server: {}", .{err}); return; }; defer server.deinit(); try testing.expect(server.socket != 0); }