414 lines
15 KiB
Zig
414 lines
15 KiB
Zig
const std = @import("std");
|
|
const posix = std.posix;
|
|
const net = std.net;
|
|
const Allocator = std.mem.Allocator;
|
|
const handler = @import("../server/handler.zig");
|
|
const doh = @import("doh.zig");
|
|
const dot = @import("dot.zig");
|
|
const connection_pool = @import("connection_pool.zig");
|
|
|
|
/// Upstream protocol type
|
|
pub const Protocol = enum {
|
|
udp, // Plain UDP (IP:port)
|
|
doh, // DNS-over-HTTPS (https://...)
|
|
dot, // DNS-over-TLS (tls://...)
|
|
};
|
|
|
|
/// Persistent UDP socket for upstream queries
|
|
/// Uses mutex to prevent response mixing between concurrent queries
|
|
const UdpUpstream = struct {
|
|
socket: posix.socket_t,
|
|
mutex: std.Thread.Mutex,
|
|
timeout_ms: u32,
|
|
|
|
fn init(ip: [4]u8, port: u16, timeout_ms: u32) !UdpUpstream {
|
|
const sock = try posix.socket(posix.AF.INET, posix.SOCK.DGRAM, 0);
|
|
errdefer posix.close(sock);
|
|
|
|
const addr = net.Address.initIp4(ip, port);
|
|
|
|
// Connect the socket to the upstream - allows send/recv and ICMP errors
|
|
try posix.connect(sock, &addr.any, addr.getOsSockLen());
|
|
|
|
return UdpUpstream{
|
|
.socket = sock,
|
|
.mutex = .{},
|
|
.timeout_ms = timeout_ms,
|
|
};
|
|
}
|
|
|
|
fn deinit(self: *UdpUpstream) void {
|
|
posix.close(self.socket);
|
|
}
|
|
|
|
fn query(self: *UdpUpstream, dns_packet: []const u8, allocator: Allocator) ![]const u8 {
|
|
self.mutex.lock();
|
|
defer self.mutex.unlock();
|
|
|
|
// Set timeout for this query
|
|
const timeout = posix.timeval{
|
|
.sec = @intCast(self.timeout_ms / 1000),
|
|
.usec = @intCast((self.timeout_ms % 1000) * 1000),
|
|
};
|
|
posix.setsockopt(self.socket, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch {};
|
|
posix.setsockopt(self.socket, posix.SOL.SOCKET, posix.SO.SNDTIMEO, std.mem.asBytes(&timeout)) catch {};
|
|
|
|
// Send using connected socket (no address needed)
|
|
_ = try posix.send(self.socket, dns_packet, 0);
|
|
|
|
// Receive response
|
|
var response_buf: [4096]u8 = undefined;
|
|
const n = try posix.recv(self.socket, &response_buf, 0);
|
|
|
|
if (n < 12) return error.InvalidResponse;
|
|
|
|
const response = try allocator.alloc(u8, n);
|
|
@memcpy(response, response_buf[0..n]);
|
|
return response;
|
|
}
|
|
};
|
|
|
|
/// Upstream DNS server configuration
|
|
pub const UpstreamConfig = struct {
|
|
/// Server address/URL
|
|
/// - "8.8.8.8" or "8.8.8.8:53" for plain UDP
|
|
/// - "https://cloudflare-dns.com/dns-query" for DoH
|
|
/// - "tls://cloudflare-dns.com" for DoT
|
|
address: []const u8,
|
|
port: u16 = 53,
|
|
enabled: bool = true,
|
|
timeout_ms: u32 = 1000, // 1 second default (reduced from 2s for faster failover)
|
|
|
|
/// Detect protocol from address
|
|
pub fn getProtocol(self: UpstreamConfig) Protocol {
|
|
if (std.mem.startsWith(u8, self.address, "https://")) {
|
|
return .doh;
|
|
} else if (std.mem.startsWith(u8, self.address, "tls://")) {
|
|
return .dot;
|
|
} else {
|
|
return .udp;
|
|
}
|
|
}
|
|
};
|
|
|
|
/// Health state for an upstream server
|
|
const UpstreamHealth = struct {
|
|
failures: u32 = 0,
|
|
last_failure: i64 = 0,
|
|
|
|
const FAILURE_THRESHOLD: u32 = 3;
|
|
const COOLDOWN_SECONDS: i64 = 30;
|
|
|
|
fn isHealthy(self: *const UpstreamHealth) bool {
|
|
if (self.failures < FAILURE_THRESHOLD) return true;
|
|
// Allow retry after cooldown
|
|
const now = std.time.timestamp();
|
|
return now - self.last_failure > COOLDOWN_SECONDS;
|
|
}
|
|
|
|
fn recordFailure(self: *UpstreamHealth) void {
|
|
self.failures +|= 1; // Saturating add
|
|
self.last_failure = std.time.timestamp();
|
|
}
|
|
|
|
fn recordSuccess(self: *UpstreamHealth) void {
|
|
self.failures = 0;
|
|
}
|
|
};
|
|
|
|
/// Pool of upstream DNS servers with failover
|
|
pub const UpstreamPool = struct {
|
|
configs: []UpstreamConfig,
|
|
allocator: Allocator,
|
|
/// Connection pool for DoT - maintains persistent TLS connections
|
|
dot_pool: connection_pool.DotConnectionPool,
|
|
/// Connection pool for DoH - maintains persistent HTTP connections
|
|
doh_pool: doh.DohConnectionPool,
|
|
/// Persistent UDP sockets - one per UDP upstream config (null for non-UDP)
|
|
udp_upstreams: []?UdpUpstream,
|
|
/// Health state per upstream for fast failover
|
|
health: []UpstreamHealth,
|
|
|
|
pub fn init(configs: []const UpstreamConfig, allocator: Allocator) !UpstreamPool {
|
|
const configs_copy = try allocator.alloc(UpstreamConfig, configs.len);
|
|
@memcpy(configs_copy, configs);
|
|
|
|
const udp_upstreams = try allocator.alloc(?UdpUpstream, configs.len);
|
|
for (configs_copy, 0..) |config, i| {
|
|
if (config.getProtocol() == .udp) {
|
|
if (parseIpv4(config.address)) |ip| {
|
|
udp_upstreams[i] = UdpUpstream.init(ip, config.port, config.timeout_ms) catch |err| {
|
|
std.log.warn("Failed to create UDP socket for {s}:{d}: {} (will use per-query fallback)", .{ config.address, config.port, err });
|
|
udp_upstreams[i] = null;
|
|
continue;
|
|
};
|
|
} else {
|
|
udp_upstreams[i] = null;
|
|
}
|
|
} else {
|
|
udp_upstreams[i] = null;
|
|
}
|
|
}
|
|
|
|
const health = try allocator.alloc(UpstreamHealth, configs.len);
|
|
for (health) |*h| {
|
|
h.* = UpstreamHealth{};
|
|
}
|
|
|
|
return UpstreamPool{
|
|
.configs = configs_copy,
|
|
.allocator = allocator,
|
|
.dot_pool = connection_pool.DotConnectionPool.init(allocator),
|
|
.doh_pool = doh.DohConnectionPool.init(allocator),
|
|
.udp_upstreams = udp_upstreams,
|
|
.health = health,
|
|
};
|
|
}
|
|
|
|
pub fn deinit(self: *UpstreamPool) void {
|
|
for (self.udp_upstreams) |*maybe_udp| {
|
|
if (maybe_udp.*) |*udp| {
|
|
udp.deinit();
|
|
}
|
|
}
|
|
self.allocator.free(self.udp_upstreams);
|
|
self.allocator.free(self.health);
|
|
self.dot_pool.deinit();
|
|
self.doh_pool.deinit();
|
|
self.allocator.free(self.configs);
|
|
}
|
|
|
|
/// Convert to handler-compatible Upstream interface
|
|
pub fn toHandlerUpstream(self: *UpstreamPool) handler.Upstream {
|
|
return handler.Upstream{
|
|
.context = self,
|
|
.queryFn = queryWrapper,
|
|
};
|
|
}
|
|
|
|
fn queryWrapper(ctx: *anyopaque, dns_packet: []const u8, allocator: Allocator) ?[]const u8 {
|
|
const self: *UpstreamPool = @ptrCast(@alignCast(ctx));
|
|
return self.query(dns_packet, allocator);
|
|
}
|
|
|
|
/// Query upstream DNS servers, trying each until one succeeds.
|
|
/// Skips unhealthy upstreams (>3 consecutive failures) for 30 seconds.
|
|
pub fn query(self: *UpstreamPool, dns_packet: []const u8, allocator: Allocator) ?[]const u8 {
|
|
// First pass: try healthy upstreams only
|
|
for (self.configs, 0..) |config, i| {
|
|
if (!config.enabled) continue;
|
|
if (!self.health[i].isHealthy()) continue;
|
|
|
|
const result = self.queryUpstream(config, i, dns_packet, allocator);
|
|
if (result) |response| {
|
|
self.health[i].recordSuccess();
|
|
return response;
|
|
}
|
|
self.health[i].recordFailure();
|
|
}
|
|
|
|
// Second pass: try unhealthy upstreams as last resort
|
|
for (self.configs, 0..) |config, i| {
|
|
if (!config.enabled) continue;
|
|
if (self.health[i].isHealthy()) continue; // Already tried
|
|
|
|
const result = self.queryUpstream(config, i, dns_packet, allocator);
|
|
if (result) |response| {
|
|
self.health[i].recordSuccess();
|
|
return response;
|
|
}
|
|
self.health[i].recordFailure();
|
|
}
|
|
|
|
std.log.err("All upstream DNS servers failed", .{});
|
|
return null;
|
|
}
|
|
|
|
/// Query a single upstream based on its protocol
|
|
fn queryUpstream(self: *UpstreamPool, config: UpstreamConfig, index: usize, dns_packet: []const u8, allocator: Allocator) ?[]const u8 {
|
|
const protocol = config.getProtocol();
|
|
|
|
switch (protocol) {
|
|
.doh => {
|
|
// Use connection pool with slot-based sharding for parallelism
|
|
return self.doh_pool.query(config.address, dns_packet) catch |err| {
|
|
std.log.warn("DoH query failed for {s}: {}", .{ config.address, err });
|
|
return null;
|
|
};
|
|
},
|
|
.dot => {
|
|
// Use persistent connection pool for DoT
|
|
const parsed = parseDotUrl(config.address) orelse {
|
|
std.log.warn("DoT invalid URL: {s}", .{config.address});
|
|
return null;
|
|
};
|
|
|
|
const conn = self.dot_pool.getConnection(parsed.host, parsed.port) catch |err| {
|
|
std.log.warn("DoT pool failed for {s}: {}", .{ config.address, err });
|
|
return null;
|
|
};
|
|
conn.setTimeout(config.timeout_ms);
|
|
|
|
return conn.query(dns_packet) catch |err| {
|
|
std.log.warn("DoT query failed for {s}: {}", .{ config.address, err });
|
|
return null;
|
|
};
|
|
},
|
|
.udp => {
|
|
// Use pooled socket if available, fall back to per-query socket
|
|
if (self.udp_upstreams[index]) |*udp| {
|
|
return udp.query(dns_packet, allocator) catch |err| {
|
|
std.log.warn("UDP upstream {s}:{d} failed: {}", .{ config.address, config.port, err });
|
|
return null;
|
|
};
|
|
} else if (parseIpv4(config.address)) |ip| {
|
|
// Fallback: per-query socket (socket creation failed at init)
|
|
return queryUdp(ip, config.port, dns_packet, config.timeout_ms, allocator) catch |err| {
|
|
std.log.warn("UDP upstream {s}:{d} failed: {}", .{ config.address, config.port, err });
|
|
return null;
|
|
};
|
|
} else {
|
|
std.log.warn("Invalid IPv4 address: {s}", .{config.address});
|
|
return null;
|
|
}
|
|
},
|
|
}
|
|
}
|
|
};
|
|
|
|
/// Parse a tls:// URL into host and port
|
|
fn parseDotUrl(url: []const u8) ?struct { host: []const u8, port: u16 } {
|
|
const prefix = "tls://";
|
|
if (!std.mem.startsWith(u8, url, prefix)) {
|
|
return null;
|
|
}
|
|
|
|
const host_port = url[prefix.len..];
|
|
|
|
// Check for port
|
|
if (std.mem.lastIndexOfScalar(u8, host_port, ':')) |colon_idx| {
|
|
const host = host_port[0..colon_idx];
|
|
const port_str = host_port[colon_idx + 1 ..];
|
|
const port = std.fmt.parseInt(u16, port_str, 10) catch return null;
|
|
return .{ .host = host, .port = port };
|
|
}
|
|
|
|
return .{ .host = host_port, .port = 853 };
|
|
}
|
|
|
|
/// Query a DNS server using plain UDP
|
|
fn queryUdp(ip: [4]u8, port: u16, dns_packet: []const u8, timeout_ms: u32, allocator: Allocator) ![]const u8 {
|
|
const sock = try posix.socket(posix.AF.INET, posix.SOCK.DGRAM, 0);
|
|
defer posix.close(sock);
|
|
|
|
// Set receive timeout - this is critical to avoid indefinite blocking
|
|
const timeout = posix.timeval{
|
|
.sec = @intCast(timeout_ms / 1000),
|
|
.usec = @intCast((timeout_ms % 1000) * 1000),
|
|
};
|
|
posix.setsockopt(sock, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch |err| {
|
|
std.log.warn("Failed to set socket receive timeout (queries may hang): {}", .{err});
|
|
// Continue anyway - the query might still work, just without timeout protection
|
|
};
|
|
|
|
// Also set send timeout
|
|
posix.setsockopt(sock, posix.SOL.SOCKET, posix.SO.SNDTIMEO, std.mem.asBytes(&timeout)) catch |err| {
|
|
std.log.warn("Failed to set socket send timeout: {}", .{err});
|
|
};
|
|
|
|
const addr = net.Address.initIp4(ip, port);
|
|
_ = try posix.sendto(sock, dns_packet, 0, &addr.any, addr.getOsSockLen());
|
|
|
|
var response_buf: [4096]u8 = undefined;
|
|
const n = try posix.recvfrom(sock, &response_buf, 0, null, null);
|
|
|
|
// DNS header is 12 bytes minimum
|
|
if (n < 12) return error.InvalidResponse;
|
|
|
|
const response = try allocator.alloc(u8, n);
|
|
@memcpy(response, response_buf[0..n]);
|
|
return response;
|
|
}
|
|
|
|
/// Parse an IPv4 address string like "8.8.8.8" into bytes
|
|
fn parseIpv4(addr: []const u8) ?[4]u8 {
|
|
// Strip any port suffix
|
|
const host = if (std.mem.indexOf(u8, addr, ":")) |idx| addr[0..idx] else addr;
|
|
|
|
var result: [4]u8 = undefined;
|
|
var parts = std.mem.splitScalar(u8, host, '.');
|
|
var i: usize = 0;
|
|
|
|
while (parts.next()) |part| {
|
|
if (i >= 4) return null;
|
|
result[i] = std.fmt.parseInt(u8, part, 10) catch return null;
|
|
i += 1;
|
|
}
|
|
|
|
if (i != 4) return null;
|
|
return result;
|
|
}
|
|
|
|
test "parseIpv4" {
|
|
const testing = std.testing;
|
|
|
|
try testing.expectEqual([4]u8{ 8, 8, 8, 8 }, parseIpv4("8.8.8.8").?);
|
|
try testing.expectEqual([4]u8{ 1, 1, 1, 1 }, parseIpv4("1.1.1.1").?);
|
|
try testing.expectEqual([4]u8{ 192, 168, 1, 1 }, parseIpv4("192.168.1.1:53").?);
|
|
try testing.expect(parseIpv4("invalid") == null);
|
|
try testing.expect(parseIpv4("256.0.0.1") == null);
|
|
}
|
|
|
|
test "UpstreamConfig protocol detection" {
|
|
const testing = std.testing;
|
|
|
|
const doh_config = UpstreamConfig{ .address = "https://cloudflare-dns.com/dns-query" };
|
|
try testing.expectEqual(Protocol.doh, doh_config.getProtocol());
|
|
|
|
const dot_config = UpstreamConfig{ .address = "tls://cloudflare-dns.com" };
|
|
try testing.expectEqual(Protocol.dot, dot_config.getProtocol());
|
|
|
|
const udp_config = UpstreamConfig{ .address = "8.8.8.8" };
|
|
try testing.expectEqual(Protocol.udp, udp_config.getProtocol());
|
|
|
|
const udp_with_port = UpstreamConfig{ .address = "1.1.1.1:53" };
|
|
try testing.expectEqual(Protocol.udp, udp_with_port.getProtocol());
|
|
}
|
|
|
|
test "UpstreamPool initialization" {
|
|
const testing = std.testing;
|
|
const allocator = testing.allocator;
|
|
|
|
const configs = [_]UpstreamConfig{
|
|
.{ .address = "https://cloudflare-dns.com/dns-query" },
|
|
.{ .address = "tls://1.1.1.1" },
|
|
.{ .address = "8.8.8.8" },
|
|
};
|
|
|
|
var pool = try UpstreamPool.init(&configs, allocator);
|
|
defer pool.deinit();
|
|
|
|
try testing.expectEqual(@as(usize, 3), pool.configs.len);
|
|
}
|
|
|
|
test "parseDotUrl" {
|
|
const testing = std.testing;
|
|
|
|
// Valid URLs
|
|
const r1 = parseDotUrl("tls://cloudflare-dns.com").?;
|
|
try testing.expectEqualStrings("cloudflare-dns.com", r1.host);
|
|
try testing.expectEqual(@as(u16, 853), r1.port);
|
|
|
|
const r2 = parseDotUrl("tls://1.1.1.1:853").?;
|
|
try testing.expectEqualStrings("1.1.1.1", r2.host);
|
|
try testing.expectEqual(@as(u16, 853), r2.port);
|
|
|
|
const r3 = parseDotUrl("tls://dns.google:8853").?;
|
|
try testing.expectEqualStrings("dns.google", r3.host);
|
|
try testing.expectEqual(@as(u16, 8853), r3.port);
|
|
|
|
// Invalid URLs
|
|
try testing.expect(parseDotUrl("https://example.com") == null);
|
|
try testing.expect(parseDotUrl("example.com") == null);
|
|
}
|