initial commit
This commit is contained in:
@@ -0,0 +1,512 @@
|
||||
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);
|
||||
}
|
||||
Reference in New Issue
Block a user