From d8d9ddfc538195aa0d51478e9a409841d0ac8626 Mon Sep 17 00:00:00 2001 From: m5r Date: Fri, 26 Dec 2025 18:42:04 +0100 Subject: [PATCH] initial commit --- .gitignore | 5 + AGENTS.md | 122 +++ CLAUDE.md | 1 + PLAN.md | 1645 ++++++++++++++++++++++++++++++ build.zig | 74 ++ build.zig.zon | 12 + config.example.toml | 38 + src/cache/dns_cache.zig | 349 +++++++ src/config/config.zig | 506 +++++++++ src/config/toml.zig | 294 ++++++ src/config/watcher.zig | 230 +++++ src/dns/edns.zig | 242 +++++ src/dns/header.zig | 266 +++++ src/dns/name.zig | 540 ++++++++++ src/dns/packet.zig | 518 ++++++++++ src/dns/question.zig | 214 ++++ src/dns/record.zig | 983 ++++++++++++++++++ src/dns/types.zig | 223 ++++ src/events.zig | 84 ++ src/filter/denylist.zig | 666 ++++++++++++ src/filter/fetcher.zig | 549 ++++++++++ src/filter/safe_search.zig | 142 +++ src/logging/logger.zig | 197 ++++ src/main.zig | 533 ++++++++++ src/memtest.zig | 86 ++ src/server/handler.zig | 655 ++++++++++++ src/server/rate_limiter.zig | 241 +++++ src/server/shutdown.zig | 117 +++ src/server/tcp.zig | 373 +++++++ src/server/udp.zig | 512 ++++++++++ src/storage/db.zig | 305 ++++++ src/storage/logger.zig | 388 +++++++ src/storage/schema.zig | 539 ++++++++++ src/upstream/connection_pool.zig | 340 ++++++ src/upstream/doh.zig | 234 +++++ src/upstream/dot.zig | 189 ++++ src/upstream/pool.zig | 413 ++++++++ src/util.zig | 24 + src/web/api/clients.zig | 85 ++ src/web/api/denylists.zig | 229 +++++ src/web/api/groups.zig | 183 ++++ src/web/api/queries.zig | 47 + src/web/api/rules.zig | 180 ++++ src/web/api/settings.zig | 71 ++ src/web/api/stats.zig | 49 + src/web/auth.zig | 141 +++ src/web/index.html | 537 ++++++++++ src/web/json.zig | 179 ++++ src/web/response.zig | 117 +++ src/web/server.zig | 813 +++++++++++++++ tests/dns/protocol_tests.zig | 466 +++++++++ tests/integration_tests.zig | 917 +++++++++++++++++ 52 files changed, 16863 insertions(+) create mode 100644 .gitignore create mode 100644 AGENTS.md create mode 120000 CLAUDE.md create mode 100644 PLAN.md create mode 100644 build.zig create mode 100644 build.zig.zon create mode 100644 config.example.toml create mode 100644 src/cache/dns_cache.zig create mode 100644 src/config/config.zig create mode 100644 src/config/toml.zig create mode 100644 src/config/watcher.zig create mode 100644 src/dns/edns.zig create mode 100644 src/dns/header.zig create mode 100644 src/dns/name.zig create mode 100644 src/dns/packet.zig create mode 100644 src/dns/question.zig create mode 100644 src/dns/record.zig create mode 100644 src/dns/types.zig create mode 100644 src/events.zig create mode 100644 src/filter/denylist.zig create mode 100644 src/filter/fetcher.zig create mode 100644 src/filter/safe_search.zig create mode 100644 src/logging/logger.zig create mode 100644 src/main.zig create mode 100644 src/memtest.zig create mode 100644 src/server/handler.zig create mode 100644 src/server/rate_limiter.zig create mode 100644 src/server/shutdown.zig create mode 100644 src/server/tcp.zig create mode 100644 src/server/udp.zig create mode 100644 src/storage/db.zig create mode 100644 src/storage/logger.zig create mode 100644 src/storage/schema.zig create mode 100644 src/upstream/connection_pool.zig create mode 100644 src/upstream/doh.zig create mode 100644 src/upstream/dot.zig create mode 100644 src/upstream/pool.zig create mode 100644 src/util.zig create mode 100644 src/web/api/clients.zig create mode 100644 src/web/api/denylists.zig create mode 100644 src/web/api/groups.zig create mode 100644 src/web/api/queries.zig create mode 100644 src/web/api/rules.zig create mode 100644 src/web/api/settings.zig create mode 100644 src/web/api/stats.zig create mode 100644 src/web/auth.zig create mode 100644 src/web/index.html create mode 100644 src/web/json.zig create mode 100644 src/web/response.zig create mode 100644 src/web/server.zig create mode 100644 tests/dns/protocol_tests.zig create mode 100644 tests/integration_tests.zig diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..3a39e9a --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +.zig-cache/ +zig-out/ +*.o + +.ignore/ diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..4934c76 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,122 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## Project Overview + +nxdns is a DNS sinkhole written in Zig. It denies ads/trackers/malware by returning 0.0.0.0 for denied domains and forwards allowed queries to upstream DoH/DoT servers. + +## Build Commands + +```bash +zig build # Build the project +zig build run # Run nxdns +zig build test # Run all tests +zig fmt src/ # Format code +``` + +Validate config without running: +```bash +NXDNS_CONFIG=./config.example.toml ./zig-out/bin/nxdns check +``` + +## Architecture + +``` +src/ +├── main.zig # Entry point, CLI, server orchestration +├── events.zig # Event signaling (denylist reload via eventfd) +├── util.zig # Utility functions +├── dns/ # DNS protocol (RFC 1035) +│ ├── packet.zig # Full packet parse/encode +│ ├── header.zig # 12-byte DNS header +│ ├── name.zig # Domain name with compression +│ ├── question.zig # Query section +│ ├── record.zig # Resource records (A, AAAA, CNAME, etc.) +│ ├── edns.zig # EDNS (OPT record) support +│ └── types.zig # Enums: QType, QClass, RCode +├── server/ +│ ├── udp.zig # UDP listener (port 53) +│ ├── tcp.zig # TCP listener (length-prefixed) +│ ├── handler.zig # Core logic: deny check → cache → upstream +│ ├── shutdown.zig # Graceful shutdown coordination +│ └── rate_limiter.zig # Per-client rate limiting +├── upstream/ +│ ├── doh.zig # DNS-over-HTTPS client +│ ├── dot.zig # DNS-over-TLS client +│ ├── pool.zig # Failover across upstreams +│ └── connection_pool.zig # Connection pooling +├── filter/ +│ ├── denylist.zig # HashMap with parent-walking lookup +│ ├── fetcher.zig # Download and parse denylist URLs +│ └── safe_search.zig # Rewrite domains to force safe search +├── cache/ +│ └── dns_cache.zig # TTL-based response cache +├── storage/ +│ ├── db.zig # SQLite wrapper +│ ├── schema.zig # Tables and migrations +│ └── logger.zig # Async batched query logging +├── config/ +│ ├── config.zig # TOML config loading +│ ├── toml.zig # TOML parser +│ └── watcher.zig # inotify + signalfd for hot reload +├── logging/ +│ └── logger.zig # Structured logging +└── web/ + ├── server.zig # HTTP server + ├── response.zig # HTTP response helpers + ├── json.zig # JSON serialization + ├── auth.zig # Session-based authentication + └── api/ # REST API handlers + ├── stats.zig + ├── queries.zig + ├── clients.zig + ├── groups.zig + ├── denylists.zig + ├── rules.zig + └── settings.zig +``` + +## Key Data Flow + +1. Query arrives (UDP/TCP) → `handler.handle()` +2. Check denylist (parent-walking: `ads.google.com` → `google.com` → `com`) +3. If denied: return 0.0.0.0 response +4. If allowed: check cache → query upstream (DoH/DoT) → cache response +5. Log query to SQLite asynchronously + +## Code Patterns + +**Database transactions** - Use RAII pattern for safety: +```zig +var tx = try db.begin(); +defer tx.deinit(); // Auto-rollback if not committed +// ... do work ... +try tx.commit(); +``` + +## Design Principles + +- **Lean**: No bloat, only what's needed +- **Pragmatic**: Real-world patterns that work +- **Simple**: Easy to understand and extend + +## Implementation Status + +See PLAN.md for the full spec. Current state: +- ✅ DNS protocol parsing/encoding (with EDNS) +- ✅ UDP/TCP servers +- ✅ Denylist with parent-walking +- ✅ DNS cache +- ✅ DoH client +- ✅ DoT client (with connection pooling) +- ✅ Web UI with REST API +- ✅ Config hot reload (inotify + signalfd) +- ✅ Denylist fetching and auto-reload +- ✅ Safe search enforcement +- ✅ Log retention enforcement (startup + daily cleanup) +- ✅ Integration tests + +## Reference + +Zig 0.15.2 stdlib source is at `.ignore/zig/` for API verification. diff --git a/CLAUDE.md b/CLAUDE.md new file mode 120000 index 0000000..55bf822 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +./AGENTS.md \ No newline at end of file diff --git a/PLAN.md b/PLAN.md new file mode 100644 index 0000000..44791fd --- /dev/null +++ b/PLAN.md @@ -0,0 +1,1645 @@ +# nxdns — Complete Implementation Plan + +## Instructions for Claude Code + +You are implementing **nxdns**, a self-hosted DNS sinkhole written in Zig. This document contains everything you need to build it from scratch. Work autonomously, implementing one component at a time, testing as you go. + +**Autonomy guidelines:** +- If something is ambiguous or underspecified, make a reasonable decision and document it +- If you genuinely need human input (architectural decision, external service credentials, etc.), create a file called `QUESTIONS.md` in the project root with your questions, then continue working on other components +- Run tests frequently — after each module is complete +- Commit logically (one feature/fix per commit) with descriptive messages +- If a dependency or approach doesn't work, try alternatives before asking + +**Do not stop until the project is complete and all tests pass.** + +--- + +## Project Overview + +**nxdns** is a DNS server that: +- Listens on UDP/TCP port 53 +- Blocks ads, trackers, and malware by returning `0.0.0.0` for blocked domains +- Forwards allowed queries to upstream DNS-over-HTTPS (DoH) or DNS-over-TLS (DoT) servers +- Provides a web UI for configuration and monitoring +- Stores query logs in SQLite +- Uses TOML for configuration + +Target platform: Linux (primarily Raspberry Pi 5, but any Linux works) + +--- + +## Technology Stack + +| Component | Technology | +|-----------|------------| +| Language | Zig (latest stable) | +| Database | SQLite3 | +| Config | TOML | +| Frontend | React + React Router (static build) + Tremor (charts) | +| HTTP Server | Zig (custom or std.http.Server) | +| DNS | Custom implementation (this is the core of the project) | + +--- + +## Project Structure + +``` +nxdns/ +├── src/ +│ ├── main.zig # Entry point, CLI, server orchestration +│ ├── events.zig # Event signaling (denylist reload via eventfd) +│ ├── util.zig # Utility functions +│ ├── dns/ +│ │ ├── packet.zig # Full DNS packet parsing/encoding +│ │ ├── header.zig # 12-byte DNS header +│ │ ├── name.zig # Domain name with compression +│ │ ├── question.zig # Query section +│ │ ├── record.zig # Resource records (A, AAAA, CNAME, etc.) +│ │ ├── edns.zig # EDNS (OPT record) support +│ │ └── types.zig # Enums: QType, QClass, RCode +│ ├── server/ +│ │ ├── udp.zig # UDP listener (port 53) +│ │ ├── tcp.zig # TCP listener (length-prefixed) +│ │ ├── handler.zig # Core logic: deny check → cache → upstream +│ │ ├── shutdown.zig # Graceful shutdown coordination +│ │ └── rate_limiter.zig # Per-client rate limiting +│ ├── upstream/ +│ │ ├── doh.zig # DNS-over-HTTPS client +│ │ ├── dot.zig # DNS-over-TLS client +│ │ ├── pool.zig # Failover across upstreams with health tracking +│ │ └── connection_pool.zig # Persistent connection pooling for DoT/DoH +│ ├── filter/ +│ │ ├── denylist.zig # HashMap with parent-walking lookup +│ │ ├── fetcher.zig # Download and parse denylist URLs +│ │ └── safe_search.zig # Rewrite domains to force safe search +│ ├── cache/ +│ │ └── dns_cache.zig # TTL-based response cache with eviction +│ ├── storage/ +│ │ ├── db.zig # SQLite wrapper with RAII Transaction +│ │ ├── schema.zig # Tables and migrations +│ │ └── logger.zig # Async batched query logging +│ ├── config/ +│ │ ├── config.zig # TOML config loading with validation +│ │ ├── toml.zig # TOML parser +│ │ └── watcher.zig # inotify + signalfd for hot reload +│ ├── logging/ +│ │ └── logger.zig # Structured logging +│ └── web/ +│ ├── server.zig # HTTP server with SSE for live logs +│ ├── response.zig # HTTP response helpers +│ ├── json.zig # JSON serialization +│ ├── auth.zig # Session-based authentication +│ └── api/ # REST API handlers +│ ├── stats.zig +│ ├── queries.zig +│ ├── clients.zig +│ ├── groups.zig +│ ├── denylists.zig +│ ├── rules.zig +│ └── settings.zig +├── web/ # React frontend (separate build) +│ └── ... +├── tests/ +│ ├── integration_tests.zig # End-to-end server tests +│ ├── integration/ # Integration test helpers +│ └── dns/ +│ └── protocol_tests.zig # DNS protocol conformance tests +├── build.zig +├── build.zig.zon +├── config.example.toml +├── README.md +└── LICENSE +``` + +**Note:** Unit tests are inline in each module using Zig's `test` blocks. Integration and protocol tests are in the `tests/` directory. + +--- + +## Phase 1: DNS Protocol Implementation + +This is the foundation. Implement RFC 1035 with extensions. + +### 1.1 DNS Types and Constants (`src/dns/types.zig`) + +```zig +// Query/Response types +pub const QType = enum(u16) { + A = 1, + NS = 2, + CNAME = 5, + SOA = 6, + PTR = 12, + MX = 15, + TXT = 16, + AAAA = 28, + SRV = 33, + OPT = 41, // EDNS + DS = 43, // DNSSEC + RRSIG = 46, // DNSSEC + DNSKEY = 48, // DNSSEC + ANY = 255, + _, // Allow unknown types +}; + +pub const QClass = enum(u16) { + IN = 1, // Internet + CH = 3, // Chaos + HS = 4, // Hesiod + ANY = 255, + _, +}; + +pub const RCode = enum(u4) { + NoError = 0, + FormErr = 1, + ServFail = 2, + NXDomain = 3, + NotImp = 4, + Refused = 5, + // ... others as needed +}; + +pub const OpCode = enum(u4) { + Query = 0, + IQuery = 1, // Inverse query (obsolete) + Status = 2, + // ... others as needed +}; +``` + +### 1.2 DNS Header (`src/dns/header.zig`) + +DNS header is exactly 12 bytes: + +``` + 1 1 1 1 1 1 + 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| ID | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +|QR| Opcode |AA|TC|RD|RA| Z | RCODE | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| QDCOUNT | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| ANCOUNT | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| NSCOUNT | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| ARCOUNT | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +``` + +Implement: +- `Header` struct with all fields +- `Header.parse(buffer: []const u8) !Header` +- `Header.encode(self: Header, buffer: []u8) void` + +### 1.3 Domain Name Encoding (`src/dns/name.zig`) + +DNS names use label format: `\x03www\x06google\x03com\x00` + +**Label compression**: A pointer (2 bytes starting with `0xC0`) can reference an earlier name in the packet to save space. + +``` +Pointer format: 11PPPPPP PPPPPPPP +- First 2 bits are 1,1 (0xC0 mask) +- Remaining 14 bits are offset from start of packet +``` + +Implement: +- `Name` struct (store as slice of labels or as string) +- `Name.parse(buffer: []const u8, packet_start: []const u8) !struct { name: Name, bytes_read: usize }` +- `Name.encode(self: Name, buffer: []u8, compression_map: *CompressionMap) !usize` +- `Name.toString(self: Name, allocator: Allocator) ![]const u8` +- Handle compression on both read and write + +### 1.4 Questions (`src/dns/question.zig`) + +``` ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| QNAME | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| QTYPE | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| QCLASS | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +``` + +Implement: +- `Question` struct +- `Question.parse(buffer: []const u8, packet_start: []const u8) !struct { question: Question, bytes_read: usize }` +- `Question.encode(...) !usize` + +### 1.5 Resource Records (`src/dns/record.zig`) + +``` ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| NAME | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| TYPE | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| CLASS | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| TTL | +| | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| RDLENGTH | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +| RDATA | ++--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +``` + +RDATA format depends on TYPE: +- **A**: 4 bytes (IPv4 address) +- **AAAA**: 16 bytes (IPv6 address) +- **CNAME**: Compressed domain name +- **MX**: 2 bytes preference + domain name +- **TXT**: Length-prefixed strings +- **SOA**: Multiple fields (see RFC 1035) +- **NS**: Domain name +- **PTR**: Domain name +- **SRV**: Priority, weight, port, target + +Implement: +- `ResourceRecord` struct with tagged union for RDATA +- Parse and encode methods + +### 1.6 EDNS (`src/dns/edns.zig`) + +EDNS uses OPT pseudo-record in Additional section: +- NAME: 0 (root) +- TYPE: 41 (OPT) +- CLASS: Requestor's UDP payload size +- TTL: Extended RCODE and flags (including DO bit for DNSSEC) +- RDATA: Attribute-value pairs (options) + +Key options: +- DNSSEC OK (DO) bit in flags + +Implement: +- `EdnsOption` struct +- Parse/encode OPT records +- Helper to check if query supports DNSSEC + +### 1.7 Complete Packet (`src/dns/packet.zig`) + +```zig +pub const Packet = struct { + header: Header, + questions: []Question, + answers: []ResourceRecord, + authority: []ResourceRecord, + additional: []ResourceRecord, + + pub fn parse(buffer: []const u8, allocator: Allocator) !Packet { ... } + pub fn encode(self: Packet, buffer: []u8) !usize { ... } + pub fn deinit(self: *Packet, allocator: Allocator) void { ... } +}; +``` + +### 1.8 Testing DNS Implementation + +Create comprehensive tests in `tests/dns/`: + +```zig +// Test with real captured DNS packets +// Use Wireshark to capture real queries/responses +// Store as hex strings and verify parsing + +test "parse simple A query" { + const query = "\x00\x01\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00" ++ + "\x03www\x06google\x03com\x00" ++ + "\x00\x01\x00\x01"; + const packet = try Packet.parse(query, testing.allocator); + defer packet.deinit(testing.allocator); + + try testing.expectEqual(packet.header.id, 1); + try testing.expectEqual(packet.questions.len, 1); + try testing.expectEqualStrings(packet.questions[0].name.toString(), "www.google.com"); +} +``` + +Test cases to cover: +- Simple queries (A, AAAA) +- Responses with multiple answers +- Label compression (pointer in middle of name) +- EDNS queries +- CNAME chains +- Malformed packets (should return errors, not crash) +- Maximum length names (253 chars) +- Maximum label length (63 chars) + +--- + +## Phase 2: DNS Server + +### 2.1 UDP Server (`src/server/udp.zig`) + +```zig +pub const UdpServer = struct { + socket: std.posix.socket_t, + handler: *Handler, + + pub fn init(bind_addr: std.net.Address, handler: *Handler) !UdpServer { ... } + pub fn run(self: *UdpServer) !void { + var buffer: [512]u8 = undefined; // Standard DNS UDP size + while (true) { + const result = try std.posix.recvfrom(self.socket, &buffer, ...); + // Spawn task or handle inline + self.handler.handle(buffer[0..result.len], result.src_addr); + } + } +}; +``` + +Notes: +- UDP DNS messages are typically max 512 bytes (or larger with EDNS) +- Each query is independent, no connection state +- Consider using async or thread pool for high throughput + +### 2.2 TCP Server (`src/server/tcp.zig`) + +TCP DNS uses 2-byte length prefix before each message: + +``` ++--+--+ +|Length| (2 bytes, big-endian) ++--+--+ +| | +| DNS | +|Packet| +| | ++--+--+ +``` + +```zig +pub const TcpServer = struct { + listener: std.net.Server, + handler: *Handler, + + pub fn run(self: *TcpServer) !void { + while (true) { + const conn = try self.listener.accept(); + // Spawn handler for this connection + try self.handleConnection(conn); + } + } + + fn handleConnection(self: *TcpServer, conn: std.net.Server.Connection) !void { + defer conn.stream.close(); + while (true) { + // Read 2-byte length + var len_buf: [2]u8 = undefined; + _ = try conn.stream.readAll(&len_buf); + const length = std.mem.readInt(u16, &len_buf, .big); + + // Read message + var buffer: [65535]u8 = undefined; + _ = try conn.stream.readAll(buffer[0..length]); + + // Handle and respond + const response = try self.handler.handle(buffer[0..length]); + + // Write length-prefixed response + var resp_len: [2]u8 = undefined; + std.mem.writeInt(u16, &resp_len, @intCast(response.len), .big); + try conn.stream.writeAll(&resp_len); + try conn.stream.writeAll(response); + } + } +}; +``` + +### 2.3 Request Handler (`src/server/handler.zig`) + +Core logic flow: + +``` +1. Parse incoming packet +2. Extract query name and type +3. Check denylist (with parent-walking: ads.google.com → google.com → com) +4. If denied: return 0.0.0.0 (or NXDOMAIN based on config) +5. If allowed: + a. Check cache + b. If cache hit: return cached response (with TTL adjustment) + c. If cache miss: query upstream (with health-aware failover) + d. Cache response + e. Return response +6. Log query asynchronously (batched writes) +``` + +```zig +pub const Handler = struct { + blocklist: *Blocklist, + cache: *DnsCache, + upstream: *UpstreamPool, + logger: *QueryLogger, + config: *Config, + + pub fn handle(self: *Handler, query_bytes: []const u8, client_addr: std.net.Address) ![]const u8 { + const query = try Packet.parse(query_bytes, self.allocator); + defer query.deinit(self.allocator); + + if (query.questions.len == 0) return error.NoQuestion; + + const question = query.questions[0]; + const domain = question.name.toString(); + + // Get client's group + const group = self.getClientGroup(client_addr); + + // Check blocklist + if (self.blocklist.isBlocked(domain, group)) { + // Log and return blocked response + self.logger.log(.{ + .domain = domain, + .client = client_addr, + .blocked = true, + }); + return self.createBlockedResponse(query); + } + + // Check cache + if (self.cache.get(domain, question.qtype)) |cached| { + return cached; + } + + // Forward to upstream + const response = try self.upstream.query(query_bytes); + + // CNAME uncloaking: check if any CNAME in response is blocked + const parsed_response = try Packet.parse(response, self.allocator); + for (parsed_response.answers) |answer| { + if (answer.type == .CNAME) { + const cname_target = answer.rdata.cname.toString(); + if (self.blocklist.isBlocked(cname_target, group)) { + self.logger.log(.{ + .domain = domain, + .client = client_addr, + .blocked = true, + .reason = "CNAME uncloaking", + }); + return self.createBlockedResponse(query); + } + } + } + + // Cache and return + self.cache.put(domain, question.qtype, response, parsed_response.answers[0].ttl); + + self.logger.log(.{ + .domain = domain, + .client = client_addr, + .blocked = false, + }); + + return response; + } + + fn createBlockedResponse(self: *Handler, query: Packet) ![]const u8 { + // Create response with same ID, set QR=1, RA=1 + // Add A record with 0.0.0.0 (or return NXDOMAIN based on config) + ... + } +}; +``` + +--- + +## Phase 3: Upstream DNS Clients + +### 3.1 DoH Client (`src/upstream/doh.zig`) + +DNS-over-HTTPS sends DNS wire format via HTTP POST: + +``` +POST /dns-query HTTP/1.1 +Host: cloudflare-dns.com +Content-Type: application/dns-message +Content-Length: + + +``` + +Response is also `application/dns-message` with DNS wire format body. + +```zig +pub const DohClient = struct { + url: []const u8, // e.g., "https://cloudflare-dns.com/dns-query" + http_client: std.http.Client, + + pub fn query(self: *DohClient, dns_packet: []const u8) ![]const u8 { + var request = try self.http_client.request(.POST, self.url, ...); + request.headers.append("Content-Type", "application/dns-message"); + request.headers.append("Accept", "application/dns-message"); + + try request.writer().writeAll(dns_packet); + try request.finish(); + + try request.wait(); + + const response = try request.reader().readAllAlloc(self.allocator, 65535); + return response; + } +}; +``` + +Upstream servers to support: +- `https://cloudflare-dns.com/dns-query` +- `https://dns.google/dns-query` +- `https://dns.quad9.net/dns-query` + +### 3.2 DoT Client (`src/upstream/dot.zig`) + +DNS-over-TLS is regular DNS over TLS on port 853, using same 2-byte length prefix as TCP: + +```zig +pub const DotClient = struct { + host: []const u8, + port: u16 = 853, + + pub fn query(self: *DotClient, dns_packet: []const u8) ![]const u8 { + // Establish TLS connection + var tls_stream = try std.crypto.tls.Client.init( + std.net.tcpConnectToHost(self.host, self.port), + self.host, + ); + defer tls_stream.close(); + + // Send length-prefixed query + var len_buf: [2]u8 = undefined; + std.mem.writeInt(u16, &len_buf, @intCast(dns_packet.len), .big); + try tls_stream.writeAll(&len_buf); + try tls_stream.writeAll(dns_packet); + + // Read length-prefixed response + _ = try tls_stream.readAll(&len_buf); + const resp_len = std.mem.readInt(u16, &len_buf, .big); + + var response = try self.allocator.alloc(u8, resp_len); + _ = try tls_stream.readAll(response); + + return response; + } +}; +``` + +### 3.3 Upstream Pool (`src/upstream/pool.zig`) + +Manages multiple upstreams with priority-order failover: + +```zig +pub const UpstreamPool = struct { + upstreams: []Upstream, // Ordered by priority + + pub fn query(self: *UpstreamPool, dns_packet: []const u8) ![]const u8 { + for (self.upstreams) |upstream| { + const result = upstream.query(dns_packet) catch |err| { + log.warn("Upstream {} failed: {}", .{upstream.url, err}); + continue; // Try next upstream + }; + return result; + } + return error.AllUpstreamsFailed; + } +}; +``` + +--- + +## Phase 4: Filtering + +### 4.1 Denylist HashMap (`src/filter/denylist.zig`) + +```zig +pub const Blocklist = struct { + /// Maps group_id -> set of blocked domains + groups: std.AutoHashMap(u32, std.StringHashMap(void)), + allocator: Allocator, + + pub fn isBlocked(self: *Blocklist, domain: []const u8, group_id: u32) bool { + const group_blocklist = self.groups.get(group_id) orelse return false; + + // Parent-walking: check domain and all parent domains + var d = domain; + while (true) { + if (group_blocklist.contains(d)) return true; + + // Move to parent domain + if (std.mem.indexOfScalar(u8, d, '.')) |idx| { + d = d[idx + 1..]; + } else { + return false; + } + } + } + + pub fn reload(self: *Blocklist, db: *Database) !void { + // Clear existing + // Load from SQLite + // Group by group_id + } +}; +``` + +### 4.2 Denylist Fetcher (`src/filter/fetcher.zig`) + +Load domains from SQLite into HashMap: + +```zig +pub fn loadBlocklist(db: *Database, allocator: Allocator) !Blocklist { + var blocklist = Blocklist.init(allocator); + + // Load all domains grouped by their group associations + const stmt = try db.prepare( + \\SELECT d.domain, gl.group_id + \\FROM domains d + \\JOIN group_lists gl ON d.list_id = gl.list_id + \\WHERE gl.enabled = 1 + ); + + while (try stmt.step()) { + const domain = stmt.getText(0); + const group_id = stmt.getInt(1); + try blocklist.addDomain(domain, group_id); + } + + return blocklist; +} +``` + +### 4.3 Safe Search (`src/filter/safesearch.zig`) + +Rewrite queries to force safe search: + +```zig +const safe_search_rewrites = .{ + // Google + .{ "www.google.com", "forcesafesearch.google.com" }, + .{ "www.google.co.uk", "forcesafesearch.google.com" }, + // ... other Google TLDs + + // Bing + .{ "www.bing.com", "strict.bing.com" }, + + // YouTube + .{ "www.youtube.com", "restrictmoderate.youtube.com" }, + .{ "m.youtube.com", "restrictmoderate.youtube.com" }, + .{ "youtubei.googleapis.com", "restrictmoderate.youtube.com" }, + .{ "youtube.googleapis.com", "restrictmoderate.youtube.com" }, + .{ "www.youtube-nocookie.com", "restrictmoderate.youtube.com" }, + + // DuckDuckGo + .{ "duckduckgo.com", "safe.duckduckgo.com" }, +}; + +pub fn applySafeSearch(domain: []const u8) ?[]const u8 { + inline for (safe_search_rewrites) |rewrite| { + if (std.mem.eql(u8, domain, rewrite[0])) { + return rewrite[1]; + } + } + return null; +} +``` + +### 4.4 CNAME Uncloaking (`src/filter/cname.zig`) + +Check CNAME targets in responses against blocklist: + +```zig +pub fn checkCnameChain(response: *Packet, blocklist: *Blocklist, group_id: u32) bool { + for (response.answers) |answer| { + if (answer.type == .CNAME) { + const target = answer.rdata.cname.toString(); + if (blocklist.isBlocked(target, group_id)) { + return true; // Blocked via CNAME + } + } + } + return false; +} +``` + +--- + +## Phase 5: DNS Cache + +### 5.1 Cache Implementation (`src/cache/dns_cache.zig`) + +```zig +const CacheEntry = struct { + response: []const u8, + expires_at: i64, + qtype: QType, +}; + +pub const DnsCache = struct { + entries: std.StringHashMap(std.ArrayList(CacheEntry)), + allocator: Allocator, + max_entries: usize, + + pub fn get(self: *DnsCache, domain: []const u8, qtype: QType) ?[]const u8 { + const entries = self.entries.get(domain) orelse return null; + const now = std.time.timestamp(); + + for (entries.items) |entry| { + if (entry.qtype == qtype and entry.expires_at > now) { + return entry.response; + } + } + return null; + } + + pub fn put(self: *DnsCache, domain: []const u8, qtype: QType, response: []const u8, ttl: u32) !void { + const expires_at = std.time.timestamp() + ttl; + + // Evict if at capacity + if (self.entries.count() >= self.max_entries) { + self.evictOldest(); + } + + // Store + const entry = CacheEntry{ + .response = try self.allocator.dupe(u8, response), + .expires_at = expires_at, + .qtype = qtype, + }; + + // ... add to entries + } + + fn evictOldest(self: *DnsCache) void { + // LRU or random eviction + } +}; +``` + +--- + +## Phase 6: Storage + +### 6.1 SQLite Wrapper (`src/storage/db.zig`) + +Use Zig's SQLite bindings or call C API directly: + +```zig +pub const Database = struct { + conn: *c.sqlite3, + + pub fn open(path: []const u8) !Database { ... } + pub fn close(self: *Database) void { ... } + pub fn exec(self: *Database, sql: []const u8) !void { ... } + pub fn prepare(self: *Database, sql: []const u8) !Statement { ... } +}; + +pub const Statement = struct { + stmt: *c.sqlite3_stmt, + + pub fn bind(self: *Statement, index: usize, value: anytype) !void { ... } + pub fn step(self: *Statement) !bool { ... } + pub fn getText(self: *Statement, col: usize) []const u8 { ... } + pub fn getInt(self: *Statement, col: usize) i64 { ... } + pub fn reset(self: *Statement) void { ... } +}; +``` + +### 6.2 Schema (`src/storage/schema.zig`) + +```sql +-- String interning for query log +CREATE TABLE domains ( + id INTEGER PRIMARY KEY, + domain TEXT NOT NULL UNIQUE +); + +CREATE TABLE clients ( + id INTEGER PRIMARY KEY, + ip TEXT NOT NULL UNIQUE, + name TEXT, + group_id INTEGER REFERENCES groups(id) +); + +CREATE TABLE groups ( + id INTEGER PRIMARY KEY, + name TEXT NOT NULL UNIQUE +); + +-- Blocklist sources +CREATE TABLE blocklist_sources ( + id INTEGER PRIMARY KEY, + url TEXT NOT NULL, + name TEXT, + category TEXT, -- 'ads', 'malware', 'adult', etc. + enabled INTEGER DEFAULT 1, + last_updated INTEGER, + domain_count INTEGER DEFAULT 0 +); + +-- Domains from blocklists +CREATE TABLE blocklist_domains ( + id INTEGER PRIMARY KEY, + domain TEXT NOT NULL, + source_id INTEGER REFERENCES blocklist_sources(id), + UNIQUE(domain, source_id) +); + +CREATE INDEX idx_blocklist_domains_domain ON blocklist_domains(domain); + +-- Which groups use which blocklist sources +CREATE TABLE group_sources ( + group_id INTEGER REFERENCES groups(id), + source_id INTEGER REFERENCES blocklist_sources(id), + PRIMARY KEY (group_id, source_id) +); + +-- Custom rules (per-group allow/block) +CREATE TABLE rules ( + id INTEGER PRIMARY KEY, + domain TEXT NOT NULL, + group_id INTEGER REFERENCES groups(id), + action TEXT NOT NULL, -- 'allow' or 'block' + created_at INTEGER +); + +-- Query log +CREATE TABLE query_log ( + id INTEGER PRIMARY KEY, + timestamp INTEGER NOT NULL, + domain_id INTEGER NOT NULL REFERENCES domains(id), + client_id INTEGER NOT NULL REFERENCES clients(id), + qtype INTEGER, + blocked INTEGER NOT NULL, + response_time_us INTEGER, + upstream TEXT +); + +CREATE INDEX idx_query_log_timestamp ON query_log(timestamp); +CREATE INDEX idx_query_log_client ON query_log(client_id); + +-- Settings +CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT +); +``` + +### 6.3 Async Query Logger (`src/storage/logger.zig`) + +Batch writes to avoid blocking DNS resolution: + +```zig +pub const QueryLogger = struct { + buffer: std.ArrayList(QueryLogEntry), + db: *Database, + mutex: std.Thread.Mutex, + last_flush: i64, + + const BATCH_SIZE = 100; + const FLUSH_INTERVAL_MS = 100; + + pub fn log(self: *QueryLogger, entry: QueryLogEntry) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + self.buffer.append(entry) catch return; + + const now = std.time.milliTimestamp(); + if (self.buffer.items.len >= BATCH_SIZE or + now - self.last_flush >= FLUSH_INTERVAL_MS) { + self.flushLocked(); + } + } + + fn flushLocked(self: *QueryLogger) void { + if (self.buffer.items.len == 0) return; + + self.db.exec("BEGIN IMMEDIATE") catch return; + + for (self.buffer.items) |entry| { + // Insert domain if new, get ID + // Insert client if new, get ID + // Insert log entry + } + + self.db.exec("COMMIT") catch { + self.db.exec("ROLLBACK") catch {}; + return; + }; + + self.buffer.clearRetainingCapacity(); + self.last_flush = std.time.milliTimestamp(); + } +}; +``` + +--- + +## Phase 7: Configuration + +### 7.1 TOML Parser + +Either use an existing Zig TOML library or implement a simple one. The config is not complex. + +### 7.2 Config Structure (`src/config/config.zig`) + +```zig +pub const Config = struct { + upstream: struct { + servers: []const []const u8, + }, + + blocking: struct { + response: enum { zero, nxdomain } = .zero, + }, + + safe_search: struct { + enabled: bool = true, + }, + + web: struct { + port: u16 = 8080, + bind: []const u8 = "127.0.0.1", + password: ?[]const u8 = null, + }, + + logging: struct { + retention: []const u8 = "30 days", + level: LogLevel = .info, + output: []const u8 = "stderr", + }, + + dns: struct { + port: u16 = 53, + bind: []const u8 = "0.0.0.0", + cache_size: usize = 10000, + }, + + pub fn load(path: []const u8) !Config { ... } + pub fn save(self: Config, path: []const u8) !void { ... } +}; +``` + +### 7.3 Config File Example (`config.example.toml`) + +```toml +[upstream] +servers = [ + "https://cloudflare-dns.com/dns-query", + "https://dns.google/dns-query" +] + +[blocking] +response = "zero" # or "nxdomain" + +[safe_search] +enabled = true + +[web] +port = 8080 +bind = "127.0.0.1" +# password = "your-password-here" # Uncomment to enable auth + +[logging] +retention = "30 days" # e.g., "7 days", "1 week", "3 months", "1 year", "forever" +level = "info" # debug, info, warn, error +output = "stderr" # stderr, syslog, or /path/to/file + +[dns] +port = 53 +bind = "0.0.0.0" +cache_size = 10000 +``` + +### 7.4 File Watcher (`src/config/watcher.zig`) + +Use inotify to watch config file for changes: + +```zig +pub const ConfigWatcher = struct { + inotify_fd: i32, + watch_fd: i32, + config_path: []const u8, + on_change: *const fn() void, + + pub fn init(path: []const u8, callback: *const fn() void) !ConfigWatcher { + const fd = try std.posix.inotify_init1(0); + const wd = try std.posix.inotify_add_watch(fd, path, std.posix.IN.MODIFY); + return .{ + .inotify_fd = fd, + .watch_fd = wd, + .config_path = path, + .on_change = callback, + }; + } + + pub fn poll(self: *ConfigWatcher) !void { + var buf: [4096]u8 = undefined; + const len = try std.posix.read(self.inotify_fd, &buf); + if (len > 0) { + self.on_change(); + } + } +}; +``` + +--- + +## Phase 8: Web Interface + +### 8.1 HTTP Server (`src/web/server.zig`) + +Use Zig's std.http.Server or implement a simple one: + +```zig +pub const WebServer = struct { + server: std.net.Server, + router: Router, + static_files: StaticFiles, + + pub fn run(self: *WebServer) !void { + while (true) { + const conn = try self.server.accept(); + try self.handleRequest(conn); + } + } + + fn handleRequest(self: *WebServer, conn: std.net.Server.Connection) !void { + var buffer: [8192]u8 = undefined; + var server = std.http.Server.init(conn, &buffer); + + const request = try server.receiveHead(); + + // Try API routes first + if (std.mem.startsWith(u8, request.target, "/api/")) { + try self.router.handle(request, &server); + return; + } + + // Serve static files + try self.static_files.serve(request, &server); + } +}; +``` + +### 8.2 API Routes (`src/web/api.zig`) + +```zig +pub fn handleStats(request: *Request, response: *Response) !void { + const stats = getStats(); + try response.json(.{ + .queries_today = stats.queries_today, + .blocked_today = stats.blocked_today, + .percent_blocked = stats.percent_blocked, + .top_blocked = stats.top_blocked, + .top_clients = stats.top_clients, + }); +} + +pub fn handleQueries(request: *Request, response: *Response) !void { + const limit = request.queryParam("limit") orelse "100"; + const offset = request.queryParam("offset") orelse "0"; + const client = request.queryParam("client"); + const blocked = request.queryParam("blocked"); + + const queries = db.getQueries(.{ + .limit = std.fmt.parseInt(usize, limit, 10) catch 100, + .offset = std.fmt.parseInt(usize, offset, 10) catch 0, + .client = client, + .blocked = blocked, + }); + + try response.json(queries); +} + +// ... implement all endpoints from API spec +``` + +### 8.3 Server-Sent Events for Live Logs + +Live query logs use SSE (Server-Sent Events) instead of WebSocket for simpler implementation. +The HTTP server maintains a list of SSE subscribers that receive real-time query notifications +via the QueryLogger's subscriber callback mechanism. Endpoint: `GET /api/queries/live`. + +### 8.4 Static File Server (`src/web/static.zig`) + +Embed frontend files at compile time: + +```zig +const index_html = @embedFile("../../web/dist/index.html"); +const app_js = @embedFile("../../web/dist/assets/app.js"); +const app_css = @embedFile("../../web/dist/assets/app.css"); + +pub const StaticFiles = struct { + pub fn serve(self: *StaticFiles, request: *Request, response: *Response) !void { + const path = if (std.mem.eql(u8, request.target, "/")) + "/index.html" + else + request.target; + + const content = self.getFile(path) orelse { + // SPA fallback: serve index.html for unknown routes + response.status = .ok; + response.headers.append("Content-Type", "text/html"); + try response.send(index_html); + return; + }; + + response.headers.append("Content-Type", self.getMimeType(path)); + try response.send(content); + } +}; +``` + +### 8.5 Authentication (`src/web/auth.zig`) + +Simple session-based auth: + +```zig +pub const Auth = struct { + password_hash: ?[]const u8, + sessions: std.StringHashMap(i64), // token -> expires_at + + pub fn isEnabled(self: *Auth) bool { + return self.password_hash != null; + } + + pub fn login(self: *Auth, password: []const u8) ?[]const u8 { + if (!self.verifyPassword(password)) return null; + + const token = generateToken(); + const expires = std.time.timestamp() + 86400; // 24 hours + self.sessions.put(token, expires) catch return null; + return token; + } + + pub fn validateRequest(self: *Auth, request: *Request) bool { + if (!self.isEnabled()) return true; + + const cookie = request.headers.get("Cookie") orelse return false; + const token = parseCookie(cookie, "session") orelse return false; + + const expires = self.sessions.get(token) orelse return false; + return expires > std.time.timestamp(); + } +}; +``` + +--- + +## Phase 9: Frontend (React) + +### 9.1 Setup + +```bash +cd web +npm create vite@latest . -- --template react-ts +npm install react-router-dom @tremor/react tailwindcss +``` + +### 9.2 Key Components + +**Dashboard.tsx** — Stats cards, charts for queries over time, top blocked domains +**QueryLog.tsx** — Table with filters, pagination, links to live view +**LiveLog.tsx** — WebSocket-connected real-time log +**Clients.tsx** — List clients, assign to groups +**Groups.tsx** — Manage groups +**Denylists.tsx** — Add/remove denylist sources, view by category +**Rules.tsx** — Custom allow/block rules +**Settings.tsx** — Config editor + +### 9.3 Tremor Charts + +```tsx +import { AreaChart, Card, Title } from "@tremor/react"; + +function QueriesChart({ data }) { + return ( + + Queries over time + + + ); +} +``` + +### 9.4 Build for Embedding + +```bash +npm run build +# Output in dist/ will be embedded by Zig build +``` + +--- + +## Phase 10: CLI + +### 10.1 Main Entry Point (`src/main.zig`) + +```zig +pub fn main() !void { + var args = std.process.args(); + _ = args.next(); // Skip program name + + const command = args.next() orelse { + printUsage(); + return; + }; + + if (std.mem.eql(u8, command, "run")) { + try runServer(); + } else if (std.mem.eql(u8, command, "check")) { + try checkConfig(); + } else if (std.mem.eql(u8, command, "migrate")) { + try runMigrations(); + } else if (std.mem.eql(u8, command, "upgrade")) { + try selfUpgrade(); + } else if (std.mem.eql(u8, command, "version")) { + printVersion(); + } else { + printUsage(); + } +} + +fn runServer() !void { + const config = try Config.load("/etc/nxdns/config.toml"); + + // Initialize components + var db = try Database.open("/etc/nxdns/nxdns.db"); + try db.migrate(); + + var blocklist = try Blocklist.load(&db); + var cache = DnsCache.init(config.dns.cache_size); + var upstream = try UpstreamPool.init(config.upstream.servers); + var logger = QueryLogger.init(&db); + + var handler = Handler{ + .blocklist = &blocklist, + .cache = &cache, + .upstream = &upstream, + .logger = &logger, + .config = &config, + }; + + // Start servers + var udp = try UdpServer.init(config.dns.port, &handler); + var tcp = try TcpServer.init(config.dns.port, &handler); + var web = try WebServer.init(config.web.port, &handler, &db); + + // Start config watcher + var watcher = try ConfigWatcher.init("/etc/nxdns/config.toml", reload); + + log.info("nxdns started", .{}); + log.info("DNS server listening on :{}", .{config.dns.port}); + log.info("Web UI at http://localhost:{}", .{config.web.port}); + + // Run (spawn threads or use async) + const threads = [_]std.Thread{ + try std.Thread.spawn(.{}, UdpServer.run, .{&udp}), + try std.Thread.spawn(.{}, TcpServer.run, .{&tcp}), + try std.Thread.spawn(.{}, WebServer.run, .{&web}), + try std.Thread.spawn(.{}, ConfigWatcher.poll, .{&watcher}), + }; + + for (threads) |t| t.join(); +} + +fn selfUpgrade() !void { + // 1. Fetch latest release from GitHub API + // 2. Download binary for current architecture + // 3. Verify checksum + // 4. Replace current binary + // 5. Print instructions to restart +} +``` + +--- + +## Phase 11: Logging + +### 11.1 Logger (`src/util/log.zig`) + +```zig +pub const LogLevel = enum { debug, info, warn, err }; + +pub const Logger = struct { + level: LogLevel, + output: Output, + + const Output = union(enum) { + stderr, + syslog, + file: std.fs.File, + }; + + pub fn log(self: *Logger, level: LogLevel, comptime fmt: []const u8, args: anytype) void { + if (@intFromEnum(level) < @intFromEnum(self.level)) return; + + const timestamp = formatTimestamp(); + const level_str = switch (level) { + .debug => "DEBUG", + .info => "INFO ", + .warn => "WARN ", + .err => "ERROR", + }; + + switch (self.output) { + .stderr => { + std.io.getStdErr().writer().print( + "{s} {s} " ++ fmt ++ "\n", + .{timestamp, level_str} ++ args, + ) catch {}; + }, + .syslog => { + // Use syslog(3) + }, + .file => |f| { + f.writer().print(...) catch {}; + }, + } + } + + pub fn info(self: *Logger, comptime fmt: []const u8, args: anytype) void { + self.log(.info, fmt, args); + } + + // ... debug, warn, err +}; +``` + +--- + +## Phase 12: Build System + +### 12.1 build.zig + +```zig +const std = @import("std"); + +pub fn build(b: *std.Build) void { + const target = b.standardTargetOptions(.{}); + const optimize = b.standardOptimizeOption(.{}); + + const exe = b.addExecutable(.{ + .name = "nxdns", + .root_source_file = b.path("src/main.zig"), + .target = target, + .optimize = optimize, + }); + + // Link SQLite + exe.linkSystemLibrary("sqlite3"); + exe.linkLibC(); + + // Embed static files + exe.addAnonymousModule("static", .{ + .root_source_file = b.path("web/dist/embed.zig"), + }); + + b.installArtifact(exe); + + // Run command + const run_cmd = b.addRunArtifact(exe); + run_cmd.step.dependOn(b.getInstallStep()); + + const run_step = b.step("run", "Run nxdns"); + run_step.dependOn(&run_cmd.step); + + // Tests + const unit_tests = b.addTest(.{ + .root_source_file = b.path("src/main.zig"), + .target = target, + .optimize = optimize, + }); + + const run_unit_tests = b.addRunArtifact(unit_tests); + const test_step = b.step("test", "Run unit tests"); + test_step.dependOn(&run_unit_tests.step); +} +``` + +--- + +## Implementation Order + +1. **DNS packet parsing** (types, header, name, question, record) +2. **DNS packet encoding** (reverse of above) +3. **Tests for parsing/encoding** (use captured real packets) +4. **UDP server** (receive query, echo back) +5. **Basic handler** (parse, create response, send) +6. **TCP server** +7. **DoH client** +8. **DoT client** +9. **Upstream pool** with failover +10. **SQLite wrapper** +11. **Database schema** +12. **Denylist loading** from SQLite +13. **Denylist HashMap** with parent-walking +14. **Deny logic** in handler +15. **CNAME uncloaking** +16. **DNS cache** +17. **Safe search** +18. **Query logging** (async batched) +19. **TOML config** parser +20. **Config loading/saving** +21. **Config file watcher** +22. **EDNS support** +23. **HTTP server** +24. **REST API** endpoints +25. **Static file serving** +26. **Authentication** +27. **WebSocket** for live logs +28. **React frontend** +29. **CLI commands** (check, migrate, upgrade, version) +30. **Logging** (stderr, syslog, file) +31. **Graceful shutdown** +32. **Integration tests** +33. **Documentation** + +--- + +## Testing Strategy + +### Unit Tests +- Every DNS parsing function +- Denylist lookups (exact match, parent walking) +- Cache operations +- Config parsing + +### Integration Tests +- Full query flow (UDP → handler → upstream → response) +- Blocking behavior +- CNAME uncloaking +- API endpoints + +### Manual Testing +- Use `dig` to test DNS resolution +- Compare responses with real DNS servers +- Test blocking with known ad domains +- Verify safe search works + +### Test Domains +- `google.com` — should resolve +- `doubleclick.net` — should be blocked (if in denylist) +- `tracker.example.com` → CNAME → `blocked-tracker.net` — test uncloaking + +--- + +## Blocklist Sources to Include + +Pre-configured categories with suggested lists: + +### Ads & Trackers +- OISD: `https://big.oisd.nl/domainswild` +- StevenBlack: `https://raw.githubusercontent.com/StevenBlack/hosts/master/hosts` +- AdGuard DNS: `https://adguardteam.github.io/AdGuardSDNSFilter/Filters/filter.txt` + +### Malware & Phishing +- URLhaus: `https://urlhaus.abuse.ch/downloads/hostfile/` +- PhishTank: (requires parsing) +- abuse.ch: `https://threatfox.abuse.ch/downloads/hostfile/` + +### Adult Content +- OISD NSFW: `https://nsfw.oisd.nl/domainswild` + +### Native Telemetry +- NextDNS native tracking: `https://raw.githubusercontent.com/nextdns/native-tracking-domains/main/domains/*` + +### Gambling +- Various community lists + +### Social Media +- Block Facebook, Twitter, TikTok, etc. + +--- + +## Performance Targets + +- Handle 100 queries/second sustained +- < 1ms for denylist lookup +- < 5ms total latency for cached responses +- < 100MB memory with 1M blocked domains +- < 50MB database per month of logs (with retention) + +--- + +## Security Considerations + +- Bind to specific interface (not 0.0.0.0 by default for web) +- Rate limiting to prevent DNS amplification +- Validate all input (DNS packets, API requests) +- SQL injection prevention (use prepared statements) +- Password hashing (argon2 or bcrypt) for web auth + +--- + +## Architecture Patterns + +### RAII Database Transactions +Database transactions use RAII pattern for automatic rollback on scope exit: + +```zig +var tx = try db.begin(); +defer tx.deinit(); // Auto-rollback if not committed + +// ... do work ... + +try tx.commit(); // Only on success +``` + +### Graceful Shutdown with signalfd +Instead of traditional signal handlers (which have async-signal-safety constraints), +use Linux signalfd to handle SIGINT/SIGTERM as regular file descriptor events. +This allows safe cleanup without signal handler restrictions. + +### Upstream Health Tracking +Upstreams track consecutive failures with 30-second cooldown. Healthy upstreams +are tried first, unhealthy ones only as last resort. This provides fast failover +without hammering failing servers. + +### Connection Pooling +DoT and DoH use persistent connection pools to avoid TLS handshake overhead. +UDP uses pre-connected sockets with mutex protection for thread safety. + +### Backpressure via SERVFAIL +When UDP queue is full, send SERVFAIL response instead of silent drop. +This provides feedback to clients for proper retry behavior. + +### Atomic Connection Limits +TCP connection limits use atomic fetch-add before accepting to prevent TOCTOU races: +```zig +const old = count.fetchAdd(1, .acq_rel); +if (old >= max) { + _ = count.fetchSub(1, .release); + // reject +} +``` + +--- + +## Default Configuration + +On first run, create: +- `/etc/nxdns/config.toml` — from example +- `/etc/nxdns/nxdns.db` — empty database with schema +- Default group "default" +- One denylist source (OISD or StevenBlack) + +--- + +## Error Handling + +- DNS parse errors: Return FORMERR +- Upstream timeout: Try next upstream, eventually return SERVFAIL +- SQLite errors: Log and continue (don't crash) +- Config errors: Log and use defaults where possible + +--- + +## Questions File + +If you need clarification on anything, create `QUESTIONS.md`: + +```markdown +# Questions for Human Review + +## Question 1: [Topic] +[Your question here] + +**Context:** [Why you need this answered] + +**Your current assumption:** [What you're doing in the meantime] + +--- + +## Question 2: ... +``` + +Continue working on other parts while waiting for answers. + +--- + +## Success Criteria + +The project is complete when: + +1. `nxdns run` starts successfully +2. DNS queries resolve correctly (test with `dig @localhost google.com`) +3. Blocked domains return 0.0.0.0 (test with `dig @localhost doubleclick.net`) +4. Web UI loads at `http://localhost:8080` +5. Query log shows in web UI +6. All tests pass (`zig build test`) +7. Can add/remove denylists via web UI +8. Safe search works +9. CNAME uncloaking works +10. Upstream failover works + +--- + +Good luck. Build something great. diff --git a/build.zig b/build.zig new file mode 100644 index 0000000..aa4583f --- /dev/null +++ b/build.zig @@ -0,0 +1,74 @@ +const std = @import("std"); + +pub fn build(b: *std.Build) void { + const target = b.standardTargetOptions(.{}); + const optimize = b.standardOptimizeOption(.{}); + + const exe = b.addExecutable(.{ + .name = "nxdns", + .root_module = b.createModule(.{ + .root_source_file = b.path("src/main.zig"), + .target = target, + .optimize = optimize, + }), + }); + + // Link SQLite + exe.linkSystemLibrary("sqlite3"); + exe.linkLibC(); + + b.installArtifact(exe); + + // Run command + const run_cmd = b.addRunArtifact(exe); + run_cmd.step.dependOn(b.getInstallStep()); + + if (b.args) |args| { + run_cmd.addArgs(args); + } + + const run_step = b.step("run", "Run nxdns"); + run_step.dependOn(&run_cmd.step); + + // Unit tests for main (includes all modules via imports - runs all embedded tests) + const main_tests = b.addTest(.{ + .root_module = b.createModule(.{ + .root_source_file = b.path("src/main.zig"), + .target = target, + .optimize = optimize, + }), + }); + main_tests.linkSystemLibrary("sqlite3"); + main_tests.linkLibC(); + + const run_main_tests = b.addRunArtifact(main_tests); + + // Unit tests for DNS packet module (self-contained) + const dns_tests = b.addTest(.{ + .root_module = b.createModule(.{ + .root_source_file = b.path("src/dns/packet.zig"), + .target = target, + .optimize = optimize, + }), + }); + + const run_dns_tests = b.addRunArtifact(dns_tests); + + const test_step = b.step("test", "Run unit tests"); + test_step.dependOn(&run_main_tests.step); + test_step.dependOn(&run_dns_tests.step); + + // Memory test binary + const memtest = b.addExecutable(.{ + .name = "memtest", + .root_module = b.createModule(.{ + .root_source_file = b.path("src/memtest.zig"), + .target = target, + .optimize = .ReleaseFast, + }), + }); + + const run_memtest = b.addRunArtifact(memtest); + const memtest_step = b.step("memtest", "Run memory test with real blocklists"); + memtest_step.dependOn(&run_memtest.step); +} diff --git a/build.zig.zon b/build.zig.zon new file mode 100644 index 0000000..c604301 --- /dev/null +++ b/build.zig.zon @@ -0,0 +1,12 @@ +.{ + .name = .nxdns, + .version = "0.1.0", + .fingerprint = 0x3307b31156a24fc8, + .minimum_zig_version = "0.14.0", + .dependencies = .{}, + .paths = .{ + "build.zig", + "build.zig.zon", + "src", + }, +} diff --git a/config.example.toml b/config.example.toml new file mode 100644 index 0000000..96b4930 --- /dev/null +++ b/config.example.toml @@ -0,0 +1,38 @@ +# nxdns configuration + +# Database path +database = "/etc/nxdns/nxdns.db" + +[upstream] +# Upstream DNS servers (DoH or DoT) +servers = ["https://cloudflare-dns.com/dns-query", "https://dns.google/dns-query"] + +[blocking] +# Response for denied domains: "zero" (0.0.0.0) or "nxdomain" +response = "zero" +# Automatically fetch denylist updates every 5 minutes +auto_update = true + +[safe_search] +# Enforce safe search on Google, Bing, YouTube, etc. +enabled = true + +[web] +# Web interface settings +port = 8080 +bind = "127.0.0.1" +# password = "your-password-here" + +[logging] +# Query log retention (e.g., "7 days", "1 week", "3 months", "1 year", "forever") +retention = "30 days" +# Log level: debug, info, warn, error +level = "info" +# Output: stderr, syslog, or /path/to/file +output = "stderr" + +[dns] +# DNS server settings +port = 53 +bind = "0.0.0.0" +cache_size = 10000 diff --git a/src/cache/dns_cache.zig b/src/cache/dns_cache.zig new file mode 100644 index 0000000..5c2937e --- /dev/null +++ b/src/cache/dns_cache.zig @@ -0,0 +1,349 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const types = @import("../dns/types.zig"); +const handler = @import("../server/handler.zig"); + +/// Cache entry with TTL and query type +const CacheEntry = struct { + response: []const u8, + qtype: u16, + expires_at: i64, + created_at: i64, +}; + +/// DNS response cache with TTL +pub const DnsCache = struct { + /// Cache entries: domain -> list of entries (for different qtypes) + entries: std.StringHashMapUnmanaged(std.ArrayListUnmanaged(CacheEntry)), + /// Total number of entries + entry_count: usize, + /// Maximum number of entries + max_entries: usize, + /// Minimum TTL (in seconds) + min_ttl: u32, + /// Maximum TTL (in seconds) + max_ttl: u32, + allocator: Allocator, + mutex: std.Thread.Mutex, + + pub const DEFAULT_MAX_ENTRIES: usize = 10000; + pub const DEFAULT_MIN_TTL: u32 = 60; + pub const DEFAULT_MAX_TTL: u32 = 86400; + + pub fn init(allocator: Allocator) DnsCache { + return initWithConfig(allocator, DEFAULT_MAX_ENTRIES, DEFAULT_MIN_TTL, DEFAULT_MAX_TTL); + } + + pub fn initWithConfig(allocator: Allocator, max_entries: usize, min_ttl: u32, max_ttl: u32) DnsCache { + return DnsCache{ + .entries = std.StringHashMapUnmanaged(std.ArrayListUnmanaged(CacheEntry)){}, + .entry_count = 0, + .max_entries = max_entries, + .min_ttl = min_ttl, + .max_ttl = max_ttl, + .allocator = allocator, + .mutex = std.Thread.Mutex{}, + }; + } + + pub fn deinit(self: *DnsCache) void { + var iter = self.entries.iterator(); + while (iter.next()) |entry| { + // Free domain key + self.allocator.free(entry.key_ptr.*); + + // Free all cached responses + for (entry.value_ptr.items) |cache_entry| { + self.allocator.free(cache_entry.response); + } + entry.value_ptr.deinit(self.allocator); + } + self.entries.deinit(self.allocator); + } + + /// Get a cached response and return a mutable copy (caller owns the memory) + /// Filters by both domain and query type + /// Returns error.OutOfMemory if allocation fails, null if not in cache + pub fn getCopy(self: *DnsCache, domain: []const u8, qtype: types.QType) error{OutOfMemory}!?[]u8 { + self.mutex.lock(); + defer self.mutex.unlock(); + + const qtype_val = @intFromEnum(qtype); + const entries_list = self.entries.get(domain) orelse return null; + const now = std.time.timestamp(); + + for (entries_list.items) |entry| { + if (entry.qtype == qtype_val and entry.expires_at > now) { + return try self.allocator.dupe(u8, entry.response); + } + } + + return null; + } + + /// Put a response in the cache + /// Note: Responses with TTL=0 are not cached per RFC 2308 (negative caching) + pub fn put(self: *DnsCache, domain: []const u8, qtype: types.QType, response: []const u8, ttl: u32) void { + // Don't cache responses with TTL=0 - these indicate "do not cache" per RFC + if (ttl == 0) { + return; + } + + self.mutex.lock(); + defer self.mutex.unlock(); + + const qtype_val = @intFromEnum(qtype); + + // Clamp TTL to configured bounds + const clamped_ttl = @max(self.min_ttl, @min(ttl, self.max_ttl)); + + const now = std.time.timestamp(); + const expires_at = now + clamped_ttl; + + // Evict until we have space + while (self.entry_count >= self.max_entries) { + self.evictExpiredLocked(); + if (self.entry_count >= self.max_entries) { + self.evictOldestLocked(); + } + } + + // Get or create entry list for this domain + const gop = self.entries.getOrPut(self.allocator, domain) catch |err| { + std.log.warn("Cache: failed to create entry: {}", .{err}); + return; + }; + + if (!gop.found_existing) { + gop.key_ptr.* = self.allocator.dupe(u8, domain) catch |err| { + std.log.warn("Cache: failed to allocate domain key: {}", .{err}); + return; + }; + gop.value_ptr.* = std.ArrayListUnmanaged(CacheEntry){}; + } + + // Check if we already have an entry for this qtype and update it + for (gop.value_ptr.items) |*entry| { + if (entry.qtype == qtype_val) { + // Update existing entry + self.allocator.free(entry.response); + entry.response = self.allocator.dupe(u8, response) catch |err| { + std.log.warn("Cache: failed to allocate response copy: {}", .{err}); + return; + }; + entry.expires_at = expires_at; + entry.created_at = now; + return; + } + } + + // Copy response and add new entry + const response_copy = self.allocator.dupe(u8, response) catch |err| { + std.log.warn("Cache: failed to allocate response copy: {}", .{err}); + return; + }; + + gop.value_ptr.append(self.allocator, CacheEntry{ + .response = response_copy, + .qtype = qtype_val, + .expires_at = expires_at, + .created_at = now, + }) catch |err| { + std.log.warn("Cache: failed to append entry: {}", .{err}); + self.allocator.free(response_copy); + return; + }; + + self.entry_count += 1; + } + + /// Clear all cached entries + pub fn clear(self: *DnsCache) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + var iter = self.entries.iterator(); + while (iter.next()) |entry| { + self.allocator.free(entry.key_ptr.*); + for (entry.value_ptr.items) |cache_entry| { + self.allocator.free(cache_entry.response); + } + entry.value_ptr.deinit(self.allocator); + } + self.entries.clearRetainingCapacity(); + self.entry_count = 0; + } + + /// Evict expired entries + fn evictExpiredLocked(self: *DnsCache) void { + const now = std.time.timestamp(); + var to_remove = std.ArrayListUnmanaged([]const u8){}; + defer to_remove.deinit(self.allocator); + + var iter = self.entries.iterator(); + while (iter.next()) |entry| { + var i: usize = 0; + while (i < entry.value_ptr.items.len) { + if (entry.value_ptr.items[i].expires_at <= now) { + self.allocator.free(entry.value_ptr.items[i].response); + _ = entry.value_ptr.swapRemove(i); + self.entry_count -= 1; + } else { + i += 1; + } + } + + // Mark empty domains for removal + if (entry.value_ptr.items.len == 0) { + to_remove.append(self.allocator, entry.key_ptr.*) catch continue; + } + } + + // Remove empty domains + for (to_remove.items) |domain| { + if (self.entries.fetchRemove(domain)) |kv| { + self.allocator.free(kv.key); + var value = kv.value; + value.deinit(self.allocator); + } + } + } + + /// Evict the oldest entry + fn evictOldestLocked(self: *DnsCache) void { + var oldest_domain: ?[]const u8 = null; + var oldest_time: i64 = std.math.maxInt(i64); + var oldest_idx: usize = 0; + + var iter = self.entries.iterator(); + while (iter.next()) |entry| { + for (entry.value_ptr.items, 0..) |cache_entry, i| { + if (cache_entry.created_at < oldest_time) { + oldest_time = cache_entry.created_at; + oldest_domain = entry.key_ptr.*; + oldest_idx = i; + } + } + } + + if (oldest_domain) |domain| { + if (self.entries.getPtr(domain)) |list| { + self.allocator.free(list.items[oldest_idx].response); + _ = list.swapRemove(oldest_idx); + self.entry_count -= 1; + + // Remove domain if empty + if (list.items.len == 0) { + if (self.entries.fetchRemove(domain)) |kv| { + self.allocator.free(kv.key); + var value = kv.value; + value.deinit(self.allocator); + } + } + } + } + } + + /// Get cache statistics + pub fn getStats(self: *DnsCache) CacheStats { + self.mutex.lock(); + defer self.mutex.unlock(); + + return CacheStats{ + .entry_count = self.entry_count, + .domain_count = self.entries.count(), + .max_entries = self.max_entries, + }; + } + + /// Convert to handler-compatible Cache interface + pub fn toHandlerCache(self: *DnsCache) handler.Cache { + return handler.Cache{ + .context = self, + .getFn = getWrapper, + .putFn = putWrapper, + }; + } + + fn getWrapper(ctx: *anyopaque, domain: []const u8, qtype: types.QType) ?[]u8 { + const self: *DnsCache = @ptrCast(@alignCast(ctx)); + return self.getCopy(domain, qtype) catch |err| { + std.log.warn("Cache lookup failed due to memory allocation error: {}", .{err}); + return null; + }; + } + + fn putWrapper(ctx: *anyopaque, domain: []const u8, qtype: types.QType, response: []const u8, ttl: u32) void { + const self: *DnsCache = @ptrCast(@alignCast(ctx)); + self.put(domain, qtype, response, ttl); + } +}; + +/// Cache statistics +pub const CacheStats = struct { + entry_count: usize, + domain_count: usize, + max_entries: usize, +}; + +test "DnsCache basic operations" { + const testing = std.testing; + const allocator = testing.allocator; + + var cache = DnsCache.init(allocator); + defer cache.deinit(); + + const response = "test response data"; + cache.put("example.com", types.QType.A, response, 300); + + const cached = try cache.getCopy("example.com", types.QType.A); + defer if (cached) |c| allocator.free(c); + + try testing.expect(cached != null); + try testing.expectEqualStrings(response, cached.?); +} + +test "DnsCache expiration" { + const testing = std.testing; + const allocator = testing.allocator; + + // min_ttl=1 ensures entry doesn't expire immediately + var cache = DnsCache.initWithConfig(allocator, 100, 1, 60); + defer cache.deinit(); + + const response = "test response data"; + cache.put("example.com", types.QType.A, response, 1); // TTL of 1 second (clamped to min) + + // Should be cached immediately (TTL is 1 second, so still valid) + const cached1 = try cache.getCopy("example.com", types.QType.A); + defer if (cached1) |c| allocator.free(c); + try testing.expect(cached1 != null); + + // Note: Can't easily test actual expiration without time manipulation +} + +test "DnsCache miss" { + const testing = std.testing; + const allocator = testing.allocator; + + var cache = DnsCache.init(allocator); + defer cache.deinit(); + + const cached = try cache.getCopy("nonexistent.com", types.QType.A); + try testing.expect(cached == null); +} + +test "DnsCache stats" { + const testing = std.testing; + const allocator = testing.allocator; + + var cache = DnsCache.init(allocator); + defer cache.deinit(); + + cache.put("example1.com", types.QType.A, "response1", 300); + cache.put("example2.com", types.QType.A, "response2", 300); + + const stats = cache.getStats(); + try testing.expectEqual(@as(usize, 2), stats.entry_count); + try testing.expectEqual(@as(usize, 2), stats.domain_count); +} diff --git a/src/config/config.zig b/src/config/config.zig new file mode 100644 index 0000000..452b5ce --- /dev/null +++ b/src/config/config.zig @@ -0,0 +1,506 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const toml = @import("toml.zig"); +const fs = std.fs; + +/// Application configuration +pub const Config = struct { + /// Upstream DNS servers + upstream: UpstreamConfig, + /// Blocking settings + blocking: BlockingConfig, + /// Safe search settings + safe_search: SafeSearchConfig, + /// Web interface settings + web: WebConfig, + /// Logging settings + logging: LoggingConfig, + /// DNS server settings + dns: DnsConfig, + /// Database path + database_path: []const u8, + + allocator: Allocator, + + pub const UpstreamConfig = struct { + servers: [][]const u8, + }; + + pub const BlockingConfig = struct { + response: BlockingResponse, + auto_update: bool, + + pub const BlockingResponse = enum { + zero, + nxdomain, + }; + }; + + pub const SafeSearchConfig = struct { + enabled: bool, + }; + + pub const WebConfig = struct { + port: u16, + bind: []const u8, + password: ?[]const u8, + }; + + pub const LoggingConfig = struct { + retention: []const u8, + level: LogLevel, + output: []const u8, + + pub const LogLevel = enum { + debug, + info, + warn, + err, + }; + + /// Get retention duration in seconds + pub fn retentionSeconds(self: *const LoggingConfig) i64 { + return parseDuration(self.retention) catch 30 * std.time.s_per_day; + } + }; + + /// Parse a human-readable duration string to seconds + /// Supports: "30 days", "1 week", "6 hours", "3 months", "1 year", "forever" + pub fn parseDuration(input: []const u8) !i64 { + const trimmed = std.mem.trim(u8, input, " \t\n\r"); + + if (std.mem.eql(u8, trimmed, "forever")) { + return std.math.maxInt(i64); + } + + // Find where digits end + var num_end: usize = 0; + for (trimmed) |c| { + if (!std.ascii.isDigit(c)) break; + num_end += 1; + } + if (num_end == 0) return error.InvalidDuration; + + const value = std.fmt.parseInt(i64, trimmed[0..num_end], 10) catch return error.InvalidDuration; + if (value < 0) return error.InvalidDuration; + + const unit = std.mem.trim(u8, trimmed[num_end..], " "); + + const multiplier: i64 = if (unit.len == 0) + std.time.s_per_day // bare number = days + else if (startsWith(unit, "sec")) + 1 + else if (startsWith(unit, "min")) + std.time.s_per_min + else if (startsWith(unit, "hour")) + std.time.s_per_hour + else if (startsWith(unit, "day")) + std.time.s_per_day + else if (startsWith(unit, "week")) + std.time.s_per_week + else if (startsWith(unit, "month")) + 30 * std.time.s_per_day + else if (startsWith(unit, "year")) + 365 * std.time.s_per_day + else + return error.InvalidDuration; + + return value * multiplier; + } + + fn startsWith(haystack: []const u8, needle: []const u8) bool { + return std.mem.startsWith(u8, haystack, needle); + } + + pub const DnsConfig = struct { + port: u16, + bind: []const u8, + cache_size: usize, + workers: u32, + }; + + /// Default configuration + pub fn default(allocator: Allocator) !Config { + const servers = try allocator.alloc([]const u8, 2); + servers[0] = try allocator.dupe(u8, "https://cloudflare-dns.com/dns-query"); + servers[1] = try allocator.dupe(u8, "https://dns.google/dns-query"); + + return Config{ + .upstream = .{ + .servers = servers, + }, + .blocking = .{ + .response = .zero, + .auto_update = true, + }, + .safe_search = .{ + .enabled = true, + }, + .web = .{ + .port = 8080, + .bind = try allocator.dupe(u8, "127.0.0.1"), + .password = null, + }, + .logging = .{ + .retention = try allocator.dupe(u8, "30 days"), + .level = .info, + .output = try allocator.dupe(u8, "stderr"), + }, + .dns = .{ + .port = 53, + .bind = try allocator.dupe(u8, "0.0.0.0"), + .cache_size = 10000, + .workers = 32, + }, + .database_path = try allocator.dupe(u8, "/etc/nxdns/nxdns.db"), + .allocator = allocator, + }; + } + + /// Load configuration from a file + pub fn load(path: []const u8, allocator: Allocator) !Config { + // Read file + const file = fs.cwd().openFile(path, .{}) catch |err| { + if (err == error.FileNotFound) { + // Return default config if file doesn't exist + return default(allocator); + } + return err; + }; + defer file.close(); + + const content = try file.readToEndAlloc(allocator, 1024 * 1024); + defer allocator.free(content); + + return parseContent(content, allocator); + } + + /// Parse configuration from a string + pub fn parseContent(content: []const u8, allocator: Allocator) !Config { + var parser = toml.TomlParser.init(allocator); + const toml_data = try parser.parse(content); + defer toml.freeValue(allocator, .{ .table = toml_data }); + + var config = try default(allocator); + errdefer config.deinit(); + + // Parse upstream section + if (toml.getTable(toml_data, "upstream")) |upstream| { + if (toml.getStringArray(upstream, "servers")) |servers_arr| { + // Free default servers + for (config.upstream.servers) |s| { + config.allocator.free(s); + } + config.allocator.free(config.upstream.servers); + + // Parse new servers + var servers = try allocator.alloc([]const u8, servers_arr.len); + for (servers_arr, 0..) |v, i| { + switch (v) { + .string => |s| servers[i] = try allocator.dupe(u8, s), + else => servers[i] = try allocator.dupe(u8, ""), + } + } + config.upstream.servers = servers; + } + } + + // Parse blocking section + if (toml.getTable(toml_data, "blocking")) |blocking| { + if (toml.getString(blocking, "response")) |resp| { + if (std.mem.eql(u8, resp, "nxdomain")) { + config.blocking.response = .nxdomain; + } else { + config.blocking.response = .zero; + } + } + if (toml.getBool(blocking, "auto_update")) |auto_update| { + config.blocking.auto_update = auto_update; + } + } + + // Parse safe_search section + if (toml.getTable(toml_data, "safe_search")) |ss| { + if (toml.getBool(ss, "enabled")) |enabled| { + config.safe_search.enabled = enabled; + } + } + + // Parse web section + if (toml.getTable(toml_data, "web")) |web| { + if (toml.getInt(web, "port")) |port| { + config.web.port = @intCast(port); + } + if (toml.getString(web, "bind")) |bind| { + config.allocator.free(config.web.bind); + config.web.bind = try allocator.dupe(u8, bind); + } + if (toml.getString(web, "password")) |password| { + if (config.web.password) |p| { + config.allocator.free(p); + } + config.web.password = try allocator.dupe(u8, password); + } + } + + // Parse logging section + if (toml.getTable(toml_data, "logging")) |logging| { + if (toml.getString(logging, "retention")) |retention| { + config.allocator.free(config.logging.retention); + config.logging.retention = try allocator.dupe(u8, retention); + } + if (toml.getString(logging, "level")) |level| { + if (std.mem.eql(u8, level, "debug")) { + config.logging.level = .debug; + } else if (std.mem.eql(u8, level, "warn")) { + config.logging.level = .warn; + } else if (std.mem.eql(u8, level, "error")) { + config.logging.level = .err; + } else { + config.logging.level = .info; + } + } + if (toml.getString(logging, "output")) |output| { + config.allocator.free(config.logging.output); + config.logging.output = try allocator.dupe(u8, output); + } + } + + // Parse dns section + if (toml.getTable(toml_data, "dns")) |dns| { + if (toml.getInt(dns, "port")) |port| { + config.dns.port = @intCast(port); + } + if (toml.getString(dns, "bind")) |bind| { + config.allocator.free(config.dns.bind); + config.dns.bind = try allocator.dupe(u8, bind); + } + if (toml.getInt(dns, "cache_size")) |size| { + config.dns.cache_size = @intCast(size); + } + if (toml.getInt(dns, "workers")) |workers| { + config.dns.workers = @intCast(workers); + } + } + + // Parse database path + if (toml.getString(toml_data, "database")) |db_path| { + config.allocator.free(config.database_path); + config.database_path = try allocator.dupe(u8, db_path); + } + + return config; + } + + /// Free configuration memory + pub fn deinit(self: *Config) void { + for (self.upstream.servers) |s| { + self.allocator.free(s); + } + self.allocator.free(self.upstream.servers); + self.allocator.free(self.web.bind); + if (self.web.password) |p| { + self.allocator.free(p); + } + self.allocator.free(self.logging.retention); + self.allocator.free(self.logging.output); + self.allocator.free(self.dns.bind); + self.allocator.free(self.database_path); + } + + /// Validate configuration + pub fn validate(self: *const Config) !void { + // Upstream validation + if (self.upstream.servers.len == 0) { + return error.NoUpstreamServers; + } + for (self.upstream.servers) |server| { + if (!isValidUpstreamServer(server)) { + return error.InvalidUpstreamServer; + } + } + + // DNS validation + if (!isValidPort(self.dns.port)) { + return error.InvalidDnsPort; + } + if (!isValidIpAddress(self.dns.bind)) { + return error.InvalidDnsBindAddress; + } + + // Web validation + if (!isValidPort(self.web.port)) { + return error.InvalidWebPort; + } + if (!isValidIpAddress(self.web.bind)) { + return error.InvalidWebBindAddress; + } + + // Port conflict check (only if binding to same interface) + if (std.mem.eql(u8, self.dns.bind, self.web.bind) and self.dns.port == self.web.port) { + return error.PortConflict; + } + } + + fn isValidPort(port: u16) bool { + return port >= 1 and port <= 65535; + } + + fn isValidIpAddress(addr: []const u8) bool { + // Check for IPv4 + var parts = std.mem.splitScalar(u8, addr, '.'); + var count: usize = 0; + while (parts.next()) |part| { + if (count >= 4) return false; + _ = std.fmt.parseInt(u8, part, 10) catch return false; + count += 1; + } + if (count == 4) return true; + + // Check for common IPv6 formats (::, ::1, etc) + if (std.mem.indexOf(u8, addr, ":") != null) { + return true; // Basic check - let the network layer validate fully + } + + return false; + } + + fn isValidUpstreamServer(server: []const u8) bool { + // Must not be empty + if (server.len == 0) return false; + + // IPv4 address format (with optional port) + if (isValidIpAddress(server) or isValidIpAddressWithPort(server)) { + return true; + } + + // URL format (for future DoH support) + if (std.mem.startsWith(u8, server, "https://") or std.mem.startsWith(u8, server, "tls://")) { + return true; + } + + return false; + } + + fn isValidIpAddressWithPort(addr: []const u8) bool { + const colon_idx = std.mem.lastIndexOf(u8, addr, ":") orelse return false; + if (colon_idx == 0 or colon_idx >= addr.len - 1) return false; + + const ip_part = addr[0..colon_idx]; + const port_part = addr[colon_idx + 1 ..]; + + if (!isValidIpAddress(ip_part)) return false; + const port = std.fmt.parseInt(u16, port_part, 10) catch return false; + return isValidPort(port); + } +}; + +/// Example configuration file content +pub const EXAMPLE_CONFIG = + \\# nxdns configuration + \\ + \\# Database path + \\database = "/etc/nxdns/nxdns.db" + \\ + \\[upstream] + \\# Upstream DNS servers (DoH or DoT) + \\servers = [ + \\ "https://cloudflare-dns.com/dns-query", + \\ "https://dns.google/dns-query" + \\] + \\ + \\[blocking] + \\# Response for blocked domains: "zero" (0.0.0.0) or "nxdomain" + \\response = "zero" + \\# Automatically fetch blocklist updates every 5 minutes + \\auto_update = true + \\ + \\[safe_search] + \\# Enforce safe search on Google, Bing, YouTube, etc. + \\enabled = true + \\ + \\[web] + \\# Web interface settings + \\port = 8080 + \\bind = "127.0.0.1" + \\# password = "your-password-here" # Uncomment to enable auth + \\ + \\[logging] + \\# Query log retention (e.g., "7 days", "1 week", "3 months", "1 year", "forever") + \\retention = "30 days" + \\# Log level: debug, info, warn, error + \\level = "info" + \\# Output: stderr, syslog, or /path/to/file + \\output = "stderr" + \\ + \\[dns] + \\# DNS server settings + \\port = 53 + \\bind = "0.0.0.0" + \\cache_size = 10000 + \\workers = 32 +; + +test "Config default" { + const testing = std.testing; + const allocator = testing.allocator; + + var config = try Config.default(allocator); + defer config.deinit(); + + try testing.expectEqual(@as(u16, 53), config.dns.port); + try testing.expectEqual(@as(u16, 8080), config.web.port); + try testing.expect(config.safe_search.enabled); +} + +test "Config parse" { + const testing = std.testing; + const allocator = testing.allocator; + + const content = + \\[dns] + \\port = 5353 + \\ + \\[web] + \\port = 9090 + \\ + \\[blocking] + \\response = "nxdomain" + ; + + var config = try Config.parseContent(content, allocator); + defer config.deinit(); + + try testing.expectEqual(@as(u16, 5353), config.dns.port); + try testing.expectEqual(@as(u16, 9090), config.web.port); + try testing.expectEqual(Config.BlockingConfig.BlockingResponse.nxdomain, config.blocking.response); +} + +test "parseDuration" { + const testing = std.testing; + + // Basic units + try testing.expectEqual(@as(i64, 60), Config.parseDuration("1 minute")); + try testing.expectEqual(@as(i64, 3600), Config.parseDuration("1 hour")); + try testing.expectEqual(@as(i64, 86400), Config.parseDuration("1 day")); + try testing.expectEqual(@as(i64, 604800), Config.parseDuration("1 week")); + + // Plurals work too + try testing.expectEqual(@as(i64, 7 * 86400), Config.parseDuration("7 days")); + try testing.expectEqual(@as(i64, 2 * 604800), Config.parseDuration("2 weeks")); + + // Months and years (approximate) + try testing.expectEqual(@as(i64, 30 * 86400), Config.parseDuration("1 month")); + try testing.expectEqual(@as(i64, 365 * 86400), Config.parseDuration("1 year")); + + // Bare number = days + try testing.expectEqual(@as(i64, 30 * 86400), Config.parseDuration("30")); + + // Forever + try testing.expectEqual(std.math.maxInt(i64), Config.parseDuration("forever")); + + // Invalid + try testing.expectError(error.InvalidDuration, Config.parseDuration("")); + try testing.expectError(error.InvalidDuration, Config.parseDuration("abc")); + try testing.expectError(error.InvalidDuration, Config.parseDuration("1 parsec")); +} diff --git a/src/config/toml.zig b/src/config/toml.zig new file mode 100644 index 0000000..44e764d --- /dev/null +++ b/src/config/toml.zig @@ -0,0 +1,294 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; + +/// Simple TOML parser for configuration files +pub const TomlParser = struct { + allocator: Allocator, + + pub const Value = union(enum) { + string: []const u8, + integer: i64, + float: f64, + boolean: bool, + array: []Value, + table: std.StringHashMapUnmanaged(Value), + }; + + pub const ParseError = error{ + InvalidSyntax, + UnterminatedString, + InvalidNumber, + InvalidKey, + OutOfMemory, + }; + + pub fn init(allocator: Allocator) TomlParser { + return TomlParser{ + .allocator = allocator, + }; + } + + /// Parse a TOML string + pub fn parse(self: *TomlParser, content: []const u8) ParseError!std.StringHashMapUnmanaged(Value) { + var result = std.StringHashMapUnmanaged(Value){}; + errdefer freeValue(self.allocator, .{ .table = result }); + + var current_table: *std.StringHashMapUnmanaged(Value) = &result; + var lines = std.mem.splitScalar(u8, content, '\n'); + + while (lines.next()) |line| { + const trimmed = std.mem.trim(u8, line, " \t\r"); + + // Skip empty lines and comments + if (trimmed.len == 0 or trimmed[0] == '#') continue; + + // Section header [section.name] + if (trimmed[0] == '[') { + if (std.mem.indexOfScalar(u8, trimmed, ']')) |end| { + const section_name = std.mem.trim(u8, trimmed[1..end], " \t"); + current_table = try self.getOrCreateTable(&result, section_name); + } else { + return error.InvalidSyntax; + } + continue; + } + + // Key-value pair + if (std.mem.indexOfScalar(u8, trimmed, '=')) |eq_pos| { + const key = std.mem.trim(u8, trimmed[0..eq_pos], " \t"); + const value_str = std.mem.trim(u8, trimmed[eq_pos + 1 ..], " \t"); + + if (key.len == 0) return error.InvalidKey; + + const key_copy = try self.allocator.dupe(u8, key); + errdefer self.allocator.free(key_copy); + + const value = try self.parseValue(value_str); + try current_table.put(self.allocator, key_copy, value); + } + } + + return result; + } + + fn getOrCreateTable(self: *TomlParser, root: *std.StringHashMapUnmanaged(Value), path: []const u8) ParseError!*std.StringHashMapUnmanaged(Value) { + var current = root; + var parts = std.mem.splitScalar(u8, path, '.'); + + while (parts.next()) |part| { + const gop = try current.getOrPut(self.allocator, part); + if (!gop.found_existing) { + gop.key_ptr.* = try self.allocator.dupe(u8, part); + gop.value_ptr.* = Value{ .table = std.StringHashMapUnmanaged(Value){} }; + } + switch (gop.value_ptr.*) { + .table => |*t| current = t, + else => return error.InvalidSyntax, + } + } + + return current; + } + + fn parseValue(self: *TomlParser, value_str: []const u8) ParseError!Value { + if (value_str.len == 0) return error.InvalidSyntax; + + // String (quoted) + if (value_str[0] == '"') { + if (value_str.len < 2 or value_str[value_str.len - 1] != '"') { + return error.UnterminatedString; + } + const str = try self.allocator.dupe(u8, value_str[1 .. value_str.len - 1]); + return Value{ .string = str }; + } + + // Boolean + if (std.mem.eql(u8, value_str, "true")) { + return Value{ .boolean = true }; + } + if (std.mem.eql(u8, value_str, "false")) { + return Value{ .boolean = false }; + } + + // Array + if (value_str[0] == '[') { + return try self.parseArray(value_str); + } + + // Number (integer or float) + if (std.mem.indexOfScalar(u8, value_str, '.')) |_| { + const float = std.fmt.parseFloat(f64, value_str) catch return error.InvalidNumber; + return Value{ .float = float }; + } else { + const int = std.fmt.parseInt(i64, value_str, 10) catch return error.InvalidNumber; + return Value{ .integer = int }; + } + } + + fn parseArray(self: *TomlParser, value_str: []const u8) ParseError!Value { + if (value_str.len < 2 or value_str[0] != '[' or value_str[value_str.len - 1] != ']') { + return error.InvalidSyntax; + } + + const content = std.mem.trim(u8, value_str[1 .. value_str.len - 1], " \t"); + if (content.len == 0) { + return Value{ .array = try self.allocator.alloc(Value, 0) }; + } + + var elements = std.ArrayListUnmanaged(Value){}; + errdefer { + for (elements.items) |v| freeValue(self.allocator, v); + elements.deinit(self.allocator); + } + + var iter = std.mem.splitScalar(u8, content, ','); + while (iter.next()) |element| { + const trimmed = std.mem.trim(u8, element, " \t\r\n"); + if (trimmed.len == 0) continue; + const value = try self.parseValue(trimmed); + try elements.append(self.allocator, value); + } + + return Value{ .array = try elements.toOwnedSlice(self.allocator) }; + } +}; + +/// Free a TOML value and all nested values +pub fn freeValue(allocator: Allocator, value: TomlParser.Value) void { + switch (value) { + .string => |s| allocator.free(s), + .array => |arr| { + for (arr) |v| freeValue(allocator, v); + allocator.free(arr); + }, + .table => |t| { + var table = t; + var iter = table.iterator(); + while (iter.next()) |entry| { + allocator.free(entry.key_ptr.*); + freeValue(allocator, entry.value_ptr.*); + } + table.deinit(allocator); + }, + else => {}, + } +} + +/// Get a string value from a table +pub fn getString(table: std.StringHashMapUnmanaged(TomlParser.Value), key: []const u8) ?[]const u8 { + if (table.get(key)) |value| { + return switch (value) { + .string => |s| s, + else => null, + }; + } + return null; +} + +/// Get an integer value from a table +pub fn getInt(table: std.StringHashMapUnmanaged(TomlParser.Value), key: []const u8) ?i64 { + if (table.get(key)) |value| { + return switch (value) { + .integer => |i| i, + else => null, + }; + } + return null; +} + +/// Get a boolean value from a table +pub fn getBool(table: std.StringHashMapUnmanaged(TomlParser.Value), key: []const u8) ?bool { + if (table.get(key)) |value| { + return switch (value) { + .boolean => |b| b, + else => null, + }; + } + return null; +} + +/// Get a nested table +pub fn getTable(table: std.StringHashMapUnmanaged(TomlParser.Value), key: []const u8) ?std.StringHashMapUnmanaged(TomlParser.Value) { + if (table.get(key)) |value| { + return switch (value) { + .table => |t| t, + else => null, + }; + } + return null; +} + +/// Get an array of strings +pub fn getStringArray(table: std.StringHashMapUnmanaged(TomlParser.Value), key: []const u8) ?[]TomlParser.Value { + if (table.get(key)) |value| { + return switch (value) { + .array => |arr| arr, + else => null, + }; + } + return null; +} + +test "TOML parse simple values" { + const testing = std.testing; + const allocator = testing.allocator; + + var parser = TomlParser.init(allocator); + + const content = + \\name = "test" + \\port = 8080 + \\enabled = true + \\ratio = 0.5 + ; + + const result = try parser.parse(content); + defer freeValue(allocator, .{ .table = result }); + + try testing.expectEqualStrings("test", getString(result, "name").?); + try testing.expectEqual(@as(i64, 8080), getInt(result, "port").?); + try testing.expectEqual(true, getBool(result, "enabled").?); +} + +test "TOML parse sections" { + const testing = std.testing; + const allocator = testing.allocator; + + var parser = TomlParser.init(allocator); + + const content = + \\[server] + \\host = "localhost" + \\port = 53 + \\ + \\[web] + \\port = 8080 + ; + + const result = try parser.parse(content); + defer freeValue(allocator, .{ .table = result }); + + const server = getTable(result, "server").?; + try testing.expectEqualStrings("localhost", getString(server, "host").?); + try testing.expectEqual(@as(i64, 53), getInt(server, "port").?); + + const web = getTable(result, "web").?; + try testing.expectEqual(@as(i64, 8080), getInt(web, "port").?); +} + +test "TOML parse array" { + const testing = std.testing; + const allocator = testing.allocator; + + var parser = TomlParser.init(allocator); + + const content = + \\servers = ["server1", "server2", "server3"] + ; + + const result = try parser.parse(content); + defer freeValue(allocator, .{ .table = result }); + + const servers = getStringArray(result, "servers").?; + try testing.expectEqual(@as(usize, 3), servers.len); +} diff --git a/src/config/watcher.zig b/src/config/watcher.zig new file mode 100644 index 0000000..80e8271 --- /dev/null +++ b/src/config/watcher.zig @@ -0,0 +1,230 @@ +const std = @import("std"); +const posix = std.posix; +const Allocator = std.mem.Allocator; +const Config = @import("config.zig").Config; +const events = @import("../events.zig"); +const ShutdownCoordinator = @import("../server/shutdown.zig").ShutdownCoordinator; + +/// Configuration file watcher using Linux inotify +/// Efficiently monitors config file for changes and triggers reload callbacks +/// Uses ShutdownCoordinator's eventfd for clean shutdown +pub const ConfigWatcher = struct { + config_path: []const u8, + allocator: Allocator, + inotify_fd: posix.fd_t, + watch_fd: i32, + coordinator: *ShutdownCoordinator, + on_config_reload: ?*const fn () void, + on_denylist_reload: ?*const fn () void, + + pub const Error = error{ + InotifyInitFailed, + WatchAddFailed, + PathTooLong, + }; + + // inotify init flags + const IN_NONBLOCK: u32 = 0x800; + const IN_CLOEXEC: u32 = 0x80000; + + // inotify watch event masks + const IN_CLOSE_WRITE: u32 = 0x00000008; + const IN_MOVED_TO: u32 = 0x00000080; + const IN_CREATE: u32 = 0x00000100; + + /// Initialize the config watcher with inotify + pub fn init(config_path: []const u8, allocator: Allocator, coordinator: *ShutdownCoordinator) Error!ConfigWatcher { + // Initialize inotify with non-blocking and close-on-exec flags + const inotify_fd = posix.inotify_init1(IN_NONBLOCK | IN_CLOEXEC) catch { + return error.InotifyInitFailed; + }; + errdefer posix.close(inotify_fd); + + // We need to watch the directory, not the file, because editors often + // create a new file and rename it (which deletes the old watch) + var dir_path_buf: [std.fs.max_path_bytes]u8 = undefined; + const dir_path = getDirectoryPath(config_path, &dir_path_buf) catch { + return error.PathTooLong; + }; + + // Add watch for the config file's directory + // Watch for CLOSE_WRITE (file saved), MOVED_TO (file renamed into place), CREATE + const watch_mask = IN_CLOSE_WRITE | IN_MOVED_TO | IN_CREATE; + const watch_fd = posix.inotify_add_watchZ( + inotify_fd, + dir_path, + watch_mask, + ) catch { + return error.WatchAddFailed; + }; + + return ConfigWatcher{ + .config_path = config_path, + .allocator = allocator, + .inotify_fd = inotify_fd, + .watch_fd = watch_fd, + .coordinator = coordinator, + .on_config_reload = null, + .on_denylist_reload = null, + }; + } + + pub fn deinit(self: *ConfigWatcher) void { + if (self.watch_fd >= 0) { + _ = posix.inotify_rm_watch(self.inotify_fd, self.watch_fd); + } + posix.close(self.inotify_fd); + } + + /// Set callback for when config is reloaded + pub fn setConfigReloadCallback(self: *ConfigWatcher, callback: *const fn () void) void { + self.on_config_reload = callback; + } + + /// Set callback for when denylist needs reload + pub fn setDenylistReloadCallback(self: *ConfigWatcher, callback: *const fn () void) void { + self.on_denylist_reload = callback; + } + + /// Start watching (blocking) - returns when shutdown requested + pub fn watch(self: *ConfigWatcher) void { + var event_buf: [4096]u8 align(@alignOf(InotifyEvent)) = undefined; + const target_filename = getFilename(self.config_path); + const denylist_event_fd = events.getDenylistEventFd(); + const shutdown_event_fd = self.coordinator.getEventFd(); + + std.log.info("Config watcher started for: {s}", .{self.config_path}); + + while (!self.coordinator.isShutdownRequested()) { + // Build poll fd array + var fds: [3]posix.pollfd = undefined; + var nfds: usize = 2; + + fds[0] = .{ .fd = self.inotify_fd, .events = posix.POLL.IN, .revents = 0 }; + fds[1] = .{ .fd = shutdown_event_fd, .events = posix.POLL.IN, .revents = 0 }; + + if (denylist_event_fd >= 0) { + fds[2] = .{ .fd = denylist_event_fd, .events = posix.POLL.IN, .revents = 0 }; + nfds = 3; + } + + // Infinite wait - only wakes on real events + const poll_result = posix.poll(fds[0..nfds], -1) catch |err| { + std.log.warn("Config watcher poll error: {}", .{err}); + continue; + }; + + if (poll_result == 0) continue; + + // Check for shutdown signal + if (fds[1].revents & posix.POLL.IN != 0) { + std.log.info("Shutdown requested, stopping watcher...", .{}); + break; + } + + // Check for denylist reload event + if (nfds == 3 and fds[2].revents & posix.POLL.IN != 0) { + if (events.consumeDenylistReload()) { + std.log.info("Denylist update detected, triggering reload...", .{}); + if (self.on_denylist_reload) |callback| { + callback(); + } + } + } + + // Check for config file changes + if (fds[0].revents & posix.POLL.IN != 0) { + // Read events + const bytes_read = posix.read(self.inotify_fd, &event_buf) catch |err| { + if (err == error.WouldBlock) continue; + std.log.warn("Config watcher read error: {}", .{err}); + continue; + }; + + if (bytes_read == 0) continue; + + // Process events + var offset: usize = 0; + while (offset < bytes_read) { + const event: *const InotifyEvent = @ptrCast(@alignCast(&event_buf[offset])); + offset += @sizeOf(InotifyEvent) + event.len; + + // Check if this event is for our config file + if (event.len > 0) { + const event_name = @as([*]const u8, @ptrCast(&event.name))[0 .. event.len - 1]; + if (std.mem.eql(u8, event_name, target_filename)) { + std.log.info("Config file changed, triggering reload...", .{}); + self.triggerConfigReload(); + } + } + } + } + } + + std.log.info("Config watcher stopped", .{}); + } + + /// Start watching in a background thread + pub fn watchAsync(self: *ConfigWatcher) !std.Thread { + return std.Thread.spawn(.{}, watch, .{self}); + } + + fn triggerConfigReload(self: *ConfigWatcher) void { + // Small delay to ensure file is fully written + std.Thread.sleep(100 * std.time.ns_per_ms); + + if (self.on_config_reload) |callback| { + callback(); + } + } + + /// Extract directory path from file path + fn getDirectoryPath(file_path: []const u8, buf: []u8) ![:0]const u8 { + if (std.mem.lastIndexOfScalar(u8, file_path, '/')) |idx| { + if (idx + 1 > buf.len) return error.PathTooLong; + @memcpy(buf[0..idx], file_path[0..idx]); + buf[idx] = 0; + return buf[0..idx :0]; + } + // No directory separator, use current directory + buf[0] = '.'; + buf[1] = 0; + return buf[0..1 :0]; + } + + /// Extract filename from path + fn getFilename(file_path: []const u8) []const u8 { + if (std.mem.lastIndexOfScalar(u8, file_path, '/')) |idx| { + return file_path[idx + 1 ..]; + } + return file_path; + } +}; + +/// Linux inotify_event structure +const InotifyEvent = extern struct { + wd: i32, + mask: u32, + cookie: u32, + len: u32, + name: [0]u8, // Variable length, access via pointer arithmetic +}; + +test "ConfigWatcher getFilename" { + const testing = std.testing; + + try testing.expectEqualStrings("config.toml", ConfigWatcher.getFilename("/etc/nxdns/config.toml")); + try testing.expectEqualStrings("config.toml", ConfigWatcher.getFilename("config.toml")); + try testing.expectEqualStrings("test.toml", ConfigWatcher.getFilename("./test.toml")); +} + +test "ConfigWatcher getDirectoryPath" { + const testing = std.testing; + var buf: [256]u8 = undefined; + + const dir1 = try ConfigWatcher.getDirectoryPath("/etc/nxdns/config.toml", &buf); + try testing.expectEqualStrings("/etc/nxdns", dir1); + + const dir2 = try ConfigWatcher.getDirectoryPath("config.toml", &buf); + try testing.expectEqualStrings(".", dir2); +} diff --git a/src/dns/edns.zig b/src/dns/edns.zig new file mode 100644 index 0000000..4563ae3 --- /dev/null +++ b/src/dns/edns.zig @@ -0,0 +1,242 @@ +const std = @import("std"); +const types = @import("types.zig"); +const Name = @import("name.zig").Name; +const ResourceRecord = @import("record.zig").ResourceRecord; +const RData = @import("record.zig").RData; +const Allocator = std.mem.Allocator; + +/// EDNS Option Codes +pub const OptionCode = enum(u16) { + LLQ = 1, // Long-Lived Queries + UL = 2, // Update Lease + NSID = 3, // Name Server Identifier + DAU = 5, // DNSSEC Algorithm Understood + DHU = 6, // DS Hash Understood + N3U = 7, // NSEC3 Hash Understood + ECS = 8, // Client Subnet + EXPIRE = 9, // Expire + COOKIE = 10, // DNS Cookie + TCP_KEEPALIVE = 11, // TCP Keepalive + PADDING = 12, // Padding + CHAIN = 13, // CHAIN + KEY_TAG = 14, // Key Tag + EXTENDED_ERROR = 15, // Extended DNS Error + CLIENT_TAG = 16, // Client Tag + SERVER_TAG = 17, // Server Tag + _, +}; + +/// EDNS Option +pub const EdnsOption = struct { + code: OptionCode, + data: []const u8, + + pub fn deinit(self: *EdnsOption, allocator: Allocator) void { + allocator.free(self.data); + } +}; + +/// EDNS (OPT Record) Information +pub const Edns = struct { + /// Requestor's UDP payload size + udp_payload_size: u16, + /// Extended RCODE (upper 8 bits) + extended_rcode: u8, + /// EDNS version + version: u8, + /// DNSSEC OK flag + dnssec_ok: bool, + /// Other flags (Z field, 15 bits) + z: u15, + /// EDNS options + options: []EdnsOption, + + allocator: ?Allocator, + + pub fn init() Edns { + return Edns{ + .udp_payload_size = types.EDNS_DEFAULT_SIZE, + .extended_rcode = 0, + .version = 0, + .dnssec_ok = false, + .z = 0, + .options = &[_]EdnsOption{}, + .allocator = null, + }; + } + + /// Parse EDNS from an OPT resource record + pub fn fromOptRecord(record: ResourceRecord, allocator: Allocator) !Edns { + if (record.rtype != types.QType.OPT) { + return error.NotOptRecord; + } + + // CLASS field contains UDP payload size + const udp_payload_size = @intFromEnum(record.class); + + // TTL field contains extended RCODE and flags + const ttl = record.ttl; + const extended_rcode: u8 = @truncate((ttl >> 24) & 0xFF); + const version: u8 = @truncate((ttl >> 16) & 0xFF); + const flags: u16 = @truncate(ttl & 0xFFFF); + const dnssec_ok = (flags >> 15) & 1 == 1; + const z: u15 = @truncate(flags & 0x7FFF); + + // Parse options from RDATA + var options = std.ArrayListUnmanaged(EdnsOption){}; + errdefer { + for (options.items) |*opt| { + opt.deinit(allocator); + } + options.deinit(allocator); + } + + const opt_data = switch (record.rdata) { + .opt => |data| data, + else => return error.InvalidOptRecord, + }; + + var pos: usize = 0; + while (pos + 4 <= opt_data.len) { + const code = std.mem.readInt(u16, opt_data[pos..][0..2], .big); + const length = std.mem.readInt(u16, opt_data[pos + 2 ..][0..2], .big); + pos += 4; + + if (pos + length > opt_data.len) break; + + const data = try allocator.dupe(u8, opt_data[pos .. pos + length]); + try options.append(allocator, EdnsOption{ + .code = @enumFromInt(code), + .data = data, + }); + pos += length; + } + + return Edns{ + .udp_payload_size = udp_payload_size, + .extended_rcode = extended_rcode, + .version = version, + .dnssec_ok = dnssec_ok, + .z = z, + .options = try options.toOwnedSlice(allocator), + .allocator = allocator, + }; + } + + /// Create an OPT resource record from this EDNS info + pub fn toOptRecord(self: Edns, allocator: Allocator) !ResourceRecord { + // Create root name (empty) + var root_name = try Name.fromString("", allocator); + errdefer root_name.deinit(); + + // Encode options to RDATA + var rdata_size: usize = 0; + for (self.options) |opt| { + rdata_size += 4 + opt.data.len; + } + + var rdata = try allocator.alloc(u8, rdata_size); + errdefer allocator.free(rdata); + + var pos: usize = 0; + for (self.options) |opt| { + std.mem.writeInt(u16, rdata[pos..][0..2], @intFromEnum(opt.code), .big); + std.mem.writeInt(u16, rdata[pos + 2 ..][0..2], @intCast(opt.data.len), .big); + @memcpy(rdata[pos + 4 .. pos + 4 + opt.data.len], opt.data); + pos += 4 + opt.data.len; + } + + // Build TTL from extended RCODE and flags + var ttl: u32 = 0; + ttl |= @as(u32, self.extended_rcode) << 24; + ttl |= @as(u32, self.version) << 16; + if (self.dnssec_ok) { + ttl |= @as(u32, 1) << 15; + } + ttl |= @as(u32, self.z); + + return ResourceRecord{ + .name = root_name, + .rtype = types.QType.OPT, + .class = @enumFromInt(self.udp_payload_size), + .ttl = ttl, + .rdata = RData{ .opt = rdata }, + }; + } + + pub fn deinit(self: *Edns) void { + if (self.allocator) |alloc| { + for (self.options) |*opt| { + var o = opt.*; + o.deinit(alloc); + } + alloc.free(self.options); + } + } + + /// Check if DNSSEC validation is requested + pub fn wantsDnssec(self: Edns) bool { + return self.dnssec_ok; + } + + /// Get the maximum UDP payload size the client can handle + pub fn getMaxPayloadSize(self: Edns) u16 { + return self.udp_payload_size; + } +}; + +/// Check if a packet has EDNS support by looking for OPT record in additional section +pub fn hasEdns(records: []const ResourceRecord) bool { + for (records) |record| { + if (record.rtype == types.QType.OPT) { + return true; + } + } + return false; +} + +/// Find and extract EDNS info from additional records +pub fn extractEdns(records: []const ResourceRecord, allocator: Allocator) !?Edns { + for (records) |record| { + if (record.rtype == types.QType.OPT) { + return try Edns.fromOptRecord(record, allocator); + } + } + return null; +} + +test "EDNS init" { + const testing = std.testing; + + const edns = Edns.init(); + + try testing.expectEqual(@as(u16, types.EDNS_DEFAULT_SIZE), edns.udp_payload_size); + try testing.expectEqual(@as(u8, 0), edns.version); + try testing.expect(!edns.dnssec_ok); +} + +test "EDNS to/from OPT record roundtrip" { + const testing = std.testing; + const allocator = testing.allocator; + + var original = Edns{ + .udp_payload_size = 1232, + .extended_rcode = 0, + .version = 0, + .dnssec_ok = true, + .z = 0, + .options = &[_]EdnsOption{}, + .allocator = null, + }; + + var opt_record = try original.toOptRecord(allocator); + defer opt_record.deinit(allocator); + + var parsed = try Edns.fromOptRecord(opt_record, allocator); + defer parsed.deinit(); + + try testing.expectEqual(original.udp_payload_size, parsed.udp_payload_size); + try testing.expectEqual(original.extended_rcode, parsed.extended_rcode); + try testing.expectEqual(original.version, parsed.version); + try testing.expectEqual(original.dnssec_ok, parsed.dnssec_ok); +} diff --git a/src/dns/header.zig b/src/dns/header.zig new file mode 100644 index 0000000..50fab59 --- /dev/null +++ b/src/dns/header.zig @@ -0,0 +1,266 @@ +const std = @import("std"); +const types = @import("types.zig"); + +/// DNS Header (12 bytes) +/// Format: +/// 1 1 1 1 1 1 +/// 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | ID | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// |QR| Opcode |AA|TC|RD|RA| Z | RCODE | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | QDCOUNT | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | ANCOUNT | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | NSCOUNT | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | ARCOUNT | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +pub const Header = struct { + /// Transaction ID + id: u16, + /// Query/Response flag (0 = query, 1 = response) + qr: bool, + /// Operation code + opcode: types.OpCode, + /// Authoritative Answer flag + aa: bool, + /// Truncation flag + tc: bool, + /// Recursion Desired flag + rd: bool, + /// Recursion Available flag + ra: bool, + /// Reserved bit (must be zero) + z: bool, + /// Authenticated Data flag (DNSSEC - response was validated) + ad: bool, + /// Checking Disabled flag (DNSSEC - client doesn't want validation) + cd: bool, + /// Response code + rcode: types.RCode, + /// Number of questions + qdcount: u16, + /// Number of answers + ancount: u16, + /// Number of authority records + nscount: u16, + /// Number of additional records + arcount: u16, + + pub const SIZE: usize = 12; + + pub const ParseError = error{ + BufferTooSmall, + }; + + /// Parse a DNS header from a buffer + pub fn parse(buffer: []const u8) ParseError!Header { + if (buffer.len < SIZE) { + return error.BufferTooSmall; + } + + const id = std.mem.readInt(u16, buffer[0..2], .big); + const flags = std.mem.readInt(u16, buffer[2..4], .big); + const qdcount = std.mem.readInt(u16, buffer[4..6], .big); + const ancount = std.mem.readInt(u16, buffer[6..8], .big); + const nscount = std.mem.readInt(u16, buffer[8..10], .big); + const arcount = std.mem.readInt(u16, buffer[10..12], .big); + + return Header{ + .id = id, + .qr = (flags >> 15) & 1 == 1, + .opcode = @enumFromInt(@as(u4, @truncate((flags >> 11) & 0xF))), + .aa = (flags >> 10) & 1 == 1, + .tc = (flags >> 9) & 1 == 1, + .rd = (flags >> 8) & 1 == 1, + .ra = (flags >> 7) & 1 == 1, + .z = (flags >> 6) & 1 == 1, + .ad = (flags >> 5) & 1 == 1, + .cd = (flags >> 4) & 1 == 1, + .rcode = @enumFromInt(@as(u4, @truncate(flags & 0xF))), + .qdcount = qdcount, + .ancount = ancount, + .nscount = nscount, + .arcount = arcount, + }; + } + + /// Encode the header into a buffer + pub fn encode(self: Header, buffer: []u8) ParseError!void { + if (buffer.len < SIZE) { + return error.BufferTooSmall; + } + + std.mem.writeInt(u16, buffer[0..2], self.id, .big); + + var flags: u16 = 0; + flags |= @as(u16, @intFromBool(self.qr)) << 15; + flags |= @as(u16, @intFromEnum(self.opcode)) << 11; + flags |= @as(u16, @intFromBool(self.aa)) << 10; + flags |= @as(u16, @intFromBool(self.tc)) << 9; + flags |= @as(u16, @intFromBool(self.rd)) << 8; + flags |= @as(u16, @intFromBool(self.ra)) << 7; + flags |= @as(u16, @intFromBool(self.z)) << 6; + flags |= @as(u16, @intFromBool(self.ad)) << 5; + flags |= @as(u16, @intFromBool(self.cd)) << 4; + flags |= @as(u16, @intFromEnum(self.rcode)); + + std.mem.writeInt(u16, buffer[2..4], flags, .big); + std.mem.writeInt(u16, buffer[4..6], self.qdcount, .big); + std.mem.writeInt(u16, buffer[6..8], self.ancount, .big); + std.mem.writeInt(u16, buffer[8..10], self.nscount, .big); + std.mem.writeInt(u16, buffer[10..12], self.arcount, .big); + } + + /// Create a response header from a query header + pub fn createResponse(query: Header, rcode: types.RCode) Header { + return Header{ + .id = query.id, + .qr = true, // Response + .opcode = query.opcode, + .aa = false, + .tc = false, + .rd = query.rd, + .ra = true, // We support recursion + .z = false, + .ad = false, // Set by DNSSEC validation + .cd = query.cd, // Preserve client's CD flag + .rcode = rcode, + .qdcount = query.qdcount, + .ancount = 0, + .nscount = 0, + .arcount = 0, + }; + } + + /// Check if the response has authenticated data (DNSSEC validated) + pub fn isAuthenticated(self: Header) bool { + return self.ad; + } + + /// Check if this is a query + pub fn isQuery(self: Header) bool { + return !self.qr; + } + + /// Check if this is a response + pub fn isResponse(self: Header) bool { + return self.qr; + } +}; + +test "parse simple header" { + const testing = std.testing; + + // Simple query header + const buffer = [_]u8{ + 0x00, 0x01, // ID = 1 + 0x01, 0x00, // Standard query, RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + }; + + const header = try Header.parse(&buffer); + + try testing.expectEqual(@as(u16, 1), header.id); + try testing.expect(!header.qr); // Query + try testing.expectEqual(types.OpCode.Query, header.opcode); + try testing.expect(!header.aa); + try testing.expect(!header.tc); + try testing.expect(header.rd); // Recursion desired + try testing.expect(!header.ra); + try testing.expect(!header.z); + try testing.expect(!header.ad); + try testing.expect(!header.cd); + try testing.expectEqual(types.RCode.NoError, header.rcode); + try testing.expectEqual(@as(u16, 1), header.qdcount); + try testing.expectEqual(@as(u16, 0), header.ancount); + try testing.expectEqual(@as(u16, 0), header.nscount); + try testing.expectEqual(@as(u16, 0), header.arcount); +} + +test "encode and decode header roundtrip" { + const testing = std.testing; + + const original = Header{ + .id = 0xABCD, + .qr = true, + .opcode = types.OpCode.Query, + .aa = true, + .tc = false, + .rd = true, + .ra = true, + .z = false, + .ad = true, + .cd = false, + .rcode = types.RCode.NoError, + .qdcount = 1, + .ancount = 2, + .nscount = 0, + .arcount = 1, + }; + + var buffer: [Header.SIZE]u8 = undefined; + try original.encode(&buffer); + + const decoded = try Header.parse(&buffer); + + try testing.expectEqual(original.id, decoded.id); + try testing.expectEqual(original.qr, decoded.qr); + try testing.expectEqual(original.opcode, decoded.opcode); + try testing.expectEqual(original.aa, decoded.aa); + try testing.expectEqual(original.tc, decoded.tc); + try testing.expectEqual(original.rd, decoded.rd); + try testing.expectEqual(original.ra, decoded.ra); + try testing.expectEqual(original.z, decoded.z); + try testing.expectEqual(original.ad, decoded.ad); + try testing.expectEqual(original.cd, decoded.cd); + try testing.expectEqual(original.rcode, decoded.rcode); + try testing.expectEqual(original.qdcount, decoded.qdcount); + try testing.expectEqual(original.ancount, decoded.ancount); + try testing.expectEqual(original.nscount, decoded.nscount); + try testing.expectEqual(original.arcount, decoded.arcount); +} + +test "create response from query" { + const testing = std.testing; + + const query = Header{ + .id = 0x1234, + .qr = false, + .opcode = types.OpCode.Query, + .aa = false, + .tc = false, + .rd = true, + .ra = false, + .z = false, + .ad = false, + .cd = false, + .rcode = types.RCode.NoError, + .qdcount = 1, + .ancount = 0, + .nscount = 0, + .arcount = 0, + }; + + const response = Header.createResponse(query, types.RCode.NoError); + + try testing.expectEqual(query.id, response.id); + try testing.expect(response.qr); // Is response + try testing.expectEqual(query.opcode, response.opcode); + try testing.expectEqual(query.rd, response.rd); + try testing.expect(response.ra); // Recursion available + try testing.expectEqual(types.RCode.NoError, response.rcode); +} + +test "buffer too small" { + const testing = std.testing; + + const buffer = [_]u8{ 0x00, 0x01 }; // Only 2 bytes + try testing.expectError(error.BufferTooSmall, Header.parse(&buffer)); +} diff --git a/src/dns/name.zig b/src/dns/name.zig new file mode 100644 index 0000000..c150ca1 --- /dev/null +++ b/src/dns/name.zig @@ -0,0 +1,540 @@ +const std = @import("std"); +const types = @import("types.zig"); +const Allocator = std.mem.Allocator; + +/// DNS Name with support for compression +pub const Name = struct { + /// Raw labels (without compression, each label prefixed with length) + labels: []const u8, + /// Allocator used for labels (null if labels are a view into packet) + allocator: ?Allocator, + + pub const ParseError = error{ + InvalidLabel, + NameTooLong, + LabelTooLong, + CompressionLoop, + InvalidPointer, + BufferTooSmall, + OutOfMemory, + }; + + pub const ParseResult = struct { + name: Name, + bytes_read: usize, + }; + + /// Parse a DNS name from a buffer with compression support (iterative, no recursion) + /// packet_start is the beginning of the entire DNS packet (for resolving pointers) + pub fn parse(buffer: []const u8, packet_start: []const u8, allocator: Allocator) ParseError!ParseResult { + // Pre-allocate to avoid multiple reallocations + var labels = try std.ArrayListUnmanaged(u8).initCapacity(allocator, types.MAX_NAME_LENGTH + 1); + errdefer labels.deinit(allocator); + + var bytes_read: usize = 0; + var final_bytes_read: ?usize = null; + var total_length: usize = 0; + var jumps: usize = 0; + const max_jumps: usize = 128; // Prevent infinite compression loops + + // Track visited offsets to detect loops more robustly + // Use a bitset for packet offsets up to 16KB (covers typical DNS packets) + // For offsets beyond 16KB, fall back to jump counter only + var visited: [2048]u8 = [_]u8{0} ** 2048; // 16384 bits = 16KB coverage + + // Current parsing position - can switch buffers on pointer jumps + var current_buffer: []const u8 = buffer; + var pos: usize = 0; + + while (true) { + if (pos >= current_buffer.len) { + return error.BufferTooSmall; + } + + const len_byte = current_buffer[pos]; + + // Check for compression pointer (top 2 bits set) + if ((len_byte & types.COMPRESSION_POINTER_MASK) == types.COMPRESSION_POINTER_MASK) { + if (pos + 1 >= current_buffer.len) { + return error.BufferTooSmall; + } + + // Store the bytes read before first pointer (only from original buffer) + if (final_bytes_read == null) { + final_bytes_read = bytes_read + 2; + } + + // Get pointer offset (14 bits) + const ptr_high: u16 = @as(u16, len_byte & 0x3F) << 8; + const ptr_low: u16 = current_buffer[pos + 1]; + const offset: usize = ptr_high | ptr_low; + + // Validate pointer - must be within packet bounds + if (offset >= packet_start.len) { + return error.InvalidPointer; + } + + // Check for compression loops using visited bitset + if (offset < 16384) { + const byte_idx = offset / 8; + const bit_idx: u3 = @intCast(offset % 8); + if ((visited[byte_idx] & (@as(u8, 1) << bit_idx)) != 0) { + return error.CompressionLoop; + } + visited[byte_idx] |= (@as(u8, 1) << bit_idx); + } + + // Also enforce jump counter as secondary protection + jumps += 1; + if (jumps > max_jumps) { + return error.CompressionLoop; + } + + // Jump to pointed location (iteratively, no recursion) + current_buffer = packet_start[offset..]; + pos = 0; + continue; + } + + // Normal label + const label_len = len_byte; + + if (label_len == 0) { + // End of name + try labels.append(allocator, 0); + // Only add to bytes_read if we haven't jumped + if (final_bytes_read == null) { + bytes_read += 1; + } + break; + } + + if (label_len > types.MAX_LABEL_LENGTH) { + return error.LabelTooLong; + } + + if (pos + 1 + label_len > current_buffer.len) { + return error.BufferTooSmall; + } + + total_length += label_len + 1; // +1 for dot + if (total_length > types.MAX_NAME_LENGTH) { + return error.NameTooLong; + } + + // Copy label length and data + try labels.append(allocator, len_byte); + try labels.appendSlice(allocator, current_buffer[pos + 1 .. pos + 1 + label_len]); + + pos += 1 + label_len; + // Only add to bytes_read if we haven't jumped + if (final_bytes_read == null) { + bytes_read += 1 + label_len; + } + } + + return ParseResult{ + .name = Name{ + .labels = try labels.toOwnedSlice(allocator), + .allocator = allocator, + }, + .bytes_read = final_bytes_read orelse bytes_read, + }; + } + + /// Parse a name from a string (e.g., "www.google.com") + pub fn fromString(str: []const u8, allocator: Allocator) ParseError!Name { + if (str.len > types.MAX_NAME_LENGTH) { + return error.NameTooLong; + } + + var labels = std.ArrayListUnmanaged(u8){}; + errdefer labels.deinit(allocator); + + // Handle root domain + if (str.len == 0 or (str.len == 1 and str[0] == '.')) { + try labels.append(allocator, 0); + return Name{ + .labels = try labels.toOwnedSlice(allocator), + .allocator = allocator, + }; + } + + var iter = std.mem.splitScalar(u8, str, '.'); + while (iter.next()) |label| { + if (label.len == 0) continue; // Skip empty labels (trailing dot) + if (label.len > types.MAX_LABEL_LENGTH) { + return error.LabelTooLong; + } + + try labels.append(allocator, @intCast(label.len)); + try labels.appendSlice(allocator, label); + } + + try labels.append(allocator, 0); // Null terminator + + return Name{ + .labels = try labels.toOwnedSlice(allocator), + .allocator = allocator, + }; + } + + /// Convert the name to a string (e.g., "www.google.com") + pub fn toString(self: Name, allocator: Allocator) ![]u8 { + var result = std.ArrayListUnmanaged(u8){}; + errdefer result.deinit(allocator); + + var pos: usize = 0; + var first = true; + + while (pos < self.labels.len) { + const label_len = self.labels[pos]; + if (label_len == 0) break; + + if (!first) { + try result.append(allocator, '.'); + } + first = false; + + const label_start = pos + 1; + const label_end = label_start + label_len; + if (label_end > self.labels.len) break; + + try result.appendSlice(allocator, self.labels[label_start..label_end]); + pos = label_end; + } + + return result.toOwnedSlice(allocator); + } + + /// Convert the name to a string using a pre-allocated buffer (no allocation) + /// Returns null if buffer is too small + pub fn toStringBuf(self: Name, buf: []u8) ?[]const u8 { + var pos: usize = 0; + var out_pos: usize = 0; + var first = true; + + while (pos < self.labels.len) { + const label_len = self.labels[pos]; + if (label_len == 0) break; + + if (!first) { + if (out_pos >= buf.len) return null; + buf[out_pos] = '.'; + out_pos += 1; + } + first = false; + + const label_start = pos + 1; + const label_end = label_start + label_len; + if (label_end > self.labels.len) break; + + if (out_pos + label_len > buf.len) return null; + @memcpy(buf[out_pos..][0..label_len], self.labels[label_start..label_end]); + out_pos += label_len; + pos = label_end; + } + + return buf[0..out_pos]; + } + + /// Encode the name into a buffer (without compression) + pub fn encode(self: Name, buffer: []u8) ParseError!usize { + if (buffer.len < self.labels.len) { + return error.BufferTooSmall; + } + + @memcpy(buffer[0..self.labels.len], self.labels); + return self.labels.len; + } + + /// Get the encoded length of the name + pub fn encodedLen(self: Name) usize { + return self.labels.len; + } + + /// Check if two names are equal (case-insensitive per DNS spec) + pub fn eql(self: Name, other: Name) bool { + if (self.labels.len != other.labels.len) return false; + + var pos: usize = 0; + while (pos < self.labels.len) { + const len = self.labels[pos]; + if (len != other.labels[pos]) return false; + + if (len == 0) break; + + const start = pos + 1; + const end = start + len; + if (end > self.labels.len) return false; + + for (self.labels[start..end], other.labels[start..end]) |a, b| { + if (std.ascii.toLower(a) != std.ascii.toLower(b)) return false; + } + + pos = end; + } + + return true; + } + + /// Get the number of labels in the name + pub fn labelCount(self: Name) usize { + var count: usize = 0; + var pos: usize = 0; + + while (pos < self.labels.len) { + const len = self.labels[pos]; + if (len == 0) break; + count += 1; + pos += 1 + len; + } + + return count; + } + + /// Check if this name is a subdomain of another + /// E.g., "www.example.com" is a subdomain of "example.com" + pub fn isSubdomainOf(self: Name, parent: Name) bool { + const self_count = self.labelCount(); + const parent_count = parent.labelCount(); + + // Self must have more labels than parent to be a subdomain + if (self_count <= parent_count) return false; + + // Skip the first (self_count - parent_count) labels in self + // Then the remaining labels should match parent exactly + const labels_to_skip = self_count - parent_count; + + var self_pos: usize = 0; + var skipped: usize = 0; + + // Skip labels in self + while (skipped < labels_to_skip and self_pos < self.labels.len) { + const len = self.labels[self_pos]; + if (len == 0) break; + self_pos += 1 + len; + skipped += 1; + } + + // Now compare remaining self labels with all parent labels + var parent_pos: usize = 0; + while (self_pos < self.labels.len and parent_pos < parent.labels.len) { + const self_len = self.labels[self_pos]; + const parent_len = parent.labels[parent_pos]; + + // Both null terminators = match + if (self_len == 0 and parent_len == 0) return true; + + // Length mismatch + if (self_len != parent_len) return false; + + // Compare label content case-insensitively + const self_end = self_pos + 1 + self_len; + const parent_end = parent_pos + 1 + parent_len; + + if (self_end > self.labels.len or parent_end > parent.labels.len) return false; + + for (self.labels[self_pos + 1 .. self_end], parent.labels[parent_pos + 1 .. parent_end]) |a, b| { + if (std.ascii.toLower(a) != std.ascii.toLower(b)) return false; + } + + self_pos = self_end; + parent_pos = parent_end; + } + + return false; + } + + /// Free the name's memory + pub fn deinit(self: *Name) void { + if (self.allocator) |alloc| { + alloc.free(self.labels); + } + } +}; + +/// Compression map for encoding names with compression +pub const CompressionMap = struct { + map: std.StringHashMapUnmanaged(u16), + allocator: Allocator, + + pub fn init(allocator: Allocator) CompressionMap { + return CompressionMap{ + .map = std.StringHashMapUnmanaged(u16){}, + .allocator = allocator, + }; + } + + pub fn deinit(self: *CompressionMap) void { + var iter = self.map.keyIterator(); + while (iter.next()) |key| { + self.allocator.free(key.*); + } + self.map.deinit(self.allocator); + } + + pub fn lookup(self: *CompressionMap, name: []const u8) ?u16 { + return self.map.get(name); + } + + pub fn insert(self: *CompressionMap, name: []const u8, offset: u16) !void { + const key = try self.allocator.dupe(u8, name); + errdefer self.allocator.free(key); + try self.map.put(self.allocator, key, offset); + } +}; + +test "parse simple name" { + const testing = std.testing; + const allocator = testing.allocator; + + // www.google.com encoded as: \x03www\x06google\x03com\x00 + const buffer = [_]u8{ 0x03, 'w', 'o', 'w', 0x06, 'g', 'o', 'o', 'g', 'l', 'e', 0x03, 'c', 'o', 'm', 0x00 }; + + const result = try Name.parse(&buffer, &buffer, allocator); + defer { + var name = result.name; + name.deinit(); + } + + try testing.expectEqual(@as(usize, 16), result.bytes_read); + + const str = try result.name.toString(allocator); + defer allocator.free(str); + + try testing.expectEqualStrings("wow.google.com", str); +} + +test "parse name with compression" { + const testing = std.testing; + const allocator = testing.allocator; + + // Packet with compression: + // Offset 0: header (12 bytes, simulated with zeros) + // Offset 12: \x03www\x06google\x03com\x00 (first name) + // Offset 28: \x03api + pointer to offset 16 (google.com) + var packet: [36]u8 = undefined; + @memset(&packet, 0); + + // First name at offset 12: www.google.com + packet[12] = 0x03; + packet[13] = 'w'; + packet[14] = 'w'; + packet[15] = 'w'; + packet[16] = 0x06; + packet[17] = 'g'; + packet[18] = 'o'; + packet[19] = 'o'; + packet[20] = 'g'; + packet[21] = 'l'; + packet[22] = 'e'; + packet[23] = 0x03; + packet[24] = 'c'; + packet[25] = 'o'; + packet[26] = 'm'; + packet[27] = 0x00; + + // Second name at offset 28: api + pointer to google.com (offset 16) + packet[28] = 0x03; + packet[29] = 'a'; + packet[30] = 'p'; + packet[31] = 'i'; + packet[32] = 0xC0; // Compression pointer + packet[33] = 16; // Points to offset 16 (google.com) + + const result = try Name.parse(packet[28..], &packet, allocator); + defer { + var name = result.name; + name.deinit(); + } + + try testing.expectEqual(@as(usize, 6), result.bytes_read); // \x03api + 2 byte pointer + + const str = try result.name.toString(allocator); + defer allocator.free(str); + + try testing.expectEqualStrings("api.google.com", str); +} + +test "name from string" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("www.example.com", allocator); + defer name.deinit(); + + const str = try name.toString(allocator); + defer allocator.free(str); + + try testing.expectEqualStrings("www.example.com", str); + try testing.expectEqual(@as(usize, 3), name.labelCount()); +} + +test "name from string with trailing dot" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("www.example.com.", allocator); + defer name.deinit(); + + const str = try name.toString(allocator); + defer allocator.free(str); + + try testing.expectEqualStrings("www.example.com", str); +} + +test "root name" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("", allocator); + defer name.deinit(); + + const str = try name.toString(allocator); + defer allocator.free(str); + + try testing.expectEqualStrings("", str); + try testing.expectEqual(@as(usize, 0), name.labelCount()); +} + +test "name encode and decode roundtrip" { + const testing = std.testing; + const allocator = testing.allocator; + + var original = try Name.fromString("test.example.org", allocator); + defer original.deinit(); + + var buffer: [256]u8 = undefined; + const encoded_len = try original.encode(&buffer); + + const result = try Name.parse(buffer[0..encoded_len], buffer[0..encoded_len], allocator); + defer { + var name = result.name; + name.deinit(); + } + + try testing.expect(original.eql(result.name)); +} + +test "name equality case insensitive" { + const testing = std.testing; + const allocator = testing.allocator; + + var name1 = try Name.fromString("WWW.EXAMPLE.COM", allocator); + defer name1.deinit(); + + var name2 = try Name.fromString("www.example.com", allocator); + defer name2.deinit(); + + try testing.expect(name1.eql(name2)); +} + +test "label too long" { + const testing = std.testing; + const allocator = testing.allocator; + + // Create a label with 64 characters (max is 63) + const long_label = "a" ** 64 ++ ".example.com"; + try testing.expectError(error.LabelTooLong, Name.fromString(long_label, allocator)); +} diff --git a/src/dns/packet.zig b/src/dns/packet.zig new file mode 100644 index 0000000..c30262b --- /dev/null +++ b/src/dns/packet.zig @@ -0,0 +1,518 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; + +pub const types = @import("types.zig"); +pub const Header = @import("header.zig").Header; +pub const Name = @import("name.zig").Name; +pub const CompressionMap = @import("name.zig").CompressionMap; +pub const Question = @import("question.zig").Question; +pub const ResourceRecord = @import("record.zig").ResourceRecord; +pub const RData = @import("record.zig").RData; +pub const createARecord = @import("record.zig").createARecord; +pub const createAAAARecord = @import("record.zig").createAAAARecord; +pub const createCnameRecord = @import("record.zig").createCnameRecord; +pub const Edns = @import("edns.zig").Edns; +pub const hasEdns = @import("edns.zig").hasEdns; +pub const extractEdns = @import("edns.zig").extractEdns; + +/// TTL for blocked/synthetic responses (in seconds) +pub const DENIED_RESPONSE_TTL: u32 = 60; + +/// Complete DNS Packet +pub const Packet = struct { + header: Header, + questions: []Question, + answers: []ResourceRecord, + authority: []ResourceRecord, + additional: []ResourceRecord, + allocator: Allocator, + + /// Raw bytes of the original packet (for forwarding) + raw_bytes: ?[]const u8 = null, + + pub const ParseError = error{ + BufferTooSmall, + InvalidLabel, + NameTooLong, + LabelTooLong, + CompressionLoop, + InvalidPointer, + OutOfMemory, + InvalidRData, + HeaderParseError, + }; + + /// Parse a complete DNS packet from a buffer + pub fn parse(buffer: []const u8, allocator: Allocator) ParseError!Packet { + // Parse header + const header = Header.parse(buffer) catch return error.HeaderParseError; + + var offset: usize = Header.SIZE; + + // Parse questions + var questions = try allocator.alloc(Question, header.qdcount); + errdefer { + for (questions) |*q| { + q.deinit(); + } + allocator.free(questions); + } + + var q_idx: usize = 0; + while (q_idx < header.qdcount) : (q_idx += 1) { + const result = try Question.parse(buffer[offset..], buffer, allocator); + questions[q_idx] = result.question; + offset += result.bytes_read; + } + + // Parse answers + var answers = try allocator.alloc(ResourceRecord, header.ancount); + errdefer { + for (answers[0..]) |*a| { + a.deinit(allocator); + } + allocator.free(answers); + } + + var a_idx: usize = 0; + while (a_idx < header.ancount) : (a_idx += 1) { + const result = try ResourceRecord.parse(buffer[offset..], buffer, allocator); + answers[a_idx] = result.record; + offset += result.bytes_read; + } + + // Parse authority + var authority = try allocator.alloc(ResourceRecord, header.nscount); + errdefer { + for (authority[0..]) |*ns| { + ns.deinit(allocator); + } + allocator.free(authority); + } + + var ns_idx: usize = 0; + while (ns_idx < header.nscount) : (ns_idx += 1) { + const result = try ResourceRecord.parse(buffer[offset..], buffer, allocator); + authority[ns_idx] = result.record; + offset += result.bytes_read; + } + + // Parse additional + var additional = try allocator.alloc(ResourceRecord, header.arcount); + errdefer { + for (additional[0..]) |*ar| { + ar.deinit(allocator); + } + allocator.free(additional); + } + + var ar_idx: usize = 0; + while (ar_idx < header.arcount) : (ar_idx += 1) { + const result = try ResourceRecord.parse(buffer[offset..], buffer, allocator); + additional[ar_idx] = result.record; + offset += result.bytes_read; + } + + // Store raw bytes + const raw = try allocator.dupe(u8, buffer); + + return Packet{ + .header = header, + .questions = questions, + .answers = answers, + .authority = authority, + .additional = additional, + .allocator = allocator, + .raw_bytes = raw, + }; + } + + /// Encode the packet into a buffer + pub fn encode(self: Packet, buffer: []u8) ParseError!usize { + var offset: usize = 0; + + // Update header counts + var header = self.header; + header.qdcount = @intCast(self.questions.len); + header.ancount = @intCast(self.answers.len); + header.nscount = @intCast(self.authority.len); + header.arcount = @intCast(self.additional.len); + + // Encode header + header.encode(buffer[offset..]) catch return error.BufferTooSmall; + offset += Header.SIZE; + + // Encode questions + for (self.questions) |q| { + const len = try q.encode(buffer[offset..]); + offset += len; + } + + // Encode answers + for (self.answers) |a| { + const len = try a.encode(buffer[offset..], self.allocator); + offset += len; + } + + // Encode authority + for (self.authority) |ns| { + const len = try ns.encode(buffer[offset..], self.allocator); + offset += len; + } + + // Encode additional + for (self.additional) |ar| { + const len = try ar.encode(buffer[offset..], self.allocator); + offset += len; + } + + return offset; + } + + /// Free all packet memory + pub fn deinit(self: *Packet) void { + for (self.questions) |*q| { + q.deinit(); + } + self.allocator.free(self.questions); + + for (self.answers) |*a| { + a.deinit(self.allocator); + } + self.allocator.free(self.answers); + + for (self.authority) |*ns| { + ns.deinit(self.allocator); + } + self.allocator.free(self.authority); + + for (self.additional) |*ar| { + ar.deinit(self.allocator); + } + self.allocator.free(self.additional); + + if (self.raw_bytes) |raw| { + self.allocator.free(raw); + } + } + + /// Check if this packet is a query + pub fn isQuery(self: Packet) bool { + return self.header.isQuery(); + } + + /// Check if this packet is a response + pub fn isResponse(self: Packet) bool { + return self.header.isResponse(); + } + + /// Get the first question's domain name as a string + pub fn getQueryName(self: Packet) ![]u8 { + if (self.questions.len == 0) return error.NoQuestion; + return self.questions[0].name.toString(self.allocator); + } + + pub const NoQuestionError = error{NoQuestion}; + + /// Get the first question's type + pub fn getQueryType(self: Packet) ?types.QType { + if (self.questions.len == 0) return null; + return self.questions[0].qtype; + } + + /// Check if EDNS is present + pub fn hasEdnsSupport(self: Packet) bool { + return hasEdns(self.additional); + } + + /// Get EDNS information if present + pub fn getEdns(self: Packet) !?Edns { + return extractEdns(self.additional, self.allocator); + } + + /// Create a response packet from a query + pub fn createResponse(query: *const Packet, rcode: types.RCode, allocator: Allocator) !Packet { + const response_header = Header.createResponse(query.header, rcode); + + // Clone questions + var questions = try allocator.alloc(Question, query.questions.len); + errdefer allocator.free(questions); + + for (query.questions, 0..) |q, i| { + questions[i] = try q.clone(allocator); + } + + return Packet{ + .header = response_header, + .questions = questions, + .answers = try allocator.alloc(ResourceRecord, 0), + .authority = try allocator.alloc(ResourceRecord, 0), + .additional = try allocator.alloc(ResourceRecord, 0), + .allocator = allocator, + }; + } + + /// Create a denied response (returns 0.0.0.0 for A queries) + pub fn createDeniedResponse(query: *const Packet, allocator: Allocator) !Packet { + var response = try createResponse(query, types.RCode.NoError, allocator); + errdefer response.deinit(); + + if (query.questions.len > 0) { + const q = query.questions[0]; + + // Clone the name for the answer + const labels = try allocator.dupe(u8, q.name.labels); + const answer_name = Name{ + .labels = labels, + .allocator = allocator, + }; + + var answers = try allocator.alloc(ResourceRecord, 1); + + if (q.qtype == types.QType.A) { + answers[0] = createARecord(answer_name, DENIED_RESPONSE_TTL, [4]u8{ 0, 0, 0, 0 }); + } else if (q.qtype == types.QType.AAAA) { + answers[0] = createAAAARecord(answer_name, DENIED_RESPONSE_TTL, [16]u8{ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 }); + } else { + // For other types, just return NXDOMAIN + answer_name.allocator.?.free(answer_name.labels); + allocator.free(answers); + response.header.rcode = types.RCode.NXDomain; + return response; + } + + allocator.free(response.answers); + response.answers = answers; + response.header.ancount = 1; + } + + return response; + } + + /// Create an NXDOMAIN response + pub fn createNxdomainResponse(query: *const Packet, allocator: Allocator) !Packet { + return createResponse(query, types.RCode.NXDomain, allocator); + } + + /// Create a safe search CNAME response that redirects to the safe domain + pub fn createSafeSearchResponse(query: *const Packet, safe_domain: []const u8, allocator: Allocator) !Packet { + var response = try createResponse(query, types.RCode.NoError, allocator); + errdefer response.deinit(); + + if (query.questions.len > 0) { + const q = query.questions[0]; + + // Clone the name for the answer + const labels = try allocator.dupe(u8, q.name.labels); + const answer_name = Name{ + .labels = labels, + .allocator = allocator, + }; + + // Create the CNAME target + const target = try Name.fromString(safe_domain, allocator); + + var answers = try allocator.alloc(ResourceRecord, 1); + answers[0] = createCnameRecord(answer_name, DENIED_RESPONSE_TTL, target); + + allocator.free(response.answers); + response.answers = answers; + response.header.ancount = 1; + } + + return response; + } + + /// Create a SERVFAIL response + pub fn createServfailResponse(query: *const Packet, allocator: Allocator) !Packet { + return createResponse(query, types.RCode.ServFail, allocator); + } + + /// Create a FORMERR response + pub fn createFormerrResponse(query: *const Packet, allocator: Allocator) !Packet { + return createResponse(query, types.RCode.FormErr, allocator); + } +}; + +/// Parse only the header from a buffer (for quick inspection) +pub fn parseHeaderOnly(buffer: []const u8) !Header { + return Header.parse(buffer); +} + +/// Get the transaction ID from a buffer without full parsing +pub fn getTransactionId(buffer: []const u8) ?u16 { + if (buffer.len < 2) return null; + return std.mem.readInt(u16, buffer[0..2], .big); +} + +test "parse simple query packet" { + const testing = std.testing; + const allocator = testing.allocator; + + // Simple A query for www.google.com + const query = [_]u8{ + // Header + 0x00, 0x01, // ID = 1 + 0x01, 0x00, // Standard query, RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + // Question + 0x03, 'w', 'w', 'w', // www + 0x06, 'g', 'o', 'o', 'g', 'l', 'e', // google + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; + + var packet = try Packet.parse(&query, allocator); + defer packet.deinit(); + + try testing.expectEqual(@as(u16, 1), packet.header.id); + try testing.expect(packet.isQuery()); + try testing.expect(!packet.isResponse()); + try testing.expectEqual(@as(usize, 1), packet.questions.len); + try testing.expectEqual(@as(usize, 0), packet.answers.len); + + const name = try packet.getQueryName(); + defer allocator.free(name); + + try testing.expectEqualStrings("www.google.com", name); + try testing.expectEqual(types.QType.A, packet.getQueryType().?); +} + +test "parse response packet with A record" { + const testing = std.testing; + const allocator = testing.allocator; + + // Response for example.com with A record 93.184.216.34 + const response = [_]u8{ + // Header + 0x00, 0x02, // ID = 2 + 0x81, 0x80, // Response, RD=1, RA=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x01, // ANCOUNT = 1 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + // Question + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + // Answer + 0xC0, 0x0C, // Pointer to offset 12 (example.com) + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x01, 0x2C, // TTL = 300 + 0x00, 0x04, // RDLENGTH = 4 + 0x5D, 0xB8, 0xD8, 0x22, // 93.184.216.34 + }; + + var packet = try Packet.parse(&response, allocator); + defer packet.deinit(); + + try testing.expectEqual(@as(u16, 2), packet.header.id); + try testing.expect(packet.isResponse()); + try testing.expectEqual(@as(usize, 1), packet.questions.len); + try testing.expectEqual(@as(usize, 1), packet.answers.len); + + const answer = packet.answers[0]; + try testing.expectEqual(types.QType.A, answer.rtype); + try testing.expectEqual(@as(u32, 300), answer.ttl); + + const ip = answer.getIPv4().?; + try testing.expectEqual([4]u8{ 93, 184, 216, 34 }, ip); +} + +test "create blocked response" { + const testing = std.testing; + const allocator = testing.allocator; + + // Create a query + const query_bytes = [_]u8{ + // Header + 0x00, 0x03, // ID = 3 + 0x01, 0x00, // Standard query, RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + // Question + 0x09, 'd', 'o', 'u', 'b', 'l', 'e', 'c', 'l', 'k', // doubleclk + 0x03, 'n', 'e', 't', // net + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; + + var query = try Packet.parse(&query_bytes, allocator); + defer query.deinit(); + + var denied = try Packet.createDeniedResponse(&query, allocator); + defer denied.deinit(); + + try testing.expect(denied.isResponse()); + try testing.expectEqual(@as(u16, 3), denied.header.id); + try testing.expectEqual(@as(usize, 1), denied.answers.len); + + const answer = denied.answers[0]; + try testing.expectEqual(types.QType.A, answer.rtype); + try testing.expectEqual([4]u8{ 0, 0, 0, 0 }, answer.getIPv4().?); +} + +test "encode and decode packet roundtrip" { + const testing = std.testing; + const allocator = testing.allocator; + + // Simple query + const original_bytes = [_]u8{ + // Header + 0xAB, 0xCD, // ID + 0x01, 0x00, // Standard query, RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + // Question + 0x04, 't', 'e', 's', 't', // test + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; + + var original = try Packet.parse(&original_bytes, allocator); + defer original.deinit(); + + // Encode + var buffer: [512]u8 = undefined; + const encoded_len = try original.encode(&buffer); + + // Decode + var decoded = try Packet.parse(buffer[0..encoded_len], allocator); + defer decoded.deinit(); + + try testing.expectEqual(original.header.id, decoded.header.id); + try testing.expectEqual(original.questions.len, decoded.questions.len); + + const orig_name = try original.getQueryName(); + defer allocator.free(orig_name); + + const dec_name = try decoded.getQueryName(); + defer allocator.free(dec_name); + + try testing.expectEqualStrings(orig_name, dec_name); +} + +test "get transaction ID" { + const testing = std.testing; + + const buffer = [_]u8{ 0x12, 0x34, 0x01, 0x00 }; + try testing.expectEqual(@as(u16, 0x1234), getTransactionId(&buffer).?); + + const short = [_]u8{0x12}; + try testing.expect(getTransactionId(&short) == null); +} diff --git a/src/dns/question.zig b/src/dns/question.zig new file mode 100644 index 0000000..5bab6ee --- /dev/null +++ b/src/dns/question.zig @@ -0,0 +1,214 @@ +const std = @import("std"); +const types = @import("types.zig"); +const Name = @import("name.zig").Name; +const Allocator = std.mem.Allocator; + +/// DNS Question Section +/// Format: +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | QNAME | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | QTYPE | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | QCLASS | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +pub const Question = struct { + /// Query name + name: Name, + /// Query type + qtype: types.QType, + /// Query class + qclass: types.QClass, + + pub const ParseError = error{ + BufferTooSmall, + InvalidLabel, + NameTooLong, + LabelTooLong, + CompressionLoop, + InvalidPointer, + OutOfMemory, + }; + + pub const ParseResult = struct { + question: Question, + bytes_read: usize, + }; + + /// Parse a question from a buffer + pub fn parse(buffer: []const u8, packet_start: []const u8, allocator: Allocator) ParseError!ParseResult { + // Parse the name first + const name_result = try Name.parse(buffer, packet_start, allocator); + errdefer { + var n = name_result.name; + n.deinit(); + } + + const remaining = buffer[name_result.bytes_read..]; + if (remaining.len < 4) { + var n = name_result.name; + n.deinit(); + return error.BufferTooSmall; + } + + const qtype = std.mem.readInt(u16, remaining[0..2], .big); + const qclass = std.mem.readInt(u16, remaining[2..4], .big); + + return ParseResult{ + .question = Question{ + .name = name_result.name, + .qtype = @enumFromInt(qtype), + .qclass = @enumFromInt(qclass), + }, + .bytes_read = name_result.bytes_read + 4, + }; + } + + /// Encode the question into a buffer + pub fn encode(self: Question, buffer: []u8) ParseError!usize { + var offset: usize = 0; + + // Encode name + const name_len = try self.name.encode(buffer[offset..]); + offset += name_len; + + if (buffer.len < offset + 4) { + return error.BufferTooSmall; + } + + // Encode type and class + std.mem.writeInt(u16, buffer[offset..][0..2], @intFromEnum(self.qtype), .big); + offset += 2; + std.mem.writeInt(u16, buffer[offset..][0..2], @intFromEnum(self.qclass), .big); + offset += 2; + + return offset; + } + + /// Get the encoded length of the question + pub fn encodedLen(self: Question) usize { + return self.name.encodedLen() + 4; // name + qtype (2) + qclass (2) + } + + /// Free the question's memory + pub fn deinit(self: *Question) void { + self.name.deinit(); + } + + /// Create a copy of the question + pub fn clone(self: Question, allocator: Allocator) !Question { + const labels = try allocator.dupe(u8, self.name.labels); + return Question{ + .name = Name{ + .labels = labels, + .allocator = allocator, + }, + .qtype = self.qtype, + .qclass = self.qclass, + }; + } +}; + +test "parse simple question" { + const testing = std.testing; + const allocator = testing.allocator; + + // www.google.com A IN + const buffer = [_]u8{ + 0x03, 'w', 'w', 'w', // www + 0x06, 'g', 'o', 'o', 'g', 'l', 'e', // google + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; + + const result = try Question.parse(&buffer, &buffer, allocator); + defer { + var q = result.question; + q.deinit(); + } + + try testing.expectEqual(@as(usize, 20), result.bytes_read); + try testing.expectEqual(types.QType.A, result.question.qtype); + try testing.expectEqual(types.QClass.IN, result.question.qclass); + + const name_str = try result.question.name.toString(allocator); + defer allocator.free(name_str); + try testing.expectEqualStrings("www.google.com", name_str); +} + +test "parse AAAA question" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com AAAA IN + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x1C, // TYPE = AAAA (28) + 0x00, 0x01, // CLASS = IN + }; + + const result = try Question.parse(&buffer, &buffer, allocator); + defer { + var q = result.question; + q.deinit(); + } + + try testing.expectEqual(types.QType.AAAA, result.question.qtype); + try testing.expectEqual(types.QClass.IN, result.question.qclass); +} + +test "encode and decode question roundtrip" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("test.example.org", allocator); + errdefer name.deinit(); + + const original = Question{ + .name = name, + .qtype = types.QType.MX, + .qclass = types.QClass.IN, + }; + defer { + var q = @as(Question, original); + q.deinit(); + } + + var buffer: [256]u8 = undefined; + const encoded_len = try original.encode(&buffer); + + const result = try Question.parse(buffer[0..encoded_len], buffer[0..encoded_len], allocator); + defer { + var q = result.question; + q.deinit(); + } + + try testing.expectEqual(original.qtype, result.question.qtype); + try testing.expectEqual(original.qclass, result.question.qclass); + try testing.expect(original.name.eql(result.question.name)); +} + +test "question encoded length" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("www.example.com", allocator); + errdefer name.deinit(); + + const q = Question{ + .name = name, + .qtype = types.QType.A, + .qclass = types.QClass.IN, + }; + defer { + var question = @as(Question, q); + question.deinit(); + } + + // www(1+3) + example(1+7) + com(1+3) + null(1) + type(2) + class(2) = 21 + try testing.expectEqual(@as(usize, 21), q.encodedLen()); +} diff --git a/src/dns/record.zig b/src/dns/record.zig new file mode 100644 index 0000000..6ac381a --- /dev/null +++ b/src/dns/record.zig @@ -0,0 +1,983 @@ +const std = @import("std"); +const types = @import("types.zig"); +const Name = @import("name.zig").Name; +const Allocator = std.mem.Allocator; + +/// Resource Record Data (tagged union) +pub const RData = union(enum) { + /// A record - IPv4 address (4 bytes) + a: [4]u8, + /// AAAA record - IPv6 address (16 bytes) + aaaa: [16]u8, + /// CNAME record - canonical name + cname: Name, + /// NS record - name server + ns: Name, + /// PTR record - pointer + ptr: Name, + /// MX record - mail exchange + mx: struct { + preference: u16, + exchange: Name, + }, + /// TXT record - text strings + txt: []const []const u8, + /// SOA record - start of authority + soa: struct { + mname: Name, + rname: Name, + serial: u32, + refresh: u32, + retry: u32, + expire: u32, + minimum: u32, + }, + /// SRV record - service locator + srv: struct { + priority: u16, + weight: u16, + port: u16, + target: Name, + }, + /// OPT record (EDNS) - stored as raw data + opt: []const u8, + + // DNSSEC record types + + /// DNSKEY record - public key for DNSSEC + dnskey: struct { + flags: u16, // Zone key flag (bit 7), SEP flag (bit 15) + protocol: u8, // Must be 3 + algorithm: u8, // DNSSEC algorithm number + public_key: []const u8, + }, + + /// DS record - Delegation Signer + ds: struct { + key_tag: u16, + algorithm: u8, + digest_type: u8, + digest: []const u8, + }, + + /// RRSIG record - signature over RRset + rrsig: struct { + type_covered: u16, + algorithm: u8, + labels: u8, + original_ttl: u32, + signature_expiration: u32, + signature_inception: u32, + key_tag: u16, + signer_name: Name, + signature: []const u8, + }, + + /// NSEC record - authenticated denial of existence + nsec: struct { + next_domain: Name, + type_bitmap: []const u8, + }, + + /// NSEC3 record - hashed authenticated denial + nsec3: struct { + hash_algorithm: u8, + flags: u8, + iterations: u16, + salt: []const u8, + next_hashed_owner: []const u8, + type_bitmap: []const u8, + }, + + /// NSEC3PARAM record - NSEC3 parameters + nsec3param: struct { + hash_algorithm: u8, + flags: u8, + iterations: u16, + salt: []const u8, + }, + + /// Unknown/raw record data + raw: []const u8, + + pub fn deinit(self: *RData, allocator: Allocator) void { + switch (self.*) { + .cname => |*name| name.deinit(), + .ns => |*name| name.deinit(), + .ptr => |*name| name.deinit(), + .mx => |*mx| mx.exchange.deinit(), + .txt => |txt| { + for (txt) |s| { + allocator.free(s); + } + allocator.free(txt); + }, + .soa => |*soa| { + soa.mname.deinit(); + soa.rname.deinit(); + }, + .srv => |*srv| srv.target.deinit(), + .opt => |data| allocator.free(data), + // DNSSEC types + .dnskey => |*dk| allocator.free(dk.public_key), + .ds => |*ds| allocator.free(ds.digest), + .rrsig => |*rrsig| { + rrsig.signer_name.deinit(); + allocator.free(rrsig.signature); + }, + .nsec => |*nsec| { + nsec.next_domain.deinit(); + allocator.free(nsec.type_bitmap); + }, + .nsec3 => |*nsec3| { + allocator.free(nsec3.salt); + allocator.free(nsec3.next_hashed_owner); + allocator.free(nsec3.type_bitmap); + }, + .nsec3param => |*np| allocator.free(np.salt), + .raw => |data| allocator.free(data), + .a, .aaaa => {}, + } + } +}; + +/// DNS Resource Record +/// Format: +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | NAME | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | TYPE | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | CLASS | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | TTL | +/// | | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | RDLENGTH | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +/// | RDATA | +/// +--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+--+ +pub const ResourceRecord = struct { + name: Name, + rtype: types.QType, + class: types.QClass, + ttl: u32, + rdata: RData, + + pub const ParseError = error{ + BufferTooSmall, + InvalidLabel, + NameTooLong, + LabelTooLong, + CompressionLoop, + InvalidPointer, + OutOfMemory, + InvalidRData, + }; + + pub const ParseResult = struct { + record: ResourceRecord, + bytes_read: usize, + }; + + /// Parse a resource record from a buffer + pub fn parse(buffer: []const u8, packet_start: []const u8, allocator: Allocator) ParseError!ParseResult { + // Parse the name + const name_result = try Name.parse(buffer, packet_start, allocator); + errdefer { + var n = name_result.name; + n.deinit(); + } + + var offset = name_result.bytes_read; + const remaining = buffer[offset..]; + + if (remaining.len < 10) { + var n = name_result.name; + n.deinit(); + return error.BufferTooSmall; + } + + const rtype_val = std.mem.readInt(u16, remaining[0..2], .big); + const class_val = std.mem.readInt(u16, remaining[2..4], .big); + const ttl = std.mem.readInt(u32, remaining[4..8], .big); + const rdlength = std.mem.readInt(u16, remaining[8..10], .big); + + offset += 10; + + const rdata_buf = buffer[offset..]; + if (rdata_buf.len < rdlength) { + var n = name_result.name; + n.deinit(); + return error.BufferTooSmall; + } + + const rtype: types.QType = @enumFromInt(rtype_val); + const rdata = try parseRData(rtype, rdata_buf[0..rdlength], packet_start, allocator); + errdefer { + var rd = rdata; + rd.deinit(allocator); + } + + return ParseResult{ + .record = ResourceRecord{ + .name = name_result.name, + .rtype = rtype, + .class = @enumFromInt(class_val), + .ttl = ttl, + .rdata = rdata, + }, + .bytes_read = offset + rdlength, + }; + } + + fn parseRData(rtype: types.QType, data: []const u8, packet_start: []const u8, allocator: Allocator) ParseError!RData { + switch (rtype) { + .A => { + if (data.len != 4) return error.InvalidRData; + return RData{ .a = data[0..4].* }; + }, + .AAAA => { + if (data.len != 16) return error.InvalidRData; + return RData{ .aaaa = data[0..16].* }; + }, + .CNAME => { + const name_result = try Name.parse(data, packet_start, allocator); + return RData{ .cname = name_result.name }; + }, + .NS => { + const name_result = try Name.parse(data, packet_start, allocator); + return RData{ .ns = name_result.name }; + }, + .PTR => { + const name_result = try Name.parse(data, packet_start, allocator); + return RData{ .ptr = name_result.name }; + }, + .MX => { + if (data.len < 3) return error.InvalidRData; + const preference = std.mem.readInt(u16, data[0..2], .big); + const name_result = try Name.parse(data[2..], packet_start, allocator); + return RData{ + .mx = .{ + .preference = preference, + .exchange = name_result.name, + }, + }; + }, + .TXT => { + var strings = std.ArrayListUnmanaged([]const u8){}; + errdefer { + for (strings.items) |s| { + allocator.free(s); + } + strings.deinit(allocator); + } + + var pos: usize = 0; + while (pos < data.len) { + const str_len = data[pos]; + pos += 1; + if (pos + str_len > data.len) return error.InvalidRData; + const str = try allocator.dupe(u8, data[pos .. pos + str_len]); + try strings.append(allocator, str); + pos += str_len; + } + + return RData{ .txt = try strings.toOwnedSlice(allocator) }; + }, + .SOA => { + var offset: usize = 0; + + const mname_result = try Name.parse(data, packet_start, allocator); + errdefer { + var n = mname_result.name; + n.deinit(); + } + offset += mname_result.bytes_read; + + const rname_result = try Name.parse(data[offset..], packet_start, allocator); + errdefer { + var n = rname_result.name; + n.deinit(); + } + offset += rname_result.bytes_read; + + if (data.len < offset + 20) return error.InvalidRData; + + const serial = std.mem.readInt(u32, data[offset..][0..4], .big); + const refresh = std.mem.readInt(u32, data[offset + 4 ..][0..4], .big); + const retry = std.mem.readInt(u32, data[offset + 8 ..][0..4], .big); + const expire = std.mem.readInt(u32, data[offset + 12 ..][0..4], .big); + const minimum = std.mem.readInt(u32, data[offset + 16 ..][0..4], .big); + + return RData{ + .soa = .{ + .mname = mname_result.name, + .rname = rname_result.name, + .serial = serial, + .refresh = refresh, + .retry = retry, + .expire = expire, + .minimum = minimum, + }, + }; + }, + .SRV => { + if (data.len < 7) return error.InvalidRData; + const priority = std.mem.readInt(u16, data[0..2], .big); + const weight = std.mem.readInt(u16, data[2..4], .big); + const port = std.mem.readInt(u16, data[4..6], .big); + const target_result = try Name.parse(data[6..], packet_start, allocator); + return RData{ + .srv = .{ + .priority = priority, + .weight = weight, + .port = port, + .target = target_result.name, + }, + }; + }, + .OPT => { + const opt_data = try allocator.dupe(u8, data); + return RData{ .opt = opt_data }; + }, + // DNSSEC record types + .DNSKEY => { + if (data.len < 4) return error.InvalidRData; + const flags = std.mem.readInt(u16, data[0..2], .big); + const protocol = data[2]; + const algorithm = data[3]; + const public_key = try allocator.dupe(u8, data[4..]); + return RData{ .dnskey = .{ + .flags = flags, + .protocol = protocol, + .algorithm = algorithm, + .public_key = public_key, + } }; + }, + .DS => { + if (data.len < 4) return error.InvalidRData; + const key_tag = std.mem.readInt(u16, data[0..2], .big); + const algorithm = data[2]; + const digest_type = data[3]; + const digest = try allocator.dupe(u8, data[4..]); + return RData{ .ds = .{ + .key_tag = key_tag, + .algorithm = algorithm, + .digest_type = digest_type, + .digest = digest, + } }; + }, + .RRSIG => { + // RRSIG format: type_covered(2) + algorithm(1) + labels(1) + original_ttl(4) + + // sig_expiration(4) + sig_inception(4) + key_tag(2) + signer_name + signature + if (data.len < 18) return error.InvalidRData; + const type_covered = std.mem.readInt(u16, data[0..2], .big); + const algorithm = data[2]; + const labels = data[3]; + const original_ttl = std.mem.readInt(u32, data[4..8], .big); + const signature_expiration = std.mem.readInt(u32, data[8..12], .big); + const signature_inception = std.mem.readInt(u32, data[12..16], .big); + const key_tag = std.mem.readInt(u16, data[16..18], .big); + + const name_result = try Name.parse(data[18..], packet_start, allocator); + errdefer { + var n = name_result.name; + n.deinit(); + } + + const sig_start = 18 + name_result.bytes_read; + const signature = try allocator.dupe(u8, data[sig_start..]); + + return RData{ .rrsig = .{ + .type_covered = type_covered, + .algorithm = algorithm, + .labels = labels, + .original_ttl = original_ttl, + .signature_expiration = signature_expiration, + .signature_inception = signature_inception, + .key_tag = key_tag, + .signer_name = name_result.name, + .signature = signature, + } }; + }, + .NSEC => { + const name_result = try Name.parse(data, packet_start, allocator); + errdefer { + var n = name_result.name; + n.deinit(); + } + const type_bitmap = try allocator.dupe(u8, data[name_result.bytes_read..]); + return RData{ .nsec = .{ + .next_domain = name_result.name, + .type_bitmap = type_bitmap, + } }; + }, + .NSEC3 => { + // NSEC3 format: hash_algorithm(1) + flags(1) + iterations(2) + salt_length(1) + + // salt + hash_length(1) + hash + type_bitmap + if (data.len < 5) return error.InvalidRData; + const hash_algorithm = data[0]; + const flags = data[1]; + const iterations = std.mem.readInt(u16, data[2..4], .big); + const salt_length = data[4]; + + var offset: usize = 5; + if (data.len < offset + salt_length) return error.InvalidRData; + const salt = try allocator.dupe(u8, data[offset .. offset + salt_length]); + errdefer allocator.free(salt); + offset += salt_length; + + if (data.len < offset + 1) return error.InvalidRData; + const hash_length = data[offset]; + offset += 1; + + if (data.len < offset + hash_length) return error.InvalidRData; + const next_hashed_owner = try allocator.dupe(u8, data[offset .. offset + hash_length]); + errdefer allocator.free(next_hashed_owner); + offset += hash_length; + + const type_bitmap = try allocator.dupe(u8, data[offset..]); + + return RData{ .nsec3 = .{ + .hash_algorithm = hash_algorithm, + .flags = flags, + .iterations = iterations, + .salt = salt, + .next_hashed_owner = next_hashed_owner, + .type_bitmap = type_bitmap, + } }; + }, + .NSEC3PARAM => { + if (data.len < 5) return error.InvalidRData; + const hash_algorithm = data[0]; + const flags = data[1]; + const iterations = std.mem.readInt(u16, data[2..4], .big); + const salt_length = data[4]; + + if (data.len < 5 + salt_length) return error.InvalidRData; + const salt = try allocator.dupe(u8, data[5 .. 5 + salt_length]); + + return RData{ .nsec3param = .{ + .hash_algorithm = hash_algorithm, + .flags = flags, + .iterations = iterations, + .salt = salt, + } }; + }, + else => { + const raw_data = try allocator.dupe(u8, data); + return RData{ .raw = raw_data }; + }, + } + } + + /// Encode the resource record into a buffer + pub fn encode(self: ResourceRecord, buffer: []u8, allocator: Allocator) ParseError!usize { + var offset: usize = 0; + + // Encode name + const name_len = try self.name.encode(buffer[offset..]); + offset += name_len; + + if (buffer.len < offset + 10) { + return error.BufferTooSmall; + } + + // Encode type, class, ttl + std.mem.writeInt(u16, buffer[offset..][0..2], @intFromEnum(self.rtype), .big); + offset += 2; + std.mem.writeInt(u16, buffer[offset..][0..2], @intFromEnum(self.class), .big); + offset += 2; + std.mem.writeInt(u32, buffer[offset..][0..4], self.ttl, .big); + offset += 4; + + // Reserve space for rdlength + const rdlength_offset = offset; + offset += 2; + + // Encode rdata + const rdata_len = try self.encodeRData(buffer[offset..], allocator); + offset += rdata_len; + + // Write rdlength + std.mem.writeInt(u16, buffer[rdlength_offset..][0..2], @intCast(rdata_len), .big); + + return offset; + } + + fn encodeRData(self: ResourceRecord, buffer: []u8, allocator: Allocator) ParseError!usize { + _ = allocator; + switch (self.rdata) { + .a => |addr| { + if (buffer.len < 4) return error.BufferTooSmall; + @memcpy(buffer[0..4], &addr); + return 4; + }, + .aaaa => |addr| { + if (buffer.len < 16) return error.BufferTooSmall; + @memcpy(buffer[0..16], &addr); + return 16; + }, + .cname => |name| { + return name.encode(buffer); + }, + .ns => |name| { + return name.encode(buffer); + }, + .ptr => |name| { + return name.encode(buffer); + }, + .mx => |mx| { + if (buffer.len < 2) return error.BufferTooSmall; + std.mem.writeInt(u16, buffer[0..2], mx.preference, .big); + const name_len = try mx.exchange.encode(buffer[2..]); + return 2 + name_len; + }, + .txt => |strings| { + var offset: usize = 0; + for (strings) |s| { + if (buffer.len < offset + 1 + s.len) return error.BufferTooSmall; + buffer[offset] = @intCast(s.len); + offset += 1; + @memcpy(buffer[offset .. offset + s.len], s); + offset += s.len; + } + return offset; + }, + .soa => |soa| { + var offset: usize = 0; + offset += try soa.mname.encode(buffer[offset..]); + offset += try soa.rname.encode(buffer[offset..]); + if (buffer.len < offset + 20) return error.BufferTooSmall; + std.mem.writeInt(u32, buffer[offset..][0..4], soa.serial, .big); + offset += 4; + std.mem.writeInt(u32, buffer[offset..][0..4], soa.refresh, .big); + offset += 4; + std.mem.writeInt(u32, buffer[offset..][0..4], soa.retry, .big); + offset += 4; + std.mem.writeInt(u32, buffer[offset..][0..4], soa.expire, .big); + offset += 4; + std.mem.writeInt(u32, buffer[offset..][0..4], soa.minimum, .big); + offset += 4; + return offset; + }, + .srv => |srv| { + if (buffer.len < 6) return error.BufferTooSmall; + std.mem.writeInt(u16, buffer[0..2], srv.priority, .big); + std.mem.writeInt(u16, buffer[2..4], srv.weight, .big); + std.mem.writeInt(u16, buffer[4..6], srv.port, .big); + const name_len = try srv.target.encode(buffer[6..]); + return 6 + name_len; + }, + .opt => |data| { + if (buffer.len < data.len) return error.BufferTooSmall; + @memcpy(buffer[0..data.len], data); + return data.len; + }, + // DNSSEC types + .dnskey => |dk| { + const min_len = 4 + dk.public_key.len; + if (buffer.len < min_len) return error.BufferTooSmall; + std.mem.writeInt(u16, buffer[0..2], dk.flags, .big); + buffer[2] = dk.protocol; + buffer[3] = dk.algorithm; + @memcpy(buffer[4 .. 4 + dk.public_key.len], dk.public_key); + return min_len; + }, + .ds => |ds| { + const min_len = 4 + ds.digest.len; + if (buffer.len < min_len) return error.BufferTooSmall; + std.mem.writeInt(u16, buffer[0..2], ds.key_tag, .big); + buffer[2] = ds.algorithm; + buffer[3] = ds.digest_type; + @memcpy(buffer[4 .. 4 + ds.digest.len], ds.digest); + return min_len; + }, + .rrsig => |rrsig| { + if (buffer.len < 18) return error.BufferTooSmall; + std.mem.writeInt(u16, buffer[0..2], rrsig.type_covered, .big); + buffer[2] = rrsig.algorithm; + buffer[3] = rrsig.labels; + std.mem.writeInt(u32, buffer[4..8], rrsig.original_ttl, .big); + std.mem.writeInt(u32, buffer[8..12], rrsig.signature_expiration, .big); + std.mem.writeInt(u32, buffer[12..16], rrsig.signature_inception, .big); + std.mem.writeInt(u16, buffer[16..18], rrsig.key_tag, .big); + const name_len = try rrsig.signer_name.encode(buffer[18..]); + const sig_start = 18 + name_len; + if (buffer.len < sig_start + rrsig.signature.len) return error.BufferTooSmall; + @memcpy(buffer[sig_start .. sig_start + rrsig.signature.len], rrsig.signature); + return sig_start + rrsig.signature.len; + }, + .nsec => |nsec| { + const name_len = try nsec.next_domain.encode(buffer); + if (buffer.len < name_len + nsec.type_bitmap.len) return error.BufferTooSmall; + @memcpy(buffer[name_len .. name_len + nsec.type_bitmap.len], nsec.type_bitmap); + return name_len + nsec.type_bitmap.len; + }, + .nsec3 => |nsec3| { + const min_len = 5 + nsec3.salt.len + 1 + nsec3.next_hashed_owner.len + nsec3.type_bitmap.len; + if (buffer.len < min_len) return error.BufferTooSmall; + buffer[0] = nsec3.hash_algorithm; + buffer[1] = nsec3.flags; + std.mem.writeInt(u16, buffer[2..4], nsec3.iterations, .big); + buffer[4] = @intCast(nsec3.salt.len); + var offset: usize = 5; + @memcpy(buffer[offset .. offset + nsec3.salt.len], nsec3.salt); + offset += nsec3.salt.len; + buffer[offset] = @intCast(nsec3.next_hashed_owner.len); + offset += 1; + @memcpy(buffer[offset .. offset + nsec3.next_hashed_owner.len], nsec3.next_hashed_owner); + offset += nsec3.next_hashed_owner.len; + @memcpy(buffer[offset .. offset + nsec3.type_bitmap.len], nsec3.type_bitmap); + return offset + nsec3.type_bitmap.len; + }, + .nsec3param => |np| { + const min_len = 5 + np.salt.len; + if (buffer.len < min_len) return error.BufferTooSmall; + buffer[0] = np.hash_algorithm; + buffer[1] = np.flags; + std.mem.writeInt(u16, buffer[2..4], np.iterations, .big); + buffer[4] = @intCast(np.salt.len); + @memcpy(buffer[5 .. 5 + np.salt.len], np.salt); + return min_len; + }, + .raw => |data| { + if (buffer.len < data.len) return error.BufferTooSmall; + @memcpy(buffer[0..data.len], data); + return data.len; + }, + } + } + + /// Free the record's memory + pub fn deinit(self: *ResourceRecord, allocator: Allocator) void { + self.name.deinit(); + self.rdata.deinit(allocator); + } + + /// Get the IPv4 address if this is an A record + pub fn getIPv4(self: ResourceRecord) ?[4]u8 { + return switch (self.rdata) { + .a => |addr| addr, + else => null, + }; + } + + /// Get the IPv6 address if this is an AAAA record + pub fn getIPv6(self: ResourceRecord) ?[16]u8 { + return switch (self.rdata) { + .aaaa => |addr| addr, + else => null, + }; + } + + /// Get the CNAME target if this is a CNAME record + pub fn getCname(self: ResourceRecord) ?Name { + return switch (self.rdata) { + .cname => |name| name, + else => null, + }; + } +}; + +/// Create an A record +pub fn createARecord(name: Name, ttl: u32, ip: [4]u8) ResourceRecord { + return ResourceRecord{ + .name = name, + .rtype = types.QType.A, + .class = types.QClass.IN, + .ttl = ttl, + .rdata = RData{ .a = ip }, + }; +} + +/// Create an AAAA record +pub fn createAAAARecord(name: Name, ttl: u32, ip: [16]u8) ResourceRecord { + return ResourceRecord{ + .name = name, + .rtype = types.QType.AAAA, + .class = types.QClass.IN, + .ttl = ttl, + .rdata = RData{ .aaaa = ip }, + }; +} + +/// Create a CNAME record +pub fn createCnameRecord(name: Name, ttl: u32, target: Name) ResourceRecord { + return ResourceRecord{ + .name = name, + .rtype = types.QType.CNAME, + .class = types.QClass.IN, + .ttl = ttl, + .rdata = RData{ .cname = target }, + }; +} + +test "parse A record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com A 1.2.3.4 TTL=300 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x01, 0x2C, // TTL = 300 + 0x00, 0x04, // RDLENGTH = 4 + 0x01, 0x02, 0x03, 0x04, // RDATA = 1.2.3.4 + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.A, result.record.rtype); + try testing.expectEqual(types.QClass.IN, result.record.class); + try testing.expectEqual(@as(u32, 300), result.record.ttl); + + const ip = result.record.getIPv4().?; + try testing.expectEqual([4]u8{ 1, 2, 3, 4 }, ip); +} + +test "parse AAAA record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com AAAA 2001:db8::1 TTL=600 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x1C, // TYPE = AAAA (28) + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x02, 0x58, // TTL = 600 + 0x00, 0x10, // RDLENGTH = 16 + 0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, // 2001:0db8:0000:0000: + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, // 0000:0000:0000:0001 + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.AAAA, result.record.rtype); + try testing.expectEqual(@as(u32, 600), result.record.ttl); + + const ip = result.record.getIPv6().?; + try testing.expectEqual([16]u8{ 0x20, 0x01, 0x0d, 0xb8, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1 }, ip); +} + +test "parse CNAME record" { + const testing = std.testing; + const allocator = testing.allocator; + + // www.example.com CNAME example.com TTL=3600 + const buffer = [_]u8{ + 0x03, 'w', 'w', 'w', // www + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x05, // TYPE = CNAME + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x0E, 0x10, // TTL = 3600 + 0x00, 0x0D, // RDLENGTH = 13 + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.CNAME, result.record.rtype); + try testing.expectEqual(@as(u32, 3600), result.record.ttl); + + const cname = result.record.getCname().?; + const cname_str = try cname.toString(allocator); + defer allocator.free(cname_str); + try testing.expectEqualStrings("example.com", cname_str); +} + +test "parse TXT record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com TXT "v=spf1 -all" TTL=300 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x10, // TYPE = TXT (16) + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x01, 0x2C, // TTL = 300 + 0x00, 0x0C, // RDLENGTH = 12 + 0x0B, // String length = 11 + 'v', '=', 's', 'p', 'f', '1', ' ', '-', 'a', 'l', 'l', + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.TXT, result.record.rtype); + + switch (result.record.rdata) { + .txt => |strings| { + try testing.expectEqual(@as(usize, 1), strings.len); + try testing.expectEqualStrings("v=spf1 -all", strings[0]); + }, + else => try testing.expect(false), + } +} + +test "create and encode A record" { + const testing = std.testing; + const allocator = testing.allocator; + + var name = try Name.fromString("test.example.com", allocator); + defer name.deinit(); + + const record = createARecord(name, 300, [4]u8{ 192, 168, 1, 1 }); + + var buffer: [256]u8 = undefined; + const encoded_len = try record.encode(&buffer, allocator); + + // Parse it back + const result = try ResourceRecord.parse(buffer[0..encoded_len], buffer[0..encoded_len], allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.A, result.record.rtype); + try testing.expectEqual(@as(u32, 300), result.record.ttl); + try testing.expectEqual([4]u8{ 192, 168, 1, 1 }, result.record.getIPv4().?); +} + +test "parse DNSKEY record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com DNSKEY 256 3 8 TTL=3600 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x30, // TYPE = DNSKEY (48) + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x0E, 0x10, // TTL = 3600 + 0x00, 0x08, // RDLENGTH = 8 + 0x01, 0x00, // flags = 256 (Zone Key) + 0x03, // protocol = 3 + 0x08, // algorithm = 8 (RSASHA256) + 0xDE, 0xAD, 0xBE, 0xEF, // public_key (dummy 4 bytes) + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.DNSKEY, result.record.rtype); + try testing.expectEqual(@as(u32, 3600), result.record.ttl); + + switch (result.record.rdata) { + .dnskey => |dk| { + try testing.expectEqual(@as(u16, 256), dk.flags); + try testing.expectEqual(@as(u8, 3), dk.protocol); + try testing.expectEqual(@as(u8, 8), dk.algorithm); + try testing.expectEqual(@as(usize, 4), dk.public_key.len); + }, + else => try testing.expect(false), + } +} + +test "parse DS record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com DS 12345 8 2 TTL=86400 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x2B, // TYPE = DS (43) + 0x00, 0x01, // CLASS = IN + 0x00, 0x01, 0x51, 0x80, // TTL = 86400 + 0x00, 0x08, // RDLENGTH = 8 + 0x30, 0x39, // key_tag = 12345 + 0x08, // algorithm = 8 + 0x02, // digest_type = 2 (SHA-256) + 0xAB, 0xCD, 0xEF, 0x01, // digest (dummy 4 bytes) + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.DS, result.record.rtype); + + switch (result.record.rdata) { + .ds => |ds| { + try testing.expectEqual(@as(u16, 12345), ds.key_tag); + try testing.expectEqual(@as(u8, 8), ds.algorithm); + try testing.expectEqual(@as(u8, 2), ds.digest_type); + try testing.expectEqual(@as(usize, 4), ds.digest.len); + }, + else => try testing.expect(false), + } +} + +test "parse NSEC3PARAM record" { + const testing = std.testing; + const allocator = testing.allocator; + + // example.com NSEC3PARAM 1 0 10 TTL=0 + const buffer = [_]u8{ + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null terminator + 0x00, 0x33, // TYPE = NSEC3PARAM (51) + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x00, 0x00, // TTL = 0 + 0x00, 0x07, // RDLENGTH = 7 + 0x01, // hash_algorithm = 1 (SHA-1) + 0x00, // flags = 0 + 0x00, 0x0A, // iterations = 10 + 0x02, // salt_length = 2 + 0xAB, 0xCD, // salt + }; + + const result = try ResourceRecord.parse(&buffer, &buffer, allocator); + defer { + var r = result.record; + r.deinit(allocator); + } + + try testing.expectEqual(types.QType.NSEC3PARAM, result.record.rtype); + + switch (result.record.rdata) { + .nsec3param => |np| { + try testing.expectEqual(@as(u8, 1), np.hash_algorithm); + try testing.expectEqual(@as(u8, 0), np.flags); + try testing.expectEqual(@as(u16, 10), np.iterations); + try testing.expectEqual(@as(usize, 2), np.salt.len); + }, + else => try testing.expect(false), + } +} diff --git a/src/dns/types.zig b/src/dns/types.zig new file mode 100644 index 0000000..3219c1d --- /dev/null +++ b/src/dns/types.zig @@ -0,0 +1,223 @@ +const std = @import("std"); + +/// DNS Query/Response types (RFC 1035 + extensions) +pub const QType = enum(u16) { + A = 1, + NS = 2, + MD = 3, // obsolete + MF = 4, // obsolete + CNAME = 5, + SOA = 6, + MB = 7, + MG = 8, + MR = 9, + NULL = 10, + WKS = 11, + PTR = 12, + HINFO = 13, + MINFO = 14, + MX = 15, + TXT = 16, + RP = 17, + AFSDB = 18, + X25 = 19, + ISDN = 20, + RT = 21, + NSAP = 22, + NSAP_PTR = 23, + SIG = 24, + KEY = 25, + PX = 26, + GPOS = 27, + AAAA = 28, + LOC = 29, + NXT = 30, + EID = 31, + NIMLOC = 32, + SRV = 33, + ATMA = 34, + NAPTR = 35, + KX = 36, + CERT = 37, + A6 = 38, + DNAME = 39, + SINK = 40, + OPT = 41, // EDNS + APL = 42, + DS = 43, // DNSSEC + SSHFP = 44, + IPSECKEY = 45, + RRSIG = 46, // DNSSEC + NSEC = 47, // DNSSEC + DNSKEY = 48, // DNSSEC + DHCID = 49, + NSEC3 = 50, + NSEC3PARAM = 51, + TLSA = 52, + SMIMEA = 53, + HIP = 55, + NINFO = 56, + RKEY = 57, + TALINK = 58, + CDS = 59, + CDNSKEY = 60, + OPENPGPKEY = 61, + CSYNC = 62, + ZONEMD = 63, + SVCB = 64, + HTTPS = 65, + SPF = 99, + UINFO = 100, + UID = 101, + GID = 102, + UNSPEC = 103, + NID = 104, + L32 = 105, + L64 = 106, + LP = 107, + EUI48 = 108, + EUI64 = 109, + TKEY = 249, + TSIG = 250, + IXFR = 251, + AXFR = 252, + MAILB = 253, + MAILA = 254, + ANY = 255, + URI = 256, + CAA = 257, + AVC = 258, + DOA = 259, + AMTRELAY = 260, + TA = 32768, + DLV = 32769, + _, + + pub fn toString(self: QType) []const u8 { + return switch (self) { + .A => "A", + .NS => "NS", + .CNAME => "CNAME", + .SOA => "SOA", + .PTR => "PTR", + .MX => "MX", + .TXT => "TXT", + .AAAA => "AAAA", + .SRV => "SRV", + .OPT => "OPT", + // DNSSEC types + .DS => "DS", + .RRSIG => "RRSIG", + .NSEC => "NSEC", + .DNSKEY => "DNSKEY", + .NSEC3 => "NSEC3", + .NSEC3PARAM => "NSEC3PARAM", + .CDS => "CDS", + .CDNSKEY => "CDNSKEY", + .ANY => "ANY", + .HTTPS => "HTTPS", + .SVCB => "SVCB", + else => "UNKNOWN", + }; + } +}; + +/// DNS Query Class +pub const QClass = enum(u16) { + IN = 1, // Internet + CS = 2, // CSNET (obsolete) + CH = 3, // Chaos + HS = 4, // Hesiod + NONE = 254, + ANY = 255, + _, + + pub fn toString(self: QClass) []const u8 { + return switch (self) { + .IN => "IN", + .CH => "CH", + .HS => "HS", + .ANY => "ANY", + else => "UNKNOWN", + }; + } +}; + +/// DNS Response Code +pub const RCode = enum(u4) { + NoError = 0, // No error + FormErr = 1, // Format error + ServFail = 2, // Server failure + NXDomain = 3, // Non-existent domain + NotImp = 4, // Not implemented + Refused = 5, // Query refused + YXDomain = 6, // Name exists when it should not + YXRRSet = 7, // RR set exists when it should not + NXRRSet = 8, // RR set does not exist + NotAuth = 9, // Not authorized + NotZone = 10, // Name not in zone + _, + + pub fn toString(self: RCode) []const u8 { + return switch (self) { + .NoError => "NOERROR", + .FormErr => "FORMERR", + .ServFail => "SERVFAIL", + .NXDomain => "NXDOMAIN", + .NotImp => "NOTIMP", + .Refused => "REFUSED", + else => "UNKNOWN", + }; + } +}; + +/// DNS Operation Code +pub const OpCode = enum(u4) { + Query = 0, // Standard query + IQuery = 1, // Inverse query (obsolete) + Status = 2, // Server status request + Notify = 4, // Zone change notification + Update = 5, // Dynamic update + _, +}; + +/// Maximum DNS name length (RFC 1035: 255 octets wire format, 253 chars in dot notation) +pub const MAX_NAME_LENGTH: usize = 253; + +/// Maximum DNS label length (RFC 1035: 63 octets, 6-bit length field) +pub const MAX_LABEL_LENGTH: usize = 63; + +/// Standard DNS UDP message size (RFC 1035 Section 4.2.1) +pub const DNS_UDP_SIZE: usize = 512; + +/// EDNS0 default UDP payload size (RFC 6891, commonly 4096 for modern resolvers) +pub const EDNS_DEFAULT_SIZE: usize = 4096; + +/// DNS header size in bytes (RFC 1035 Section 4.1.1: ID + flags + 4 counts) +pub const DNS_HEADER_SIZE: usize = 12; + +/// Compression pointer mask - top 2 bits set indicates pointer (RFC 1035 Section 4.1.4) +pub const COMPRESSION_POINTER_MASK: u8 = 0xC0; + +/// Maximum compression pointer offset (14-bit value, RFC 1035 Section 4.1.4) +pub const MAX_COMPRESSION_OFFSET: u16 = 0x3FFF; + +test "QType values" { + const testing = std.testing; + try testing.expectEqual(@as(u16, 1), @intFromEnum(QType.A)); + try testing.expectEqual(@as(u16, 28), @intFromEnum(QType.AAAA)); + try testing.expectEqual(@as(u16, 5), @intFromEnum(QType.CNAME)); + try testing.expectEqual(@as(u16, 41), @intFromEnum(QType.OPT)); +} + +test "QClass values" { + const testing = std.testing; + try testing.expectEqual(@as(u16, 1), @intFromEnum(QClass.IN)); + try testing.expectEqual(@as(u16, 255), @intFromEnum(QClass.ANY)); +} + +test "RCode values" { + const testing = std.testing; + try testing.expectEqual(@as(u4, 0), @intFromEnum(RCode.NoError)); + try testing.expectEqual(@as(u4, 3), @intFromEnum(RCode.NXDomain)); +} diff --git a/src/events.zig b/src/events.zig new file mode 100644 index 0000000..7465452 --- /dev/null +++ b/src/events.zig @@ -0,0 +1,84 @@ +const std = @import("std"); +const posix = std.posix; + +/// Event signaling using Linux eventfd +/// Allows efficient cross-thread event notification without polling +pub const EventFd = struct { + fd: posix.fd_t, + + pub const Error = error{ + EventFdCreateFailed, + }; + + /// Create a new eventfd + pub fn init() Error!EventFd { + // EFD_NONBLOCK | EFD_CLOEXEC + const flags: u32 = 0x800 | 0x80000; + const fd = std.os.linux.eventfd(0, flags); + if (@as(isize, @bitCast(fd)) < 0) { + return error.EventFdCreateFailed; + } + return .{ .fd = @intCast(fd) }; + } + + pub fn deinit(self: *EventFd) void { + posix.close(self.fd); + } + + /// Signal the event (non-blocking) + pub fn signal(self: *EventFd) void { + const val: u64 = 1; + _ = posix.write(self.fd, std.mem.asBytes(&val)) catch {}; + } + + /// Consume the event (non-blocking, returns true if event was pending) + pub fn consume(self: *EventFd) bool { + var val: u64 = undefined; + const n = posix.read(self.fd, std.mem.asBytes(&val)) catch return false; + return n == 8; + } + + /// Get the file descriptor for polling + pub fn getFd(self: *EventFd) posix.fd_t { + return self.fd; + } +}; + +/// Global eventfd for denylist reload signaling +var denylist_event: ?EventFd = null; + +/// Initialize the global denylist event +pub fn initDenylistEvent() !void { + denylist_event = try EventFd.init(); +} + +/// Deinitialize the global denylist event +pub fn deinitDenylistEvent() void { + if (denylist_event) |*ev| { + ev.deinit(); + denylist_event = null; + } +} + +/// Signal that denylist needs reload +pub fn signalDenylistReload() void { + if (denylist_event) |*ev| { + ev.signal(); + } +} + +/// Get the denylist event fd for polling (returns -1 if not initialized) +pub fn getDenylistEventFd() posix.fd_t { + if (denylist_event) |*ev| { + return ev.getFd(); + } + return -1; +} + +/// Consume the denylist reload event +pub fn consumeDenylistReload() bool { + if (denylist_event) |*ev| { + return ev.consume(); + } + return false; +} diff --git a/src/filter/denylist.zig b/src/filter/denylist.zig new file mode 100644 index 0000000..06adb07 --- /dev/null +++ b/src/filter/denylist.zig @@ -0,0 +1,666 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const Database = @import("../storage/db.zig").Database; +const handler = @import("../server/handler.zig"); +const types = @import("../dns/types.zig"); +const toLower = @import("../util.zig").toLower; + +/// Maximum domain buffer size (RFC 1035: 253 chars + null terminator) +const MAX_DOMAIN_BUF_SIZE = types.MAX_NAME_LENGTH + 1; + +/// Domain info with attribution (Gap 2) +const DomainInfo = struct { + groups: u64, // Bitmask of groups that have this domain denied + source_id: i64, // First source_id that contributed this domain +}; + +/// Rule info with attribution (Gap 2) +const RuleInfo = struct { + groups: u64, // Bitmask of groups this rule applies to + rule_id: i64, // The rule ID in the database +}; + +/// Thread-safe denylist wrapper using RwLock for concurrent access +/// Allows multiple readers during queries, exclusive access during reload +pub const ThreadSafeDenylist = struct { + inner: *Denylist, + lock: std.Thread.RwLock, + allocator: Allocator, + + pub fn init(denylist: *Denylist, allocator: Allocator) ThreadSafeDenylist { + return .{ + .inner = denylist, + .lock = .{}, + .allocator = allocator, + }; + } + + /// Check if domain is denied with full attribution (acquires read lock) + /// Gap 4: Takes groups bitmask for many-to-many support + pub fn checkWithMask(self: *ThreadSafeDenylist, domain: []const u8, groups_mask: u64) handler.DenyResult { + self.lock.lockShared(); + defer self.lock.unlockShared(); + return self.inner.checkWithMask(domain, groups_mask); + } + + /// Legacy: Check if domain is denied for a single group (acquires read lock) + pub fn check(self: *ThreadSafeDenylist, domain: []const u8, group_id: u32) handler.DenyResult { + self.lock.lockShared(); + defer self.lock.unlockShared(); + return self.inner.check(domain, group_id); + } + + /// Legacy: Check if domain is denied (acquires read lock) + pub fn isDenied(self: *ThreadSafeDenylist, domain: []const u8, group_id: u32) bool { + return self.check(domain, group_id).denied; + } + + /// Get all groups for a client as a bitmask (Gap 4: many-to-many) + pub fn getGroupsForClient(self: *ThreadSafeDenylist, client_ip: []const u8) u64 { + self.lock.lockShared(); + defer self.lock.unlockShared(); + return self.inner.getGroupsForClient(client_ip); + } + + /// Legacy: Get primary group for client (acquires read lock) + pub fn getGroupForClient(self: *ThreadSafeDenylist, client_ip: []const u8) u32 { + self.lock.lockShared(); + defer self.lock.unlockShared(); + return self.inner.getGroupForClient(client_ip); + } + + /// Atomically swap the inner denylist (acquires write lock) + /// Returns the old denylist for cleanup by the caller + pub fn swap(self: *ThreadSafeDenylist, new_denylist: Denylist) Denylist { + self.lock.lock(); + defer self.lock.unlock(); + + const old = self.inner.*; + self.inner.* = new_denylist; + return old; + } + + /// Convert to handler-compatible interface + pub fn toHandlerDenylist(self: *ThreadSafeDenylist) handler.Denylist { + return handler.Denylist{ + .context = self, + .checkFn = checkWithMaskWrapper, + .getGroupsFn = getGroupsWrapper, + }; + } + + fn checkWithMaskWrapper(ctx: *anyopaque, domain: []const u8, groups_mask: u64) handler.DenyResult { + const self: *ThreadSafeDenylist = @ptrCast(@alignCast(ctx)); + return self.checkWithMask(domain, groups_mask); + } + + fn getGroupsWrapper(ctx: *anyopaque, client_ip: []const u8) u64 { + const self: *ThreadSafeDenylist = @ptrCast(@alignCast(ctx)); + return self.getGroupsForClient(client_ip); + } +}; + +/// Denylist with parent-walking lookup +/// Uses arena allocator for fast bulk allocation and instant deallocation +pub const Denylist = struct { + /// Maps domain -> DomainInfo with groups and source attribution + denied_domains: std.StringHashMapUnmanaged(DomainInfo), + /// Maps client IP -> groups bitmask (Gap 4: supports many-to-many) + client_groups: std.StringHashMapUnmanaged(u64), + /// Custom allow rules per group: domain -> RuleInfo with groups and rule_id + allow_rules: std.StringHashMapUnmanaged(RuleInfo), + /// Custom deny rules per group: domain -> RuleInfo with groups and rule_id + deny_rules: std.StringHashMapUnmanaged(RuleInfo), + /// Allowlist domains from external sources (Gap 5): domain -> DomainInfo + allowlist_domains: std.StringHashMapUnmanaged(DomainInfo), + /// Arena for fast allocations - all memory freed at once on deinit + arena: std.heap.ArenaAllocator, + + pub fn init(backing_allocator: Allocator) Denylist { + return Denylist{ + .denied_domains = std.StringHashMapUnmanaged(DomainInfo){}, + .client_groups = std.StringHashMapUnmanaged(u64){}, + .allow_rules = std.StringHashMapUnmanaged(RuleInfo){}, + .deny_rules = std.StringHashMapUnmanaged(RuleInfo){}, + .allowlist_domains = std.StringHashMapUnmanaged(DomainInfo){}, + .arena = std.heap.ArenaAllocator.init(backing_allocator), + }; + } + + fn groupMask(group_id: u32) u64 { + if (group_id >= 64) return 0; + return @as(u64, 1) << @intCast(group_id); + } + + fn allocator(self: *Denylist) Allocator { + return self.arena.allocator(); + } + + pub fn deinit(self: *Denylist) void { + // Arena frees everything at once - no need to iterate + self.arena.deinit(); + } + + /// Check if a domain is denied for given groups (bitmask), with full attribution + /// Gap 4: groups_mask allows checking against multiple groups at once + pub fn checkWithMask(self: *Denylist, domain: []const u8, groups_mask: u64) handler.DenyResult { + // Normalize domain to lowercase + var domain_lower_buf: [MAX_DOMAIN_BUF_SIZE]u8 = undefined; + const domain_lower = toLower(domain, &domain_lower_buf) orelse return .{ + .denied = false, + .list_id = null, + .rule_id = null, + .is_rule_allow = false, + }; + + // Combine with default group (group 0 applies to everyone) + const default_mask = groupMask(0); + const combined_mask = groups_mask | default_mask; + + // 1. Check custom allow rules first (exact match only for allow) + if (self.allow_rules.get(domain_lower)) |rule_info| { + if ((rule_info.groups & combined_mask) != 0) { + return .{ + .denied = false, + .list_id = null, + .rule_id = rule_info.rule_id, + .is_rule_allow = true, + }; + } + } + + // 2. Check allowlist domains from external sources (Gap 5) + if (self.checkAllowlistWithMask(domain_lower, combined_mask)) |_| { + return .{ + .denied = false, + .list_id = null, + .rule_id = null, + .is_rule_allow = true, // Treated as an allow override + }; + } + + // 3. Check custom deny rules (with parent walking) + if (self.checkDenyRulesWithMask(domain_lower, combined_mask)) |rule_info| { + return .{ + .denied = true, + .list_id = null, + .rule_id = rule_info.rule_id, + .is_rule_allow = false, + }; + } + + // 4. Check denylists (with parent walking) + if (self.checkDenylistWithMask(domain_lower, combined_mask)) |domain_info| { + return .{ + .denied = true, + .list_id = domain_info.source_id, + .rule_id = null, + .is_rule_allow = false, + }; + } + + return .{ + .denied = false, + .list_id = null, + .rule_id = null, + .is_rule_allow = false, + }; + } + + /// Check if a domain is denied for a given group, with full attribution + /// Legacy API - converts single group_id to bitmask + pub fn check(self: *Denylist, domain: []const u8, group_id: u32) handler.DenyResult { + return self.checkWithMask(domain, groupMask(group_id)); + } + + /// Legacy: Check if a domain is denied for a given group + pub fn isDenied(self: *Denylist, domain: []const u8, group_id: u32) bool { + return self.check(domain, group_id).denied; + } + + fn checkDenyRulesWithMask(self: *Denylist, domain: []const u8, mask: u64) ?RuleInfo { + var d: []const u8 = domain; + while (true) { + if (self.deny_rules.get(d)) |rule_info| { + if ((rule_info.groups & mask) != 0) { + return rule_info; + } + } + + // Move to parent domain + if (std.mem.indexOfScalar(u8, d, '.')) |idx| { + d = d[idx + 1 ..]; + } else { + return null; + } + } + } + + fn checkDenylistWithMask(self: *Denylist, domain: []const u8, mask: u64) ?DomainInfo { + var d: []const u8 = domain; + while (true) { + if (self.denied_domains.get(d)) |domain_info| { + if ((domain_info.groups & mask) != 0) { + return domain_info; + } + } + + // Move to parent domain + if (std.mem.indexOfScalar(u8, d, '.')) |idx| { + d = d[idx + 1 ..]; + } else { + return null; + } + } + } + + /// Check allowlist domains from external sources (Gap 5) + /// Uses exact match only (no parent walking - allowlists are typically specific) + fn checkAllowlistWithMask(self: *Denylist, domain: []const u8, mask: u64) ?DomainInfo { + if (self.allowlist_domains.get(domain)) |domain_info| { + if ((domain_info.groups & mask) != 0) { + return domain_info; + } + } + return null; + } + + /// Get all groups for a client IP as a bitmask (Gap 4: many-to-many) + pub fn getGroupsForClient(self: *Denylist, client_ip: []const u8) u64 { + return self.client_groups.get(client_ip) orelse groupMask(0); + } + + /// Legacy: Get the primary group ID for a client IP (returns first set bit) + pub fn getGroupForClient(self: *Denylist, client_ip: []const u8) u32 { + const mask = self.getGroupsForClient(client_ip); + // Return the lowest set group bit (or 0 if only default group) + if (mask == groupMask(0)) return 0; + var bit: u6 = 0; + while (bit < 64) : (bit += 1) { + if ((mask & (@as(u64, 1) << bit)) != 0) return bit; + } + return 0; + } + + /// Add a denied domain for a group with source attribution + pub fn addDeniedDomainWithSource(self: *Denylist, domain: []const u8, group_id: u32, source_id: i64) !void { + var domain_lower_buf: [MAX_DOMAIN_BUF_SIZE]u8 = undefined; + const domain_lower = toLower(domain, &domain_lower_buf) orelse return; + + const alloc = self.allocator(); + const gop = try self.denied_domains.getOrPut(alloc, domain_lower); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, domain_lower); + gop.value_ptr.* = .{ .groups = 0, .source_id = source_id }; + } + gop.value_ptr.groups |= groupMask(group_id); + } + + /// Add a denied domain for a group (legacy - no source attribution) + pub fn addDeniedDomain(self: *Denylist, domain: []const u8, group_id: u32) !void { + return self.addDeniedDomainWithSource(domain, group_id, 0); + } + + /// Add an allowlist domain for a group with source attribution (Gap 5) + pub fn addAllowlistDomainWithSource(self: *Denylist, domain: []const u8, group_id: u32, source_id: i64) !void { + var domain_lower_buf: [MAX_DOMAIN_BUF_SIZE]u8 = undefined; + const domain_lower = toLower(domain, &domain_lower_buf) orelse return; + + const alloc = self.allocator(); + const gop = try self.allowlist_domains.getOrPut(alloc, domain_lower); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, domain_lower); + gop.value_ptr.* = .{ .groups = 0, .source_id = source_id }; + } + gop.value_ptr.groups |= groupMask(group_id); + } + + /// Add a group to a client's groups bitmask (Gap 4: many-to-many) + pub fn addClientToGroup(self: *Denylist, client_ip: []const u8, group_id: u32) !void { + const alloc = self.allocator(); + const gop = try self.client_groups.getOrPut(alloc, client_ip); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, client_ip); + gop.value_ptr.* = 0; + } + gop.value_ptr.* |= groupMask(group_id); + } + + /// Legacy: Set client's group (replaces all groups with single group) + pub fn setClientGroup(self: *Denylist, client_ip: []const u8, group_id: u32) !void { + const alloc = self.allocator(); + const gop = try self.client_groups.getOrPut(alloc, client_ip); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, client_ip); + } + gop.value_ptr.* = groupMask(group_id); + } + + /// Add an allow rule with rule_id attribution + pub fn addAllowRuleWithId(self: *Denylist, domain: []const u8, group_id: u32, rule_id: i64) !void { + var domain_lower_buf: [MAX_DOMAIN_BUF_SIZE]u8 = undefined; + const domain_lower = toLower(domain, &domain_lower_buf) orelse return; + + const alloc = self.allocator(); + const gop = try self.allow_rules.getOrPut(alloc, domain_lower); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, domain_lower); + gop.value_ptr.* = .{ .groups = 0, .rule_id = rule_id }; + } + gop.value_ptr.groups |= groupMask(group_id); + } + + /// Add an allow rule (legacy - no rule_id attribution) + pub fn addAllowRule(self: *Denylist, domain: []const u8, group_id: u32) !void { + return self.addAllowRuleWithId(domain, group_id, 0); + } + + /// Add a deny rule with rule_id attribution + pub fn addDenyRuleWithId(self: *Denylist, domain: []const u8, group_id: u32, rule_id: i64) !void { + var domain_lower_buf: [MAX_DOMAIN_BUF_SIZE]u8 = undefined; + const domain_lower = toLower(domain, &domain_lower_buf) orelse return; + + const alloc = self.allocator(); + const gop = try self.deny_rules.getOrPut(alloc, domain_lower); + if (!gop.found_existing) { + gop.key_ptr.* = try alloc.dupe(u8, domain_lower); + gop.value_ptr.* = .{ .groups = 0, .rule_id = rule_id }; + } + gop.value_ptr.groups |= groupMask(group_id); + } + + /// Add a deny rule (legacy - no rule_id attribution) + pub fn addDenyRule(self: *Denylist, domain: []const u8, group_id: u32) !void { + return self.addDenyRuleWithId(domain, group_id, 0); + } + + /// Convert to handler-compatible Denylist interface + pub fn toHandlerDenylist(self: *Denylist) handler.Denylist { + return handler.Denylist{ + .context = self, + .checkFn = checkWrapper, + .getGroupFn = getGroupWrapper, + }; + } + + fn checkWrapper(ctx: *anyopaque, domain: []const u8, group_id: u32) handler.DenyResult { + const self: *Denylist = @ptrCast(@alignCast(ctx)); + return self.check(domain, group_id); + } + + fn getGroupWrapper(ctx: *anyopaque, client_ip: []const u8) u32 { + const self: *Denylist = @ptrCast(@alignCast(ctx)); + return self.getGroupForClient(client_ip); + } +}; + +/// Load denylist from database +pub fn loadFromDatabase(db: *Database, allocator: Allocator) !Denylist { + var denylist = Denylist.init(allocator); + errdefer denylist.deinit(); + + // Count domains first to preallocate HashMap (avoids repeated resizing) + var count_stmt = try db.prepare( + \\SELECT COUNT(DISTINCT bd.domain) + \\FROM denylist_domains bd + \\JOIN group_sources gs ON bd.source_id = gs.source_id + \\JOIN denylist_sources bs ON bd.source_id = bs.id + \\WHERE bs.enabled = 1 + ); + defer count_stmt.finalize(); + + var domain_count: u32 = 0; + if (try count_stmt.step()) { + domain_count = @intCast(count_stmt.getInt(0)); + } + + // Preallocate capacity to avoid resizing during inserts + if (domain_count > 0) { + try denylist.denied_domains.ensureTotalCapacity(denylist.allocator(), domain_count); + } + + // Load denied domains grouped by group, with source attribution (type=0 denylists) + var stmt = try db.prepare( + \\SELECT bd.domain, gs.group_id, bd.source_id + \\FROM denylist_domains bd + \\JOIN group_sources gs ON bd.source_id = gs.source_id + \\JOIN denylist_sources bs ON bd.source_id = bs.id + \\WHERE bs.enabled = 1 AND (bs.type = 0 OR bs.type IS NULL) + ); + defer stmt.finalize(); + + while (try stmt.step()) { + const domain = stmt.getText(0) orelse continue; + const group_id: u32 = @intCast(stmt.getInt(1)); + const source_id = stmt.getInt(2); + try denylist.addDeniedDomainWithSource(domain, group_id, source_id); + } + + // Load allowlist domains from external sources (Gap 5: type=1 allowlists) + var allowlist_stmt = try db.prepare( + \\SELECT ad.domain, gs.group_id, ad.source_id + \\FROM allowlist_domains ad + \\JOIN group_sources gs ON ad.source_id = gs.source_id + \\JOIN denylist_sources bs ON ad.source_id = bs.id + \\WHERE bs.enabled = 1 AND bs.type = 1 + ); + defer allowlist_stmt.finalize(); + + while (try allowlist_stmt.step()) { + const domain = allowlist_stmt.getText(0) orelse continue; + const group_id: u32 = @intCast(allowlist_stmt.getInt(1)); + const source_id = allowlist_stmt.getInt(2); + try denylist.addAllowlistDomainWithSource(domain, group_id, source_id); + } + + // Load client groups from legacy clients.group_id column + var client_stmt = try db.prepare("SELECT ip, group_id FROM clients WHERE group_id IS NOT NULL"); + defer client_stmt.finalize(); + + while (try client_stmt.step()) { + const ip = client_stmt.getText(0) orelse continue; + const group_id: u32 = @intCast(client_stmt.getInt(1)); + try denylist.addClientToGroup(ip, group_id); + } + + // Load client groups from client_groups junction table (Gap 4: many-to-many) + var junction_stmt = try db.prepare( + \\SELECT c.ip, cg.group_id + \\FROM client_groups cg + \\JOIN clients c ON cg.client_id = c.id + ); + defer junction_stmt.finalize(); + + while (try junction_stmt.step()) { + const ip = junction_stmt.getText(0) orelse continue; + const group_id: u32 = @intCast(junction_stmt.getInt(1)); + try denylist.addClientToGroup(ip, group_id); + } + + // Load custom rules with rule_id attribution + var rule_stmt = try db.prepare("SELECT id, domain, group_id, action FROM rules"); + defer rule_stmt.finalize(); + + while (try rule_stmt.step()) { + const rule_id = rule_stmt.getInt(0); + const domain = rule_stmt.getText(1) orelse continue; + const group_id: u32 = if (rule_stmt.isNull(2)) 0 else @intCast(rule_stmt.getInt(2)); + const action = rule_stmt.getText(3) orelse continue; + + if (std.mem.eql(u8, action, "allow")) { + try denylist.addAllowRuleWithId(domain, group_id, rule_id); + } else if (std.mem.eql(u8, action, "deny")) { + try denylist.addDenyRuleWithId(domain, group_id, rule_id); + } + } + + return denylist; +} + +test "Denylist exact match" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.addDeniedDomain("ads.example.com", 0); + + try testing.expect(denylist.isDenied("ads.example.com", 0)); + try testing.expect(!denylist.isDenied("example.com", 0)); + try testing.expect(!denylist.isDenied("other.example.com", 0)); +} + +test "Denylist parent walking" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.addDeniedDomain("doubleclick.net", 0); + + // Subdomain should be denied + try testing.expect(denylist.isDenied("www.doubleclick.net", 0)); + try testing.expect(denylist.isDenied("ads.doubleclick.net", 0)); + try testing.expect(denylist.isDenied("sub.ads.doubleclick.net", 0)); + + // Parent should not be denied + try testing.expect(!denylist.isDenied("net", 0)); +} + +test "Denylist allow rules" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.addDeniedDomain("example.com", 0); + try denylist.addAllowRule("allowed.example.com", 0); + + try testing.expect(denylist.isDenied("denied.example.com", 0)); + try testing.expect(!denylist.isDenied("allowed.example.com", 0)); +} + +test "Denylist group isolation" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.addDeniedDomain("denied.com", 1); + + // Group 1 should see it denied + try testing.expect(denylist.isDenied("denied.com", 1)); + + // Group 0 and 2 should not see it denied + try testing.expect(!denylist.isDenied("denied.com", 0)); + try testing.expect(!denylist.isDenied("denied.com", 2)); +} + +test "Denylist case insensitive" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.addDeniedDomain("EXAMPLE.COM", 0); + + try testing.expect(denylist.isDenied("example.com", 0)); + try testing.expect(denylist.isDenied("EXAMPLE.COM", 0)); + try testing.expect(denylist.isDenied("Example.Com", 0)); +} + +test "Denylist scale to 1 million domains" { + const testing = std.testing; + const allocator = testing.allocator; + + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + // Preallocate for better performance + try denylist.denied_domains.ensureTotalCapacity(denylist.allocator(), 1_000_000); + + const start_load = std.time.milliTimestamp(); + + // Insert 1 million domains + var domain_buf: [64]u8 = undefined; + for (0..1_000_000) |i| { + const domain = std.fmt.bufPrint(&domain_buf, "domain{d}.example.com", .{i}) catch continue; + try denylist.addDeniedDomain(domain, 0); + } + + const load_time = std.time.milliTimestamp() - start_load; + + // Verify count + try testing.expectEqual(@as(usize, 1_000_000), denylist.denied_domains.count()); + + // Benchmark lookups (1000 random lookups) + const start_lookup = std.time.milliTimestamp(); + var found: usize = 0; + for (0..1000) |i| { + const domain = std.fmt.bufPrint(&domain_buf, "domain{d}.example.com", .{i * 1000}) catch continue; + if (denylist.isDenied(domain, 0)) { + found += 1; + } + } + const lookup_time = std.time.milliTimestamp() - start_lookup; + + // All lookups should find denied domains + try testing.expectEqual(@as(usize, 1000), found); + + // Test parent-walking lookup (subdomain of denied domain) + try denylist.addDeniedDomain("denied-parent.com", 0); + try testing.expect(denylist.isDenied("sub.denied-parent.com", 0)); + try testing.expect(denylist.isDenied("deep.sub.denied-parent.com", 0)); + + // Test non-denied domain + try testing.expect(!denylist.isDenied("not-denied.com", 0)); + + // Print performance metrics + std.debug.print("\n[Denylist Scale Test]\n", .{}); + std.debug.print(" Domains: 1,000,000\n", .{}); + std.debug.print(" Load time: {d}ms\n", .{load_time}); + std.debug.print(" 1000 lookups: {d}ms ({d}us/lookup)\n", .{ lookup_time, lookup_time }); +} + +test "Denylist load time scales linearly" { + const testing = std.testing; + const allocator = testing.allocator; + + const sizes = [_]usize{ 10_000, 50_000, 100_000, 200_000 }; + var times: [sizes.len]i64 = undefined; + + for (sizes, 0..) |size, idx| { + var denylist = Denylist.init(allocator); + defer denylist.deinit(); + + try denylist.denied_domains.ensureTotalCapacity(denylist.allocator(), @intCast(size)); + + const start = std.time.milliTimestamp(); + + var domain_buf: [64]u8 = undefined; + for (0..size) |i| { + const domain = std.fmt.bufPrint(&domain_buf, "d{d}.test.com", .{i}) catch continue; + try denylist.addDeniedDomain(domain, 0); + } + + times[idx] = std.time.milliTimestamp() - start; + try testing.expectEqual(size, denylist.denied_domains.count()); + } + + std.debug.print("\n[Denylist Linear Scaling]\n", .{}); + for (sizes, 0..) |size, idx| { + const rate = if (times[idx] > 0) @divTrunc(@as(i64, @intCast(size)), times[idx]) else 0; + std.debug.print(" {d}: {d}ms ({d}k domains/sec)\n", .{ size, times[idx], rate }); + } + + // Verify roughly linear scaling (4x size should be ~4x time, allow 2x-8x range) + if (times[1] > 0) { + const ratio = @divTrunc(times[3] * 100, times[1]); + try testing.expect(ratio >= 200 and ratio <= 800); + } +} diff --git a/src/filter/fetcher.zig b/src/filter/fetcher.zig new file mode 100644 index 0000000..733d5d0 --- /dev/null +++ b/src/filter/fetcher.zig @@ -0,0 +1,549 @@ +const std = @import("std"); +const http = std.http; +const Allocator = std.mem.Allocator; +const Writer = std.Io.Writer; +const Database = @import("../storage/db.zig").Database; +const schema = @import("../storage/schema.zig"); +const events = @import("../events.zig"); + +const DenylistStatus = schema.DenylistStatus; + +/// Maximum output size for denylist fetch (100MB) +const MAX_DENYLIST_SIZE = 100 * 1024 * 1024; + +/// Denylist format types +pub const DenylistFormat = enum { + hosts, // 0.0.0.0 domain.com or 127.0.0.1 domain.com + domains, // one domain per line + adblock, // ||domain.com^ + + /// Detect format from content sample + pub fn detect(content: []const u8) DenylistFormat { + var lines = std.mem.splitScalar(u8, content, '\n'); + var hosts_count: usize = 0; + var adblock_count: usize = 0; + var checked: usize = 0; + + while (lines.next()) |line| { + const trimmed = std.mem.trim(u8, line, " \t\r"); + if (trimmed.len == 0 or trimmed[0] == '#') continue; + + if (std.mem.startsWith(u8, trimmed, "||") and std.mem.endsWith(u8, trimmed, "^")) { + adblock_count += 1; + } else if (std.mem.startsWith(u8, trimmed, "0.0.0.0 ") or + std.mem.startsWith(u8, trimmed, "127.0.0.1 ")) + { + hosts_count += 1; + } + + checked += 1; + if (checked >= 50) break; + } + + if (adblock_count > hosts_count) return .adblock; + if (hosts_count > 0) return .hosts; + return .domains; + } +}; + +/// Result of parsing a denylist +pub const ParseResult = struct { + domains: std.ArrayListUnmanaged([]const u8), + invalid_count: usize, +}; + +/// Parse denylist content and return domains plus count of invalid lines +pub fn parseDenylist(content: []const u8, format: DenylistFormat, allocator: Allocator) !ParseResult { + var domains = std.ArrayListUnmanaged([]const u8){}; + errdefer { + for (domains.items) |d| allocator.free(d); + domains.deinit(allocator); + } + + var invalid_count: usize = 0; + var lines = std.mem.splitScalar(u8, content, '\n'); + + while (lines.next()) |line| { + const trimmed = std.mem.trim(u8, line, " \t\r"); + if (trimmed.len == 0 or trimmed[0] == '#' or trimmed[0] == '!') continue; + + const domain = switch (format) { + .hosts => parseHostsLine(trimmed), + .domains => parseDomainLine(trimmed), + .adblock => parseAdblockLine(trimmed), + }; + + if (domain) |d| { + if (isValidDomain(d)) { + const copy = try allocator.dupe(u8, d); + try domains.append(allocator, copy); + } else { + invalid_count += 1; + } + } else { + // Line matched format but couldn't be parsed (e.g., wrong IP in hosts) + invalid_count += 1; + } + } + + return .{ .domains = domains, .invalid_count = invalid_count }; +} + +fn parseHostsLine(line: []const u8) ?[]const u8 { + // Format: 0.0.0.0 domain.com or 127.0.0.1 domain.com + var parts = std.mem.splitAny(u8, line, " \t"); + + const ip = parts.first(); + if (!std.mem.eql(u8, ip, "0.0.0.0") and !std.mem.eql(u8, ip, "127.0.0.1")) { + return null; + } + + const domain = parts.next() orelse return null; + + // Skip localhost entries + if (std.mem.eql(u8, domain, "localhost") or + std.mem.eql(u8, domain, "localhost.localdomain") or + std.mem.eql(u8, domain, "local") or + std.mem.eql(u8, domain, "broadcasthost")) + { + return null; + } + + return domain; +} + +fn parseDomainLine(line: []const u8) ?[]const u8 { + // Just the domain, skip comments + if (line.len == 0) return null; + + // Check for inline comments + var domain = line; + if (std.mem.indexOf(u8, line, "#")) |idx| { + domain = std.mem.trim(u8, line[0..idx], " \t"); + } + + if (domain.len == 0) return null; + + // Handle wildcard format: *.domain.com -> domain.com + if (std.mem.startsWith(u8, domain, "*.")) { + domain = domain[2..]; + } + + return if (domain.len > 0) domain else null; +} + +fn parseAdblockLine(line: []const u8) ?[]const u8 { + // Format: ||domain.com^ or ||domain.com^$... + if (!std.mem.startsWith(u8, line, "||")) return null; + + var end = line.len; + + // Find the end marker (^ or $) + if (std.mem.indexOf(u8, line[2..], "^")) |idx| { + end = idx + 2; + } else if (std.mem.indexOf(u8, line[2..], "$")) |idx| { + end = idx + 2; + } + + const domain = line[2..end]; + return if (domain.len > 0) domain else null; +} + +fn isValidDomain(domain: []const u8) bool { + if (domain.len == 0 or domain.len > 253) return false; + + // Must contain at least one dot (except for special cases) + if (std.mem.indexOf(u8, domain, ".") == null) return false; + + // Check for valid characters + for (domain) |c| { + if (!std.ascii.isAlphanumeric(c) and c != '.' and c != '-' and c != '_') { + return false; + } + } + + // Don't start or end with hyphen or dot + if (domain[0] == '-' or domain[0] == '.') return false; + if (domain[domain.len - 1] == '-' or domain[domain.len - 1] == '.') return false; + + return true; +} + +/// Denylist fetcher for downloading and updating denylists +pub const DenylistFetcher = struct { + db: *Database, + allocator: Allocator, + + pub fn init(db: *Database, allocator: Allocator) DenylistFetcher { + return DenylistFetcher{ + .db = db, + .allocator = allocator, + }; + } + + /// Fetch and update a denylist from URL + pub fn fetchAndUpdate(self: *DenylistFetcher, source_id: i64, url: []const u8) !usize { + // Fetch content via HTTP + const content = self.httpGet(url) catch |err| { + // Fetch failed - check if we have cached data + const has_cached = self.hasCachedDomains(source_id); + const status: DenylistStatus = if (has_cached) .cached else .failed; + self.updateStatus(source_id, status); + return err; + }; + defer self.allocator.free(content); + + // Compute content hash + var hasher = std.hash.XxHash3.init(0); + hasher.update(content); + const hash = hasher.final(); + var hash_str: [16]u8 = undefined; + _ = std.fmt.bufPrint(&hash_str, "{x:0>16}", .{hash}) catch unreachable; + + // Check if content changed + const old_hash = self.getContentHash(source_id); + if (old_hash) |h| { + if (std.mem.eql(u8, &hash_str, h)) { + // Content unchanged + self.updateStatusAndHash(source_id, .unchanged, &hash_str); + return self.getDomainCount(source_id); + } + } + + // Detect format + const format = DenylistFormat.detect(content); + + // Parse domains + var result = try parseDenylist(content, format, self.allocator); + defer { + for (result.domains.items) |d| self.allocator.free(d); + result.domains.deinit(self.allocator); + } + + // Update database with new domains and invalid count + const count = try self.updateDatabase(source_id, result.domains.items, result.invalid_count, &hash_str); + + return count; + } + + fn hasCachedDomains(self: *DenylistFetcher, source_id: i64) bool { + var stmt = self.db.prepare("SELECT domain_count FROM denylist_sources WHERE id = ?") catch return false; + defer stmt.finalize(); + stmt.bindInt(1, source_id) catch return false; + if (stmt.step() catch false) { + return stmt.getInt(0) > 0; + } + return false; + } + + fn getContentHash(self: *DenylistFetcher, source_id: i64) ?[]const u8 { + var stmt = self.db.prepare("SELECT content_hash FROM denylist_sources WHERE id = ?") catch return null; + defer stmt.finalize(); + stmt.bindInt(1, source_id) catch return null; + if (stmt.step() catch false) { + return stmt.getText(0); + } + return null; + } + + fn getDomainCount(self: *DenylistFetcher, source_id: i64) usize { + var stmt = self.db.prepare("SELECT domain_count FROM denylist_sources WHERE id = ?") catch return 0; + defer stmt.finalize(); + stmt.bindInt(1, source_id) catch return 0; + if (stmt.step() catch false) { + const count = stmt.getInt(0); + return if (count >= 0) @intCast(count) else 0; + } + return 0; + } + + fn updateStatus(self: *DenylistFetcher, source_id: i64, status: DenylistStatus) void { + var stmt = self.db.prepare("UPDATE denylist_sources SET status = ? WHERE id = ?") catch return; + defer stmt.finalize(); + stmt.bindInt(1, @intFromEnum(status)) catch return; + stmt.bindInt(2, source_id) catch return; + _ = stmt.step() catch {}; + } + + fn updateStatusAndHash(self: *DenylistFetcher, source_id: i64, status: DenylistStatus, hash: []const u8) void { + var stmt = self.db.prepare( + "UPDATE denylist_sources SET status = ?, content_hash = ?, last_updated = strftime('%s', 'now') WHERE id = ?", + ) catch return; + defer stmt.finalize(); + stmt.bindInt(1, @intFromEnum(status)) catch return; + stmt.bindText(2, hash) catch return; + stmt.bindInt(3, source_id) catch return; + _ = stmt.step() catch {}; + } + + fn updateDatabase(self: *DenylistFetcher, source_id: i64, domains: []const []const u8, invalid_count: usize, hash: []const u8) !usize { + // Get source type to determine which table to use (Gap 5) + const source_type = self.getSourceType(source_id); + const is_allowlist = source_type == .allowlist; + + // Start transaction with RAII - auto-rollback if not committed + var tx = try self.db.begin(); + defer tx.deinit(); + + // Delete existing domains for this source from appropriate table + if (is_allowlist) { + var delete_stmt = try self.db.prepare("DELETE FROM allowlist_domains WHERE source_id = ?"); + defer delete_stmt.finalize(); + try delete_stmt.bindInt(1, source_id); + _ = try delete_stmt.step(); + } else { + var delete_stmt = try self.db.prepare("DELETE FROM denylist_domains WHERE source_id = ?"); + defer delete_stmt.finalize(); + try delete_stmt.bindInt(1, source_id); + _ = try delete_stmt.step(); + } + + // Insert new domains into appropriate table + const insert_sql = if (is_allowlist) + "INSERT OR IGNORE INTO allowlist_domains (domain, source_id) VALUES (?, ?)" + else + "INSERT OR IGNORE INTO denylist_domains (domain, source_id) VALUES (?, ?)"; + + var insert_stmt = try self.db.prepare(insert_sql); + defer insert_stmt.finalize(); + + for (domains) |domain| { + insert_stmt.reset(); + try insert_stmt.bindText(1, domain); + try insert_stmt.bindInt(2, source_id); + _ = try insert_stmt.step(); + } + + // Update source metadata with status = updated, including invalid_domains count (Gap 6) + var update_stmt = try self.db.prepare( + \\UPDATE denylist_sources SET + \\ domain_count = ?, + \\ invalid_domains = ?, + \\ status = ?, + \\ content_hash = ?, + \\ last_updated = strftime('%s', 'now'), + \\ date_modified = strftime('%s', 'now') + \\WHERE id = ? + ); + defer update_stmt.finalize(); + try update_stmt.bindInt(1, @intCast(domains.len)); + try update_stmt.bindInt(2, @intCast(invalid_count)); + try update_stmt.bindInt(3, @intFromEnum(DenylistStatus.updated)); + try update_stmt.bindText(4, hash); + try update_stmt.bindInt(5, source_id); + _ = try update_stmt.step(); + + try tx.commit(); + + return domains.len; + } + + fn getSourceType(self: *DenylistFetcher, source_id: i64) schema.SourceType { + var stmt = self.db.prepare("SELECT type FROM denylist_sources WHERE id = ?") catch return .denylist; + defer stmt.finalize(); + stmt.bindInt(1, source_id) catch return .denylist; + if (stmt.step() catch false) { + return schema.SourceType.fromInt(stmt.getInt(0)); + } + return .denylist; + } + + fn httpGet(self: *DenylistFetcher, url: []const u8) ![]const u8 { + var client = http.Client{ .allocator = self.allocator }; + defer client.deinit(); + + var response_writer = Writer.Allocating.init(self.allocator); + errdefer response_writer.deinit(); + + const result = client.fetch(.{ + .location = .{ .url = url }, + .response_writer = &response_writer.writer, + }) catch |err| { + std.log.warn("HTTP fetch failed for {s}: {}", .{ url, err }); + return error.FetchFailed; + }; + + if (result.status != .ok) { + std.log.warn("HTTP {d} response for {s}", .{ @intFromEnum(result.status), url }); + return error.FetchFailed; // errdefer handles cleanup + } + + return response_writer.toOwnedSlice() catch error.FetchFailed; + } + + /// Update all enabled denylists + pub fn updateAll(self: *DenylistFetcher) !void { + var stmt = try self.db.prepare("SELECT id, url, comment FROM denylist_sources WHERE enabled = 1"); + defer stmt.finalize(); + + while (try stmt.step()) { + const id = stmt.getInt(0); + const url = stmt.getText(1) orelse continue; + const label = stmt.getText(2) orelse url; + + std.log.info("Updating denylist: {s}", .{label}); + + const url_copy = try self.allocator.dupe(u8, url); + defer self.allocator.free(url_copy); + + const count = self.fetchAndUpdate(id, url_copy) catch |err| { + std.log.warn("Failed to update denylist {s}: {}", .{ label, err }); + continue; + }; + + std.log.info("Denylist {s}: {d} domains", .{ label, count }); + } + + // Signal that denylist needs to be reloaded into memory + events.signalDenylistReload(); + } + + /// Fetch denylists that have 0 domains (incomplete/interrupted fetches) + /// Called at startup to resume any denylists that were mid-fetch when server stopped + pub fn fetchIncomplete(self: *DenylistFetcher) !usize { + var stmt = try self.db.prepare( + "SELECT id, url, comment FROM denylist_sources WHERE enabled = 1 AND domain_count = 0", + ); + defer stmt.finalize(); + + var fetched: usize = 0; + + while (try stmt.step()) { + const id = stmt.getInt(0); + const url = stmt.getText(1) orelse continue; + const label = stmt.getText(2) orelse url; + + std.log.info("Resuming incomplete denylist: {s}", .{label}); + + const url_copy = try self.allocator.dupe(u8, url); + defer self.allocator.free(url_copy); + + const count = self.fetchAndUpdate(id, url_copy) catch |err| { + std.log.warn("Failed to fetch denylist {s}: {}", .{ label, err }); + continue; + }; + + std.log.info("Denylist {s}: {d} domains", .{ label, count }); + fetched += 1; + } + + if (fetched > 0) { + events.signalDenylistReload(); + } + + return fetched; + } +}; + +test "detect hosts format" { + const testing = std.testing; + + const hosts_content = + \\# Comment + \\0.0.0.0 ads.example.com + \\127.0.0.1 tracker.example.com + \\0.0.0.0 malware.example.com + ; + + const format = DenylistFormat.detect(hosts_content); + try testing.expectEqual(DenylistFormat.hosts, format); +} + +test "detect adblock format" { + const testing = std.testing; + + const adblock_content = + \\! AdBlock list + \\||ads.example.com^ + \\||tracker.example.com^$third-party + \\||malware.example.com^ + ; + + const format = DenylistFormat.detect(adblock_content); + try testing.expectEqual(DenylistFormat.adblock, format); +} + +test "parse hosts format" { + const testing = std.testing; + const allocator = testing.allocator; + + const content = + \\# Comment + \\0.0.0.0 ads.example.com + \\127.0.0.1 tracker.example.com + \\0.0.0.0 localhost + ; + + var result = try parseDenylist(content, .hosts, allocator); + defer { + for (result.domains.items) |d| allocator.free(d); + result.domains.deinit(allocator); + } + + try testing.expectEqual(@as(usize, 2), result.domains.items.len); + try testing.expectEqualStrings("ads.example.com", result.domains.items[0]); + try testing.expectEqualStrings("tracker.example.com", result.domains.items[1]); +} + +test "parse adblock format" { + const testing = std.testing; + const allocator = testing.allocator; + + const content = + \\! Comment + \\||ads.example.com^ + \\||tracker.example.com^$third-party + ; + + var result = try parseDenylist(content, .adblock, allocator); + defer { + for (result.domains.items) |d| allocator.free(d); + result.domains.deinit(allocator); + } + + try testing.expectEqual(@as(usize, 2), result.domains.items.len); + try testing.expectEqualStrings("ads.example.com", result.domains.items[0]); + try testing.expectEqualStrings("tracker.example.com", result.domains.items[1]); +} + +test "parse domain list format" { + const testing = std.testing; + const allocator = testing.allocator; + + const content = + \\# Comment + \\ads.example.com + \\tracker.example.com # inline comment + \\malware.example.com + ; + + var result = try parseDenylist(content, .domains, allocator); + defer { + for (result.domains.items) |d| allocator.free(d); + result.domains.deinit(allocator); + } + + try testing.expectEqual(@as(usize, 3), result.domains.items.len); +} + +test "parse tracks invalid domains" { + const testing = std.testing; + const allocator = testing.allocator; + + const content = + \\# Comment + \\valid.example.com + \\-invalid-start.com + \\also.valid.org + \\toolong12345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345678901234567890123456789012345.com + ; + + var result = try parseDenylist(content, .domains, allocator); + defer { + for (result.domains.items) |d| allocator.free(d); + result.domains.deinit(allocator); + } + + try testing.expectEqual(@as(usize, 2), result.domains.items.len); + try testing.expectEqual(@as(usize, 2), result.invalid_count); +} diff --git a/src/filter/safe_search.zig b/src/filter/safe_search.zig new file mode 100644 index 0000000..ea41475 --- /dev/null +++ b/src/filter/safe_search.zig @@ -0,0 +1,142 @@ +const std = @import("std"); +const toLower = @import("../util.zig").toLower; + +/// Safe search enforcement - rewrites queries to force safe search +/// Based on PLAN.md specifications + +/// Safe search rewrite rule +const SafeSearchRewrite = struct { + from: []const u8, + to: []const u8, +}; + +/// Google safe search rewrites +const GOOGLE_REWRITES = [_]SafeSearchRewrite{ + .{ .from = "www.google.com", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.uk", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.ca", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.com.au", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.de", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.fr", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.es", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.it", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.nl", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.be", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.ch", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.at", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.pl", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.se", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.no", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.dk", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.fi", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.pt", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.ru", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.jp", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.kr", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.com.br", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.com.mx", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.in", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.nz", .to = "forcesafesearch.google.com" }, + .{ .from = "www.google.co.za", .to = "forcesafesearch.google.com" }, +}; + +/// Bing safe search rewrites +const BING_REWRITES = [_]SafeSearchRewrite{ + .{ .from = "www.bing.com", .to = "strict.bing.com" }, + .{ .from = "bing.com", .to = "strict.bing.com" }, +}; + +/// YouTube safe search rewrites (restrictmoderate for moderate, restrict for strict) +const YOUTUBE_REWRITES = [_]SafeSearchRewrite{ + .{ .from = "www.youtube.com", .to = "restrictmoderate.youtube.com" }, + .{ .from = "youtube.com", .to = "restrictmoderate.youtube.com" }, + .{ .from = "m.youtube.com", .to = "restrictmoderate.youtube.com" }, + .{ .from = "youtubei.googleapis.com", .to = "restrictmoderate.youtube.com" }, + .{ .from = "youtube.googleapis.com", .to = "restrictmoderate.youtube.com" }, + .{ .from = "www.youtube-nocookie.com", .to = "restrictmoderate.youtube.com" }, +}; + +/// DuckDuckGo safe search rewrites +const DUCKDUCKGO_REWRITES = [_]SafeSearchRewrite{ + .{ .from = "duckduckgo.com", .to = "safe.duckduckgo.com" }, + .{ .from = "www.duckduckgo.com", .to = "safe.duckduckgo.com" }, +}; + +/// Pixabay safe search rewrites +const PIXABAY_REWRITES = [_]SafeSearchRewrite{ + .{ .from = "pixabay.com", .to = "safesearch.pixabay.com" }, +}; + +/// Check if a domain should be rewritten for safe search +/// Returns the safe search domain if a rewrite is needed, null otherwise +pub fn applySafeSearch(domain: []const u8) ?[]const u8 { + // Lowercase the domain for case-insensitive matching + var lower_buf: [256]u8 = undefined; + const lower = toLower(domain, &lower_buf) orelse return null; + + // Check all rewrite rules + inline for (GOOGLE_REWRITES) |rule| { + if (std.mem.eql(u8, lower, rule.from)) return rule.to; + } + inline for (BING_REWRITES) |rule| { + if (std.mem.eql(u8, lower, rule.from)) return rule.to; + } + inline for (YOUTUBE_REWRITES) |rule| { + if (std.mem.eql(u8, lower, rule.from)) return rule.to; + } + inline for (DUCKDUCKGO_REWRITES) |rule| { + if (std.mem.eql(u8, lower, rule.from)) return rule.to; + } + inline for (PIXABAY_REWRITES) |rule| { + if (std.mem.eql(u8, lower, rule.from)) return rule.to; + } + + return null; +} + +/// Get safe search CNAME target for a domain +/// Used to create a CNAME response that redirects to safe search +pub fn getSafeSearchCname(domain: []const u8) ?[]const u8 { + return applySafeSearch(domain); +} + +/// Check if safe search is applicable to this domain (for logging) +pub fn isSafeSearchDomain(domain: []const u8) bool { + return applySafeSearch(domain) != null; +} + +test "safe search Google" { + const testing = std.testing; + + try testing.expectEqualStrings("forcesafesearch.google.com", applySafeSearch("www.google.com").?); + try testing.expectEqualStrings("forcesafesearch.google.com", applySafeSearch("WWW.GOOGLE.COM").?); + try testing.expectEqualStrings("forcesafesearch.google.com", applySafeSearch("www.google.co.uk").?); +} + +test "safe search Bing" { + const testing = std.testing; + + try testing.expectEqualStrings("strict.bing.com", applySafeSearch("www.bing.com").?); + try testing.expectEqualStrings("strict.bing.com", applySafeSearch("bing.com").?); +} + +test "safe search YouTube" { + const testing = std.testing; + + try testing.expectEqualStrings("restrictmoderate.youtube.com", applySafeSearch("www.youtube.com").?); + try testing.expectEqualStrings("restrictmoderate.youtube.com", applySafeSearch("m.youtube.com").?); +} + +test "safe search DuckDuckGo" { + const testing = std.testing; + + try testing.expectEqualStrings("safe.duckduckgo.com", applySafeSearch("duckduckgo.com").?); +} + +test "safe search non-matching" { + const testing = std.testing; + + try testing.expect(applySafeSearch("example.com") == null); + try testing.expect(applySafeSearch("google.com") == null); // Note: no www + try testing.expect(applySafeSearch("mail.google.com") == null); +} diff --git a/src/logging/logger.zig b/src/logging/logger.zig new file mode 100644 index 0000000..ff3c644 --- /dev/null +++ b/src/logging/logger.zig @@ -0,0 +1,197 @@ +const std = @import("std"); +const posix = std.posix; +const fs = std.fs; + +/// Log output destination +pub const LogOutput = union(enum) { + stderr: void, + syslog: void, + file: fs.File, +}; + +/// Global logger state +pub const Logger = struct { + output: LogOutput, + level: std.log.Level, + mutex: std.Thread.Mutex, + + /// Global logger instance + var global: ?*Logger = null; + + pub fn init(output_str: []const u8, level: std.log.Level) !Logger { + const output: LogOutput = if (std.mem.eql(u8, output_str, "stderr")) + .stderr + else if (std.mem.eql(u8, output_str, "syslog")) + .syslog + else blk: { + // Treat as file path + const file = try fs.cwd().createFile(output_str, .{ + .truncate = false, + }); + // Seek to end for append - if this fails, we'll overwrite from start (not ideal but not fatal) + file.seekFromEnd(0) catch {}; + break :blk LogOutput{ .file = file }; + }; + + return Logger{ + .output = output, + .level = level, + .mutex = std.Thread.Mutex{}, + }; + } + + pub fn deinit(self: *Logger) void { + switch (self.output) { + .file => |f| f.close(), + else => {}, + } + } + + /// Set this logger as the global logger + pub fn setGlobal(self: *Logger) void { + global = self; + } + + /// Clear the global logger + pub fn clearGlobal() void { + global = null; + } + + /// Write a log message + pub fn write(self: *Logger, level: std.log.Level, scope: []const u8, message: []const u8) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + switch (self.output) { + .stderr => self.writeStderr(level, scope, message), + .syslog => self.writeSyslog(level, message), + .file => |f| self.writeFile(f, level, scope, message), + } + } + + fn writeStderr(self: *Logger, level: std.log.Level, scope: []const u8, message: []const u8) void { + _ = self; + const level_str = switch (level) { + .err => "error", + .warn => "warning", + .info => "info", + .debug => "debug", + }; + + // Get current time + const ts = std.time.timestamp(); + const epoch_secs: std.time.epoch.EpochSeconds = .{ .secs = @intCast(ts) }; + const day_secs = epoch_secs.getDaySeconds(); + const hours = day_secs.getHoursIntoDay(); + const minutes = day_secs.getMinutesIntoHour(); + const seconds = day_secs.getSecondsIntoMinute(); + + var buf: [4096]u8 = undefined; + const formatted = if (scope.len > 0 and !std.mem.eql(u8, scope, "default")) + std.fmt.bufPrint(&buf, "{d:0>2}:{d:0>2}:{d:0>2} {s}({s}): {s}\n", .{ hours, minutes, seconds, level_str, scope, message }) catch return + else + std.fmt.bufPrint(&buf, "{d:0>2}:{d:0>2}:{d:0>2} {s}: {s}\n", .{ hours, minutes, seconds, level_str, message }) catch return; + + // Intentionally silent - if stderr write fails, there's nowhere else to report it + _ = posix.write(posix.STDERR_FILENO, formatted) catch {}; + } + + fn writeFile(self: *Logger, file: fs.File, level: std.log.Level, scope: []const u8, message: []const u8) void { + _ = self; + const timestamp = std.time.timestamp(); + const level_str = switch (level) { + .err => "ERROR", + .warn => "WARN", + .info => "INFO", + .debug => "DEBUG", + }; + + var buf: [4096]u8 = undefined; + const formatted = if (scope.len > 0 and !std.mem.eql(u8, scope, "default")) + std.fmt.bufPrint(&buf, "{d} [{s}] ({s}): {s}\n", .{ timestamp, level_str, scope, message }) catch return + else + std.fmt.bufPrint(&buf, "{d} [{s}]: {s}\n", .{ timestamp, level_str, message }) catch return; + + // Intentionally silent - logging failures shouldn't crash the application + // If file write fails repeatedly, the application will continue but logs will be lost + _ = posix.write(file.handle, formatted) catch {}; + } + + fn writeSyslog(self: *Logger, level: std.log.Level, message: []const u8) void { + _ = self; + // Syslog priority: facility (user = 1) * 8 + severity + const severity: u8 = switch (level) { + .err => 3, // error + .warn => 4, // warning + .info => 6, // informational + .debug => 7, // debug + }; + const priority = (1 << 3) | severity; // USER facility + + // Create socket and send to /dev/log + const sock = posix.socket(posix.AF.UNIX, posix.SOCK.DGRAM, 0) catch return; + defer posix.close(sock); + + var addr: posix.sockaddr.un = .{ + .family = posix.AF.UNIX, + .path = undefined, + }; + const path = "/dev/log"; + @memcpy(addr.path[0..path.len], path); + addr.path[path.len] = 0; + + // Format: message + var buf: [1024]u8 = undefined; + const formatted = std.fmt.bufPrint(&buf, "<{d}>nxdns: {s}", .{ priority, message }) catch return; + + // Intentionally silent - syslog failures shouldn't crash the application + _ = posix.sendto( + sock, + formatted, + 0, + @ptrCast(&addr), + @sizeOf(posix.sockaddr.un), + ) catch {}; + } + + /// Custom log function compatible with std.options.logFn + pub fn log( + comptime level: std.log.Level, + comptime scope: @TypeOf(.enum_literal), + comptime format: []const u8, + args: anytype, + ) void { + const logger = global orelse { + // Fall back to default logging + std.log.defaultLog(level, scope, format, args); + return; + }; + + // Check log level + if (@intFromEnum(level) > @intFromEnum(logger.level)) return; + + // Format the message + var buf: [4096]u8 = undefined; + const message = std.fmt.bufPrint(&buf, format, args) catch return; + + const scope_name = @tagName(scope); + logger.write(level, scope_name, message); + } +}; + +/// Convert config log level to std.log.Level +pub fn configLevelToStd(level: anytype) std.log.Level { + return switch (level) { + .debug => .debug, + .info => .info, + .warn => .warn, + .err => .err, + }; +} + +test "Logger stderr" { + var logger = try Logger.init("stderr", .info); + defer logger.deinit(); + + logger.write(.info, "test", "Hello from test"); +} diff --git a/src/main.zig b/src/main.zig new file mode 100644 index 0000000..f3868bc --- /dev/null +++ b/src/main.zig @@ -0,0 +1,533 @@ +const std = @import("std"); +const net = std.net; +const fs = std.fs; + +// Logging module +const app_logger = @import("logging/logger.zig"); + +// Override std.log to use our custom logger +pub const std_options: std.Options = .{ + .logFn = app_logger.Logger.log, +}; + +// DNS modules +pub const dns = struct { + pub const types = @import("dns/types.zig"); + pub const header = @import("dns/header.zig"); + pub const name = @import("dns/name.zig"); + pub const question = @import("dns/question.zig"); + pub const record = @import("dns/record.zig"); + pub const edns = @import("dns/edns.zig"); + pub const packet = @import("dns/packet.zig"); +}; + +// Server modules +const udp_server = @import("server/udp.zig"); +const tcp_server = @import("server/tcp.zig"); +const handler_mod = @import("server/handler.zig"); +const shutdown = @import("server/shutdown.zig"); +const rate_limiter_mod = @import("server/rate_limiter.zig"); + +// Upstream modules +const pool = @import("upstream/pool.zig"); + +// Storage modules +const db_mod = @import("storage/db.zig"); +const schema = @import("storage/schema.zig"); +const logger_mod = @import("storage/logger.zig"); + +// Filter modules +const denylist_mod = @import("filter/denylist.zig"); +const fetcher_mod = @import("filter/fetcher.zig"); + +// Cache module +const cache_mod = @import("cache/dns_cache.zig"); + +// Config module +const config_mod = @import("config/config.zig"); +const watcher_mod = @import("config/watcher.zig"); + +// Events module +const events = @import("events.zig"); + +// Web module +const web_server = @import("web/server.zig"); + +const VERSION = "0.1.0"; + +/// Context for config reload callback (needed because callback is a simple fn pointer) +const ReloadContext = struct { + db: *db_mod.Database, + denylist: *denylist_mod.ThreadSafeDenylist, + config_path: []const u8, + allocator: std.mem.Allocator, +}; + +/// Context for log cleanup thread +const CleanupContext = struct { + db: *db_mod.Database, + retention_seconds: i64, + coordinator: *shutdown.ShutdownCoordinator, +}; + +/// Daily log cleanup thread - runs cleanup every 24 hours +fn logCleanupThread(ctx: *CleanupContext) void { + const day_ms: i32 = 24 * 60 * 60 * 1000; + + while (!ctx.coordinator.isShutdownRequested()) { + // Block until shutdown or 24 hours pass + if (ctx.coordinator.wait(day_ms)) break; + + std.log.info("Running daily log cleanup...", .{}); + logger_mod.cleanupOldLogs(ctx.db, ctx.retention_seconds) catch |err| { + std.log.warn("Daily log cleanup failed: {}", .{err}); + }; + } +} + +/// Context for denylist auto-update thread +const AutoUpdateContext = struct { + db: *db_mod.Database, + coordinator: *shutdown.ShutdownCoordinator, + allocator: std.mem.Allocator, +}; + +/// Denylist auto-update thread - fetches updates daily +fn autoUpdateThread(ctx: *AutoUpdateContext) void { + const update_interval_ms: i32 = 24 * 60 * 60 * 1000; // 24 hours + + while (!ctx.coordinator.isShutdownRequested()) { + // Wait until shutdown or 24 hours pass + if (ctx.coordinator.wait(update_interval_ms)) break; + + std.log.info("Running scheduled denylist update...", .{}); + var fetcher = fetcher_mod.DenylistFetcher.init(ctx.db, ctx.allocator); + fetcher.updateAll() catch |err| { + std.log.warn("Scheduled denylist update failed: {}", .{err}); + }; + } +} + +/// Atomic storage for reload context (accessed from multiple threads) +var config_reload_ctx: std.atomic.Value(?*ReloadContext) = .{ .raw = null }; + +/// Public function to reload denylist from database +/// Called after denylist updates via web UI +pub fn reloadDenylist() void { + const ctx = config_reload_ctx.load(.acquire) orelse { + std.log.warn("Denylist reload called but context is null", .{}); + return; + }; + + std.log.info("Reloading denylist from database...", .{}); + + const new_denylist = denylist_mod.loadFromDatabase(ctx.db, ctx.allocator) catch |err| { + std.log.err("Failed to reload denylist: {}", .{err}); + return; + }; + + var old_denylist = ctx.denylist.swap(new_denylist); + old_denylist.deinit(); + + std.log.info("Denylist reloaded successfully", .{}); +} + +/// Callback triggered when config file changes +fn configReloadCallback() void { + const ctx = config_reload_ctx.load(.acquire) orelse { + std.log.warn("Config reload callback called but context is null", .{}); + return; + }; + + std.log.info("Reloading configuration...", .{}); + + // Reload config + var new_config = config_mod.Config.load(ctx.config_path, ctx.allocator) catch |err| { + std.log.err("Failed to reload config: {}", .{err}); + return; + }; + defer new_config.deinit(); + + new_config.validate() catch |err| { + std.log.err("New config is invalid: {}", .{err}); + return; + }; + + // Reload denylist from database + const new_denylist = denylist_mod.loadFromDatabase(ctx.db, ctx.allocator) catch |err| { + std.log.err("Failed to reload denylist: {}", .{err}); + return; + }; + + // Atomically swap denylists using RwLock (thread-safe) + // The swap acquires a write lock, ensuring no readers are active + var old_denylist = ctx.denylist.swap(new_denylist); + + // Clean up old denylist after swap is complete + // At this point, no readers can access the old denylist + old_denylist.deinit(); + + std.log.info("Configuration reloaded successfully", .{}); +} + +pub fn main() !void { + var gpa = std.heap.GeneralPurposeAllocator(.{}){}; + defer _ = gpa.deinit(); + const allocator = gpa.allocator(); + + var args = std.process.args(); + _ = args.next(); // Skip program name + + const command = args.next() orelse { + printUsage(); + return; + }; + + if (std.mem.eql(u8, command, "run")) { + try runServer(allocator); + } else if (std.mem.eql(u8, command, "check")) { + try checkConfig(allocator); + } else if (std.mem.eql(u8, command, "version")) { + printVersion(); + } else if (std.mem.eql(u8, command, "help") or std.mem.eql(u8, command, "-h") or std.mem.eql(u8, command, "--help")) { + printUsage(); + } else { + std.log.err("Unknown command: {s}", .{command}); + printUsage(); + } +} + +fn printUsage() void { + std.debug.print( + \\nxdns - DNS sinkhole + \\ + \\Usage: nxdns + \\ + \\Commands: + \\ run Start the DNS server + \\ check Validate configuration + \\ version Show version information + \\ help Show this help message + \\ + \\Environment Variables: + \\ NXDNS_CONFIG Path to configuration file + \\ + , .{}); +} + +fn printVersion() void { + std.debug.print("nxdns version {s}\n", .{VERSION}); +} + +fn getConfigPath() []const u8 { + return std.posix.getenv("NXDNS_CONFIG") orelse "/etc/nxdns/config.toml"; +} + +fn runServer(allocator: std.mem.Allocator) !void { + std.debug.print("nxdns v{s} starting...\n", .{VERSION}); + + // Load configuration + const config_path = getConfigPath(); + var config = try config_mod.Config.load(config_path, allocator); + defer config.deinit(); + + try config.validate(); + std.debug.print("Configuration loaded from {s}\n", .{config_path}); + + // Initialize logging + const log_level = app_logger.configLevelToStd(config.logging.level); + var logger = app_logger.Logger.init(config.logging.output, log_level) catch |err| blk: { + std.debug.print("Warning: Failed to initialize logger: {}, using stderr\n", .{err}); + break :blk app_logger.Logger.init("stderr", log_level) catch |e| { + std.debug.print("Fatal: Cannot initialize logging to stderr: {}\n", .{e}); + return e; + }; + }; + defer logger.deinit(); + logger.setGlobal(); + defer app_logger.Logger.clearGlobal(); + + // Open database + var db = try db_mod.Database.open(config.database_path, allocator); + defer db.close(); + + // Run migrations + try schema.migrate(&db); + std.debug.print("Database initialized at {s}\n", .{config.database_path}); + + // Cleanup old logs based on retention policy + const retention_seconds = config.logging.retentionSeconds(); + logger_mod.cleanupOldLogs(&db, retention_seconds) catch |err| { + std.log.warn("Failed to cleanup old logs: {}", .{err}); + }; + + // Initialize denylist with thread-safe wrapper + var denylist = try denylist_mod.loadFromDatabase(&db, allocator); + defer denylist.deinit(); + var safe_denylist = denylist_mod.ThreadSafeDenylist.init(&denylist, allocator); + std.debug.print("Denylist loaded ({} domains)\n", .{denylist.denied_domains.count()}); + + // Resume any incomplete denylist fetches from previous run + { + var fetcher = fetcher_mod.DenylistFetcher.init(&db, allocator); + const resumed = fetcher.fetchIncomplete() catch |err| blk: { + std.log.warn("Failed to resume incomplete denylists: {}", .{err}); + break :blk 0; + }; + if (resumed > 0) { + // Reload denylist with newly fetched domains + var new_denylist = denylist_mod.loadFromDatabase(&db, allocator) catch |err| { + std.log.err("Failed to reload denylist after resume: {}", .{err}); + return err; + }; + var old = safe_denylist.swap(new_denylist); + old.deinit(); + std.debug.print("Denylist updated after resume ({} domains)\n", .{new_denylist.denied_domains.count()}); + } + } + + // Initialize cache + var cache = cache_mod.DnsCache.initWithConfig( + allocator, + config.dns.cache_size, + cache_mod.DnsCache.DEFAULT_MIN_TTL, + cache_mod.DnsCache.DEFAULT_MAX_TTL, + ); + defer cache.deinit(); + std.debug.print("DNS cache initialized (max {} entries)\n", .{config.dns.cache_size}); + + // Initialize query logger + var query_logger = logger_mod.QueryLogger.init(&db, allocator); + defer query_logger.deinit(); + + // Initialize upstream pool + const upstream_configs = try allocator.alloc(pool.UpstreamConfig, config.upstream.servers.len); + defer allocator.free(upstream_configs); + + for (config.upstream.servers, 0..) |server, i| { + upstream_configs[i] = pool.UpstreamConfig{ + .address = server, + }; + } + + var upstream = try pool.UpstreamPool.init(upstream_configs, allocator); + defer upstream.deinit(); + std.debug.print("Upstream DNS configured ({} servers)\n", .{config.upstream.servers.len}); + + // Initialize rate limiter + var limiter = rate_limiter_mod.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 20, + .enabled = true, + }); + defer limiter.deinit(); + std.debug.print("Rate limiting enabled (20 qps per client)\n", .{}); + + // Initialize handler + var handler = handler_mod.Handler.init(allocator); + handler.setDenylist(safe_denylist.toHandlerDenylist()); + handler.setCache(cache.toHandlerCache()); + handler.setUpstream(upstream.toHandlerUpstream()); + handler.setLogger(query_logger.toHandlerLogger()); + handler.setRateLimiter(&limiter); + handler.config.blocking_response = switch (config.blocking.response) { + .zero => .zero, + .nxdomain => .nxdomain, + }; + + // Start DNS servers + const dns_addr = try net.Address.parseIp(config.dns.bind, config.dns.port); + + var udp_handler = handler.toUdpHandler(); + var udp = try udp_server.UdpServer.initWithConfig(dns_addr, &udp_handler, allocator, .{ + .num_workers = config.dns.workers, + }); + defer udp.deinit(); + + var tcp_handler = handler.toTcpHandler(); + var tcp = try tcp_server.TcpServer.initWithConfig(dns_addr, &tcp_handler, allocator, .{ + .num_workers = config.dns.workers, + }); + defer tcp.deinit(); + + std.debug.print("DNS server listening on {}:{}\n", .{ dns_addr.in.sa.addr, config.dns.port }); + + // Start web server + const web_addr = try net.Address.parseIp(config.web.bind, config.web.port); + var web = try web_server.WebServer.init(web_addr, &db, allocator); + defer web.deinit(); + std.debug.print("Web UI available at http://{}:{}\n", .{ web_addr.in.sa.addr, config.web.port }); + + // Wire up SSE subscriber for real-time query log streaming + query_logger.setSubscriber(web.toLoggerSubscriber()); + + // Set up shutdown coordinator (owns signalfd for SIGINT/SIGTERM) + var coordinator = shutdown.ShutdownCoordinator.init(); + defer coordinator.deinit(); + + // Initialize event signaling for denylist reloads + try events.initDenylistEvent(); + defer events.deinitDenylistEvent(); + + // Start config file watcher for hot reload + var config_watcher: ?watcher_mod.ConfigWatcher = null; + + // Store pointers for the reload callback + var reload_context = ReloadContext{ + .db = &db, + .denylist = &safe_denylist, + .config_path = config_path, + .allocator = allocator, + }; + config_reload_ctx.store(&reload_context, .release); + + // Context for daily log cleanup + var cleanup_context = CleanupContext{ + .db = &db, + .retention_seconds = retention_seconds, + .coordinator = &coordinator, + }; + + // Context for denylist auto-updates (if enabled) + var auto_update_context = AutoUpdateContext{ + .db = &db, + .coordinator = &coordinator, + .allocator = allocator, + }; + + if (watcher_mod.ConfigWatcher.init(config_path, allocator, &coordinator)) |watcher| { + config_watcher = watcher; + config_watcher.?.setConfigReloadCallback(configReloadCallback); + config_watcher.?.setDenylistReloadCallback(reloadDenylist); + std.debug.print("Config watcher started (inotify)\n", .{}); + } else |err| { + std.log.warn("Failed to start config watcher: {} - config hot reload disabled", .{err}); + } + defer if (config_watcher) |*w| w.deinit(); + + std.debug.print("\nnxdns is running. Press Ctrl+C to stop.\n\n", .{}); + + // Start server threads + const udp_thread = try std.Thread.spawn(.{}, udp_server.UdpServer.run, .{&udp}); + const tcp_thread = try std.Thread.spawn(.{}, tcp_server.TcpServer.run, .{&tcp}); + const web_thread = try std.Thread.spawn(.{}, web_server.WebServer.run, .{&web}); + const cleanup_thread = try std.Thread.spawn(.{}, logCleanupThread, .{&cleanup_context}); + + // Start auto-update thread if enabled + var auto_update_thread: ?std.Thread = null; + if (config.blocking.auto_update) { + auto_update_thread = try std.Thread.spawn(.{}, autoUpdateThread, .{&auto_update_context}); + std.debug.print("Denylist auto-update enabled (daily)\n", .{}); + } + + // Start config watcher in background thread + var watcher_thread: ?std.Thread = null; + if (config_watcher) |*w| { + watcher_thread = try std.Thread.spawn(.{}, watcher_mod.ConfigWatcher.watch, .{w}); + } + + // Main thread waits for shutdown signal (SIGINT/SIGTERM via signalfd) + _ = coordinator.waitForSignal(); + + std.debug.print("\nShutting down...\n", .{}); + + // Stop all servers + std.debug.print(" Stopping servers...\n", .{}); + udp.stop(); + tcp.stop(); + web.stop(); + + // Wait for connections to drain + std.debug.print(" Waiting for connections...\n", .{}); + tcp.waitForConnections(2000); + web.waitForConnections(2000); + + // Flush query logs + std.debug.print(" Flushing logs...\n", .{}); + query_logger.flush() catch |err| { + std.log.warn("Failed to flush query logs during shutdown: {}", .{err}); + }; + + // Wait for threads to finish + std.debug.print(" Joining UDP...\n", .{}); + udp_thread.join(); + std.debug.print(" Joining TCP...\n", .{}); + tcp_thread.join(); + std.debug.print(" Joining web...\n", .{}); + web_thread.join(); + std.debug.print(" Joining cleanup...\n", .{}); + cleanup_thread.join(); + if (auto_update_thread) |t| { + std.debug.print(" Joining auto-update...\n", .{}); + t.join(); + } + if (watcher_thread) |t| { + std.debug.print(" Joining watcher...\n", .{}); + t.join(); + } + + std.debug.print("Goodbye!\n", .{}); +} + +fn checkConfig(allocator: std.mem.Allocator) !void { + const config_path = getConfigPath(); + + std.debug.print("Checking configuration at {s}...\n", .{config_path}); + + var config = config_mod.Config.load(config_path, allocator) catch |err| { + std.debug.print("ERROR: Failed to load config: {}\n", .{err}); + return; + }; + defer config.deinit(); + + config.validate() catch |err| { + std.debug.print("ERROR: Invalid configuration: {}\n", .{err}); + return; + }; + + std.debug.print("Configuration is valid.\n", .{}); + std.debug.print("\nSettings:\n", .{}); + std.debug.print(" DNS port: {}\n", .{config.dns.port}); + std.debug.print(" DNS bind: {s}\n", .{config.dns.bind}); + std.debug.print(" Web port: {}\n", .{config.web.port}); + std.debug.print(" Web bind: {s}\n", .{config.web.bind}); + std.debug.print(" Upstream servers: {}\n", .{config.upstream.servers.len}); + std.debug.print(" Cache size: {}\n", .{config.dns.cache_size}); + std.debug.print(" Safe search: {}\n", .{config.safe_search.enabled}); + std.debug.print(" Deny response: {s}\n", .{@tagName(config.blocking.response)}); + std.debug.print(" Denylist auto-update: {}\n", .{config.blocking.auto_update}); + std.debug.print(" Log retention: {s}\n", .{config.logging.retention}); +} + +// Tests +test { + // Run all tests from imported modules + // DNS protocol + _ = dns.types; + _ = dns.header; + _ = dns.name; + _ = dns.question; + _ = dns.record; + _ = dns.edns; + _ = dns.packet; + + // Server components + _ = handler_mod; + _ = rate_limiter_mod; + _ = udp_server; + _ = tcp_server; + + // Upstream + _ = pool; + _ = @import("upstream/doh.zig"); + _ = @import("upstream/dot.zig"); + + // Cache + _ = cache_mod; + + // Filter + _ = denylist_mod; + + // Config + _ = config_mod; + _ = watcher_mod; +} diff --git a/src/memtest.zig b/src/memtest.zig new file mode 100644 index 0000000..f158b73 --- /dev/null +++ b/src/memtest.zig @@ -0,0 +1,86 @@ +const std = @import("std"); +const fs = std.fs; +const denylist_mod = @import("filter/denylist.zig"); + +fn getVmRSS() usize { + const status_file = fs.openFileAbsolute("/proc/self/status", .{}) catch return 0; + defer status_file.close(); + + var status_buf: [4096]u8 = undefined; + const status_len = status_file.readAll(&status_buf) catch return 0; + const status = status_buf[0..status_len]; + + var lines = std.mem.splitScalar(u8, status, '\n'); + while (lines.next()) |l| { + if (std.mem.startsWith(u8, l, "VmRSS:")) { + const rest = std.mem.trim(u8, l["VmRSS:".len..], " \t"); + var parts = std.mem.splitScalar(u8, rest, ' '); + if (parts.next()) |kb_str| { + return std.fmt.parseInt(usize, kb_str, 10) catch 0; + } + } + } + return 0; +} + +pub fn main() !void { + var gpa = std.heap.GeneralPurposeAllocator(.{}){}; + defer _ = gpa.deinit(); + const allocator = gpa.allocator(); + + const mem_before = getVmRSS(); + + var denylist = denylist_mod.Denylist.init(allocator); + defer denylist.deinit(); + + const denylist_dir = ".ignore/denylists"; + var dir = try fs.cwd().openDir(denylist_dir, .{ .iterate = true }); + defer dir.close(); + + var total_domains: usize = 0; + var total_lines: usize = 0; + + const start = std.time.milliTimestamp(); + + var iter = dir.iterate(); + while (try iter.next()) |entry| { + if (entry.kind != .file) continue; + + // Read entire file + const content = dir.readFileAlloc(allocator, entry.name, 100 * 1024 * 1024) catch continue; + defer allocator.free(content); + + var lines = std.mem.splitScalar(u8, content, '\n'); + while (lines.next()) |line| { + total_lines += 1; + + // Parse Adblock format: ||domain^ + if (std.mem.startsWith(u8, line, "||")) { + const rest = line[2..]; + if (std.mem.indexOfScalar(u8, rest, '^')) |end| { + const domain = rest[0..end]; + if (domain.len > 0 and domain.len < 254) { + denylist.addDeniedDomain(domain, 0) catch continue; + total_domains += 1; + } + } + } + } + } + + const load_time = std.time.milliTimestamp() - start; + const mem_after = getVmRSS(); + const mem_delta_kb = if (mem_after > mem_before) mem_after - mem_before else 0; + const mem_delta_bytes = mem_delta_kb * 1024; + const domain_count = denylist.denied_domains.count(); + const bytes_per_domain = if (domain_count > 0) mem_delta_bytes / domain_count else 0; + + std.debug.print("\n[Memory Test - Real Denylists]\n", .{}); + std.debug.print(" Total lines: {d}\n", .{total_lines}); + std.debug.print(" Unique domains: {d}\n", .{domain_count}); + std.debug.print(" Load time: {d}ms\n", .{load_time}); + std.debug.print(" Memory before: {d} KB\n", .{mem_before}); + std.debug.print(" Memory after: {d} KB\n", .{mem_after}); + std.debug.print(" Memory delta: {d} KB ({d} MB)\n", .{ mem_delta_kb, mem_delta_kb / 1024 }); + std.debug.print(" Bytes per domain: {d}\n", .{bytes_per_domain}); +} diff --git a/src/server/handler.zig b/src/server/handler.zig new file mode 100644 index 0000000..ed8c94d --- /dev/null +++ b/src/server/handler.zig @@ -0,0 +1,655 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const packet = @import("../dns/packet.zig"); +const types = @import("../dns/types.zig"); +const safesearch = @import("../filter/safe_search.zig"); +const schema = @import("../storage/schema.zig"); +const udp = @import("udp.zig"); +const tcp = @import("tcp.zig"); +const rate_limiter = @import("rate_limiter.zig"); + +/// Query log entry with full status tracking +pub const QueryLogEntry = struct { + timestamp: i64, + domain: []const u8, + client_ip: []const u8, + qtype: types.QType, + // Status tracking (Gap 1) + status: schema.QueryStatus, + // Attribution tracking (Gap 2) + list_id: ?i64, + rule_id: ?i64, + // Reply type tracking (Gap 3) + reply_type: schema.ReplyType, + // Protocol tracking (Gap 8) + protocol: schema.ClientProtocol, + // Performance + response_time_us: u64, + upstream: ?[]const u8, + reason: ?[]const u8, + /// DNSSEC validation status from upstream (AD bit) + dnssec_validated: bool, + + // Legacy compatibility - computed from status + pub fn isDenied(self: QueryLogEntry) bool { + return self.status.isDenied(); + } +}; + +/// Handler configuration +pub const HandlerConfig = struct { + blocking_response: BlockingResponse = .zero, + safe_search_enabled: bool = true, + rate_limit_enabled: bool = true, + rate_limit_qps: u32 = 20, + + pub const BlockingResponse = enum { + zero, // Return 0.0.0.0 + nxdomain, // Return NXDOMAIN + }; +}; + +// Type-erased interfaces are used here to avoid circular imports: +// handler.zig defines interfaces, concrete types implement them. +// This allows handler to remain decoupled from specific implementations. + +/// Result of denylist check with attribution +pub const DenyResult = struct { + denied: bool, + list_id: ?i64, // Which denylist source caused the block + rule_id: ?i64, // Which rule caused the block/allow + is_rule_allow: bool, // True if an allow rule overrode a denylist match +}; + +/// Interface for denylist checking +/// Gap 4: Uses u64 groups bitmask to support many-to-many client-group relationships +pub const Denylist = struct { + context: *anyopaque, + checkFn: *const fn (*anyopaque, []const u8, u64) DenyResult, + getGroupsFn: *const fn (*anyopaque, []const u8) u64, + + /// Check if domain is denied for groups (bitmask), with full attribution + pub fn checkWithMask(self: Denylist, domain: []const u8, groups_mask: u64) DenyResult { + return self.checkFn(self.context, domain, groups_mask); + } + + /// Legacy: Check if domain is denied for a single group + pub fn check(self: Denylist, domain: []const u8, group_id: u32) DenyResult { + const mask: u64 = if (group_id >= 64) 0 else (@as(u64, 1) << @intCast(group_id)); + return self.checkWithMask(domain, mask); + } + + /// Legacy compatibility - just returns boolean + pub fn isDenied(self: Denylist, domain: []const u8, group_id: u32) bool { + return self.check(domain, group_id).denied; + } + + /// Get all groups for a client as a bitmask (Gap 4: many-to-many) + pub fn getGroupsForClient(self: Denylist, client_ip: []const u8) u64 { + return self.getGroupsFn(self.context, client_ip); + } + + /// Legacy: Get primary group ID for a client + pub fn getGroupForClient(self: Denylist, client_ip: []const u8) u32 { + const mask = self.getGroupsForClient(client_ip); + if (mask == 1) return 0; // Only default group + var bit: u6 = 0; + while (bit < 64) : (bit += 1) { + if ((mask & (@as(u64, 1) << bit)) != 0) return bit; + } + return 0; + } +}; + +/// Interface for upstream DNS +pub const Upstream = struct { + context: *anyopaque, + queryFn: *const fn (*anyopaque, []const u8, Allocator) ?[]const u8, + + pub fn query(self: Upstream, dns_packet: []const u8, allocator: Allocator) ?[]const u8 { + return self.queryFn(self.context, dns_packet, allocator); + } +}; + +/// Interface for DNS cache +/// Note: getCopy returns a mutable slice that the caller owns and must free +pub const Cache = struct { + context: *anyopaque, + getFn: *const fn (*anyopaque, []const u8, types.QType) ?[]u8, + putFn: *const fn (*anyopaque, []const u8, types.QType, []const u8, u32) void, + + /// Get a copy of a cached response. Caller owns the returned memory. + pub fn getCopy(self: Cache, domain: []const u8, qtype: types.QType) ?[]u8 { + return self.getFn(self.context, domain, qtype); + } + + pub fn put(self: Cache, domain: []const u8, qtype: types.QType, response: []const u8, ttl: u32) void { + self.putFn(self.context, domain, qtype, response, ttl); + } +}; + +/// Interface for query logging +pub const Logger = struct { + context: *anyopaque, + logFn: *const fn (*anyopaque, QueryLogEntry) void, + + pub fn log(self: Logger, entry: QueryLogEntry) void { + self.logFn(self.context, entry); + } +}; + +/// Main DNS request handler +pub const Handler = struct { + denylist: ?Denylist, + cache: ?Cache, + upstream: ?Upstream, + logger: ?Logger, + limiter: ?*rate_limiter.RateLimiter, + config: HandlerConfig, + allocator: Allocator, + + pub fn init(allocator: Allocator) Handler { + return Handler{ + .denylist = null, + .cache = null, + .upstream = null, + .logger = null, + .limiter = null, + .config = HandlerConfig{}, + .allocator = allocator, + }; + } + + pub fn setDenylist(self: *Handler, denylist: Denylist) void { + self.denylist = denylist; + } + + pub fn setCache(self: *Handler, cache: Cache) void { + self.cache = cache; + } + + pub fn setUpstream(self: *Handler, upstream: Upstream) void { + self.upstream = upstream; + } + + pub fn setLogger(self: *Handler, logger: Logger) void { + self.logger = logger; + } + + pub fn setRateLimiter(self: *Handler, limiter: *rate_limiter.RateLimiter) void { + self.limiter = limiter; + } + + /// Handle a DNS query with protocol info + pub fn handleWithProtocol( + self: *Handler, + query_bytes: []const u8, + client_addr: std.net.Address, + protocol: schema.ClientProtocol, + allocator: Allocator, + ) ?[]const u8 { + const start_time = std.time.microTimestamp(); + + // Parse the query + var query = packet.Packet.parse(query_bytes, allocator) catch |err| { + std.log.warn("Failed to parse DNS query: {}", .{err}); + return self.createErrorResponse(query_bytes, types.RCode.FormErr, allocator); + }; + defer query.deinit(); + + // Must have at least one question + if (query.questions.len == 0) { + return self.createErrorResponse(query_bytes, types.RCode.FormErr, allocator); + } + + const question = query.questions[0]; + + // Get domain name as string (stack buffer, no allocation) + var domain_buf: [254]u8 = undefined; + const domain = question.name.toStringBuf(&domain_buf) orelse { + return self.createErrorResponse(query_bytes, types.RCode.ServFail, allocator); + }; + + // Get client IP as string + var client_ip_buf: [45]u8 = undefined; + const client_ip = formatAddress(client_addr, &client_ip_buf); + + // Check rate limit + if (self.limiter) |limiter| { + if (!limiter.checkRequest(client_ip)) { + std.log.debug("Rate limited client: {s}", .{client_ip}); + // Return REFUSED for rate-limited clients + return self.createErrorResponse(query_bytes, types.RCode.Refused, allocator); + } + } + + // Get client's groups (Gap 4: many-to-many - returns bitmask) + var groups_mask: u64 = 1; // Default group (bit 0) + if (self.denylist) |dl| { + groups_mask = dl.getGroupsForClient(client_ip); + } + + // Check denylist with full attribution (Gap 4: uses groups bitmask) + if (self.denylist) |dl| { + const result = dl.checkWithMask(domain, groups_mask); + + if (result.denied) { + const response = self.createDeniedResponsePacket(&query, allocator) orelse return null; + + // Determine status based on whether it was a denylist or rule + const status: schema.QueryStatus = if (result.rule_id != null) + .denied_rule + else + .denied_denylist; + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = status, + .list_id = result.list_id, + .rule_id = result.rule_id, + .reply_type = .ip, // Returns 0.0.0.0 or NXDOMAIN + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = null, + .reason = "denylist", + .dnssec_validated = false, + }); + + return response; + } else if (result.is_rule_allow and result.rule_id != null) { + // Domain was on denylist but explicitly allowed by rule + // Continue to forward, but log the allow rule + } + } + + // Check safe search enforcement + if (self.config.safe_search_enabled) { + if (safesearch.applySafeSearch(domain)) |safe_domain| { + const response = self.createSafeSearchResponsePacket(&query, safe_domain, allocator) orelse return null; + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .denied_upstream, // Safe search is a kind of upstream block + .list_id = null, + .rule_id = null, + .reply_type = .cname, + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = null, + .reason = "safesearch", + .dnssec_validated = false, + }); + + return response; + } + } + + // Check cache + if (self.cache) |cache| { + if (cache.getCopy(domain, question.qtype)) |cached| { + // Update transaction ID in cached response + // cached is a mutable copy owned by caller + if (cached.len < 2) { + allocator.free(cached); + return null; + } + cached[0] = @intCast((query.header.id >> 8) & 0xFF); + cached[1] = @intCast(query.header.id & 0xFF); + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .cached, + .list_id = null, + .rule_id = null, + .reply_type = .ip, // Assume IP for cached + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = null, + .reason = "cache", + .dnssec_validated = false, // TODO: could store/retrieve from cache + }); + + return cached; + } + } + + // Forward to upstream + if (self.upstream) |upstream| { + const response = upstream.query(query_bytes, allocator) orelse { + std.log.warn("Upstream query failed for {s}", .{domain}); + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .upstream_error, + .list_id = null, + .rule_id = null, + .reply_type = .servfail, + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = null, + .reason = "upstream_failed", + .dnssec_validated = false, + }); + + return self.createErrorResponse(query_bytes, types.RCode.ServFail, allocator); + }; + + // Validate response is a valid DNS packet before processing + var response_pkt = packet.Packet.parse(response, allocator) catch |err| { + std.log.warn("Invalid upstream response for {s}: {}", .{ domain, err }); + allocator.free(response); + return self.createErrorResponse(query_bytes, types.RCode.ServFail, allocator); + }; + + defer response_pkt.deinit(); + + // Determine reply type from response + const reply_type = self.detectReplyType(&response_pkt); + + // Check for CNAME uncloaking - detect denied domains hiding behind CNAMEs + const dl = self.denylist orelse { + // No denylist - cache response and return it + if (self.cache) |cache| { + if (response_pkt.answers.len > 0) { + cache.put(domain, question.qtype, response, response_pkt.answers[0].ttl); + } + } + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .forwarded, + .list_id = null, + .rule_id = null, + .reply_type = reply_type, + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = "upstream", + .reason = null, + .dnssec_validated = response_pkt.header.ad, + }); + return response; + }; + + // Check if any CNAME target is denied + if (self.findDeniedCname(&response_pkt, dl, groups_mask)) |_| { + allocator.free(response); + const denied_response = self.createDeniedResponsePacket(&query, allocator) orelse return null; + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .denied_denylist, + .list_id = null, // TODO: Could track which list blocked the CNAME + .rule_id = null, + .reply_type = .ip, + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = null, + .reason = "cname_uncloaking", + .dnssec_validated = false, + }); + + return denied_response; + } + + // Cache the response + if (self.cache) |cache| { + if (response_pkt.answers.len > 0) { + cache.put(domain, question.qtype, response, response_pkt.answers[0].ttl); + } + } + + self.logQuery(.{ + .timestamp = std.time.timestamp(), + .domain = domain, + .client_ip = client_ip, + .qtype = question.qtype, + .status = .forwarded, + .list_id = null, + .rule_id = null, + .reply_type = reply_type, + .protocol = protocol, + .response_time_us = calcElapsedMicros(start_time), + .upstream = "upstream", + .reason = null, + .dnssec_validated = response_pkt.header.ad, + }); + + return response; + } + + // No upstream configured + return self.createErrorResponse(query_bytes, types.RCode.ServFail, allocator); + } + + /// Handle a DNS query (legacy - defaults to UDP protocol) + pub fn handle( + self: *Handler, + query_bytes: []const u8, + client_addr: std.net.Address, + allocator: Allocator, + ) ?[]const u8 { + return self.handleWithProtocol(query_bytes, client_addr, .udp, allocator); + } + + /// Detect reply type from response packet + fn detectReplyType(self: *Handler, response_pkt: *packet.Packet) schema.ReplyType { + _ = self; + + // Check RCODE first + switch (response_pkt.header.rcode) { + .ServFail => return .servfail, + .NXDomain => return .nxdomain, + .Refused => return .refused, + else => {}, + } + + // No answers = NODATA + if (response_pkt.answers.len == 0) { + return .nodata; + } + + // Check answer types + for (response_pkt.answers) |answer| { + if (answer.rtype == types.QType.A or answer.rtype == types.QType.AAAA) { + return .ip; + } + if (answer.rtype == types.QType.CNAME) { + return .cname; + } + } + + return .unknown; + } + + /// Check if any CNAME target in the response is denied + /// Checks ANSWER, AUTHORITY, and ADDITIONAL sections for completeness + /// Returns true if a denied CNAME was found + /// Gap 4: Takes groups bitmask for many-to-many support + fn findDeniedCname(self: *Handler, response_pkt: *packet.Packet, dl: Denylist, groups_mask: u64) ?bool { + _ = self; + + // Check all sections where CNAMEs could appear + const sections = [_][]const packet.ResourceRecord{ + response_pkt.answers, + response_pkt.authority, + response_pkt.additional, + }; + + for (sections) |section| { + for (section) |record| { + if (record.rtype != types.QType.CNAME) continue; + + const cname = record.getCname() orelse continue; + var cname_buf: [254]u8 = undefined; + const cname_str = cname.toStringBuf(&cname_buf) orelse continue; + + if (dl.checkWithMask(cname_str, groups_mask).denied) { + return true; + } + } + } + return null; + } + + /// Calculate elapsed microseconds since start_time, handling potential negative values + fn calcElapsedMicros(start_time: i64) u64 { + const elapsed = std.time.microTimestamp() - start_time; + return if (elapsed > 0) @intCast(elapsed) else 0; + } + + fn createDeniedResponsePacket(self: *Handler, query: *const packet.Packet, allocator: Allocator) ?[]const u8 { + var response = switch (self.config.blocking_response) { + .zero => packet.Packet.createDeniedResponse(query, allocator) catch return null, + .nxdomain => packet.Packet.createNxdomainResponse(query, allocator) catch return null, + }; + defer response.deinit(); + + var buffer: [types.EDNS_DEFAULT_SIZE]u8 = undefined; + const len = response.encode(&buffer) catch return null; + + return allocator.dupe(u8, buffer[0..len]) catch null; + } + + fn createSafeSearchResponsePacket(self: *Handler, query: *const packet.Packet, safe_domain: []const u8, allocator: Allocator) ?[]const u8 { + _ = self; + var response = packet.Packet.createSafeSearchResponse(query, safe_domain, allocator) catch return null; + defer response.deinit(); + + var buffer: [types.EDNS_DEFAULT_SIZE]u8 = undefined; + const len = response.encode(&buffer) catch return null; + + return allocator.dupe(u8, buffer[0..len]) catch null; + } + + fn createErrorResponse(self: *Handler, query_bytes: []const u8, rcode: types.RCode, allocator: Allocator) ?[]const u8 { + _ = self; + if (query_bytes.len < types.DNS_HEADER_SIZE) return null; + + // Create a minimal error response + var response_buf: [types.DNS_HEADER_SIZE]u8 = undefined; + @memcpy(&response_buf, query_bytes[0..types.DNS_HEADER_SIZE]); + + // Set QR = 1 (response), RA = 1 (recursion available), and RCODE + response_buf[2] |= 0x80; // QR = 1 + response_buf[3] = (response_buf[3] & 0xF0) | @intFromEnum(rcode); + response_buf[3] |= 0x80; // RA = 1 + + // Set counts to 0 for answers, authority, additional + response_buf[6] = 0; + response_buf[7] = 0; + response_buf[8] = 0; + response_buf[9] = 0; + response_buf[10] = 0; + response_buf[11] = 0; + + return allocator.dupe(u8, &response_buf) catch null; + } + + fn logQuery(self: *Handler, entry: QueryLogEntry) void { + if (self.logger) |logger| { + logger.log(entry); + } + } + + /// Create a UDP handler wrapper + pub fn toUdpHandler(self: *Handler) udp.UdpServer.Handler { + return udp.UdpServer.Handler{ + .context = self, + .handleFn = handleUdp, + }; + } + + /// Create a TCP handler wrapper + pub fn toTcpHandler(self: *Handler) tcp.TcpServer.Handler { + return tcp.TcpServer.Handler{ + .context = self, + .handleFn = handleTcp, + }; + } + + fn handleUdp(ctx: *anyopaque, query: []const u8, client_addr: std.net.Address, allocator: Allocator) ?[]const u8 { + const self: *Handler = @ptrCast(@alignCast(ctx)); + return self.handleWithProtocol(query, client_addr, .udp, allocator); + } + + fn handleTcp(ctx: *anyopaque, query: []const u8, client_addr: std.net.Address, allocator: Allocator) ?[]const u8 { + const self: *Handler = @ptrCast(@alignCast(ctx)); + return self.handleWithProtocol(query, client_addr, .tcp, allocator); + } +}; + +/// Format a network address to a string representation +/// Returns "0.0.0.0" as fallback for unknown address families (safer than "unknown" for rate limiting) +fn formatAddress(addr: std.net.Address, buf: []u8) []const u8 { + if (addr.any.family == std.posix.AF.INET) { + // IPv4 address + const bytes = @as(*const [4]u8, @ptrCast(&addr.in.sa.addr)); + const result = std.fmt.bufPrint(buf, "{d}.{d}.{d}.{d}", .{ + bytes[0], bytes[1], bytes[2], bytes[3], + }) catch return "0.0.0.0"; + return result; + } else if (addr.any.family == std.posix.AF.INET6) { + // IPv6 address - use compressed format for common cases + const bytes = @as(*const [16]u8, @ptrCast(&addr.in6.sa.addr)); + const result = std.fmt.bufPrint(buf, "{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}:{x:0>2}{x:0>2}", .{ + bytes[0], bytes[1], bytes[2], bytes[3], + bytes[4], bytes[5], bytes[6], bytes[7], + bytes[8], bytes[9], bytes[10], bytes[11], + bytes[12], bytes[13], bytes[14], bytes[15], + }) catch return "::"; + return result; + } + // Unknown address family - use safe fallback that won't break rate limiting + std.log.warn("Unknown address family: {}", .{addr.any.family}); + return "0.0.0.0"; +} + +test "Handler basic test" { + const testing = std.testing; + const allocator = testing.allocator; + + var handler = Handler.init(allocator); + + // Create a simple query + const query_bytes = [_]u8{ + 0x00, 0x01, // ID + 0x01, 0x00, // Standard query, RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', // example + 0x03, 'c', 'o', 'm', // com + 0x00, // null + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; + + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + // Without upstream, should return SERVFAIL + const response = handler.handle(&query_bytes, addr, allocator); + + if (response) |r| { + defer allocator.free(r); + try testing.expect(r.len >= types.DNS_HEADER_SIZE); + } +} diff --git a/src/server/rate_limiter.zig b/src/server/rate_limiter.zig new file mode 100644 index 0000000..9ec0ab7 --- /dev/null +++ b/src/server/rate_limiter.zig @@ -0,0 +1,241 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; + +/// Rate limiter to prevent DNS amplification attacks +/// Uses a sliding window algorithm with per-IP tracking +pub const RateLimiter = struct { + /// Per-IP request tracking + clients: std.StringHashMapUnmanaged(ClientState), + allocator: Allocator, + mutex: std.Thread.Mutex, + + /// Configuration + config: Config, + + /// Statistics + stats: Stats, + + pub const Config = struct { + /// Maximum requests per second per IP + max_qps: u32 = 20, + /// Time window for rate limiting (in milliseconds) + window_ms: u64 = 1000, + /// Maximum number of IPs to track (LRU eviction beyond this) + max_clients: usize = 10000, + /// Enable rate limiting + enabled: bool = true, + }; + + const ClientState = struct { + /// Request timestamps in the current window (circular buffer) + timestamps: [32]i64, + /// Current position in the buffer + pos: usize, + /// Number of requests in current window + count: u32, + /// Last activity timestamp (for LRU eviction) + last_seen: i64, + }; + + pub const Stats = struct { + /// Total requests processed + total_requests: u64 = 0, + /// Requests that were rate limited + rate_limited: u64 = 0, + /// Currently tracked clients + active_clients: usize = 0, + }; + + pub fn init(allocator: Allocator) RateLimiter { + return initWithConfig(allocator, Config{}); + } + + pub fn initWithConfig(allocator: Allocator, config: Config) RateLimiter { + return RateLimiter{ + .clients = std.StringHashMapUnmanaged(ClientState){}, + .allocator = allocator, + .mutex = std.Thread.Mutex{}, + .config = config, + .stats = Stats{}, + }; + } + + pub fn deinit(self: *RateLimiter) void { + var iter = self.clients.iterator(); + while (iter.next()) |entry| { + self.allocator.free(entry.key_ptr.*); + } + self.clients.deinit(self.allocator); + } + + /// Check if a request from this IP should be allowed + /// Returns true if allowed, false if rate limited + pub fn checkRequest(self: *RateLimiter, client_ip: []const u8) bool { + if (!self.config.enabled) return true; + + self.mutex.lock(); + defer self.mutex.unlock(); + + self.stats.total_requests += 1; + + const now = std.time.milliTimestamp(); + const window_start = now - @as(i64, @intCast(self.config.window_ms)); + + // Get or create client state + const gop = self.clients.getOrPut(self.allocator, client_ip) catch { + // On allocation failure, allow the request + return true; + }; + + if (!gop.found_existing) { + // New client - copy the key + gop.key_ptr.* = self.allocator.dupe(u8, client_ip) catch { + _ = self.clients.remove(client_ip); + return true; + }; + + // Initialize state + gop.value_ptr.* = ClientState{ + .timestamps = [_]i64{0} ** 32, + .pos = 0, + .count = 0, + .last_seen = now, + }; + + // Evict oldest client if at capacity + if (self.clients.count() > self.config.max_clients) { + self.evictOldestLocked(); + } + } + + var state = gop.value_ptr; + state.last_seen = now; + + // Count requests in current window + var count: u32 = 0; + for (state.timestamps) |ts| { + if (ts > window_start) { + count += 1; + } + } + + // Check rate limit + if (count >= self.config.max_qps) { + self.stats.rate_limited += 1; + return false; + } + + // Record this request + state.timestamps[state.pos] = now; + state.pos = (state.pos + 1) % state.timestamps.len; + state.count = count + 1; + + self.stats.active_clients = self.clients.count(); + + return true; + } + + /// Evict the oldest (least recently seen) client + fn evictOldestLocked(self: *RateLimiter) void { + var oldest_key: ?[]const u8 = null; + var oldest_time: i64 = std.math.maxInt(i64); + + var iter = self.clients.iterator(); + while (iter.next()) |entry| { + if (entry.value_ptr.last_seen < oldest_time) { + oldest_time = entry.value_ptr.last_seen; + oldest_key = entry.key_ptr.*; + } + } + + if (oldest_key) |key| { + if (self.clients.fetchRemove(key)) |kv| { + self.allocator.free(kv.key); + } + } + } + + /// Get current statistics + pub fn getStats(self: *RateLimiter) Stats { + self.mutex.lock(); + defer self.mutex.unlock(); + return self.stats; + } + + /// Reset statistics + pub fn resetStats(self: *RateLimiter) void { + self.mutex.lock(); + defer self.mutex.unlock(); + self.stats = Stats{}; + } + + /// Clear all tracked clients + pub fn clear(self: *RateLimiter) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + var iter = self.clients.iterator(); + while (iter.next()) |entry| { + self.allocator.free(entry.key_ptr.*); + } + self.clients.clearRetainingCapacity(); + self.stats.active_clients = 0; + } +}; + +test "RateLimiter basic operations" { + const testing = std.testing; + const allocator = testing.allocator; + + var limiter = RateLimiter.initWithConfig(allocator, .{ + .max_qps = 5, + .window_ms = 1000, + }); + defer limiter.deinit(); + + // First 5 requests should be allowed + for (0..5) |_| { + try testing.expect(limiter.checkRequest("192.168.1.1")); + } + + // 6th request should be rate limited + try testing.expect(!limiter.checkRequest("192.168.1.1")); + + // Different IP should still be allowed + try testing.expect(limiter.checkRequest("192.168.1.2")); +} + +test "RateLimiter disabled" { + const testing = std.testing; + const allocator = testing.allocator; + + var limiter = RateLimiter.initWithConfig(allocator, .{ + .max_qps = 1, + .enabled = false, + }); + defer limiter.deinit(); + + // All requests should be allowed when disabled + for (0..100) |_| { + try testing.expect(limiter.checkRequest("192.168.1.1")); + } +} + +test "RateLimiter statistics" { + const testing = std.testing; + const allocator = testing.allocator; + + var limiter = RateLimiter.initWithConfig(allocator, .{ + .max_qps = 2, + .window_ms = 1000, + }); + defer limiter.deinit(); + + _ = limiter.checkRequest("192.168.1.1"); // allowed + _ = limiter.checkRequest("192.168.1.1"); // allowed + _ = limiter.checkRequest("192.168.1.1"); // rate limited + + const stats = limiter.getStats(); + try testing.expectEqual(@as(u64, 3), stats.total_requests); + try testing.expectEqual(@as(u64, 1), stats.rate_limited); +} diff --git a/src/server/shutdown.zig b/src/server/shutdown.zig new file mode 100644 index 0000000..88edd2b --- /dev/null +++ b/src/server/shutdown.zig @@ -0,0 +1,117 @@ +const std = @import("std"); +const posix = std.posix; +const linux = std.os.linux; + +/// Unified shutdown coordinator using signalfd + eventfd +/// - signalfd: catches SIGINT/SIGTERM (main thread waits on this) +/// - eventfd: wakes all worker threads when shutdown requested +pub const ShutdownCoordinator = struct { + shutdown_requested: std.atomic.Value(bool), + signal_fd: posix.fd_t, + event_fd: posix.fd_t, + + pub fn init() ShutdownCoordinator { + // Block SIGINT/SIGTERM so they go to signalfd instead of default handler + var mask = std.mem.zeroes(linux.sigset_t); + linux.sigaddset(&mask, linux.SIG.INT); + linux.sigaddset(&mask, linux.SIG.TERM); + _ = linux.sigprocmask(linux.SIG.BLOCK, &mask, null); + + // Create signalfd + const sig_fd_raw = linux.signalfd(-1, &mask, linux.SFD.CLOEXEC); + const sig_fd: posix.fd_t = if (@as(isize, @bitCast(sig_fd_raw)) < 0) -1 else @intCast(sig_fd_raw); + + // Create eventfd for waking threads + const evt_fd = posix.eventfd(0, linux.EFD.CLOEXEC) catch -1; + + return ShutdownCoordinator{ + .shutdown_requested = std.atomic.Value(bool).init(false), + .signal_fd = sig_fd, + .event_fd = evt_fd, + }; + } + + pub fn deinit(self: *ShutdownCoordinator) void { + if (self.signal_fd != -1) posix.close(self.signal_fd); + if (self.event_fd != -1) posix.close(self.event_fd); + } + + /// Signal shutdown and wake all waiting threads + pub fn requestShutdown(self: *ShutdownCoordinator) void { + self.shutdown_requested.store(true, .release); + + // Wake all threads waiting on eventfd + if (self.event_fd != -1) { + const val: u64 = 1; + _ = posix.write(self.event_fd, std.mem.asBytes(&val)) catch {}; + } + } + + pub fn isShutdownRequested(self: *ShutdownCoordinator) bool { + return self.shutdown_requested.load(.acquire); + } + + /// Get eventfd for polling (worker threads use this) + pub fn getEventFd(self: *ShutdownCoordinator) posix.fd_t { + return self.event_fd; + } + + /// Get signalfd for polling (main thread uses this) + pub fn getSignalFd(self: *ShutdownCoordinator) posix.fd_t { + return self.signal_fd; + } + + /// Block until shutdown signal received (main thread calls this) + /// Returns true if signal was received, false on error + pub fn waitForSignal(self: *ShutdownCoordinator) bool { + if (self.signal_fd == -1) return false; + + var fds = [1]posix.pollfd{ + .{ .fd = self.signal_fd, .events = posix.POLL.IN, .revents = 0 }, + }; + + _ = posix.poll(&fds, -1) catch return false; + + if (fds[0].revents & posix.POLL.IN != 0) { + // Consume the signal + var siginfo: linux.signalfd_siginfo = undefined; + _ = posix.read(self.signal_fd, std.mem.asBytes(&siginfo)) catch {}; + self.requestShutdown(); + return true; + } + + return false; + } + + /// Block until shutdown is requested or timeout (worker threads use this) + /// timeout_ms: -1 for infinite wait + pub fn wait(self: *ShutdownCoordinator, timeout_ms: i32) bool { + if (self.isShutdownRequested()) return true; + if (self.event_fd == -1) return false; + + var fds = [1]posix.pollfd{ + .{ .fd = self.event_fd, .events = posix.POLL.IN, .revents = 0 }, + }; + + _ = posix.poll(&fds, timeout_ms) catch return self.isShutdownRequested(); + return self.isShutdownRequested(); + } +}; + +test "ShutdownCoordinator basic" { + var coordinator = ShutdownCoordinator.init(); + defer coordinator.deinit(); + + try std.testing.expect(!coordinator.isShutdownRequested()); + coordinator.requestShutdown(); + try std.testing.expect(coordinator.isShutdownRequested()); +} + +test "ShutdownCoordinator wait returns immediately after shutdown" { + var coordinator = ShutdownCoordinator.init(); + defer coordinator.deinit(); + + coordinator.requestShutdown(); + const result = coordinator.wait(1000); + try std.testing.expect(result); +} diff --git a/src/server/tcp.zig b/src/server/tcp.zig new file mode 100644 index 0000000..584d54e --- /dev/null +++ b/src/server/tcp.zig @@ -0,0 +1,373 @@ +const std = @import("std"); +const posix = std.posix; +const net = std.net; +const Allocator = std.mem.Allocator; +const packet = @import("../dns/packet.zig"); +const types = @import("../dns/types.zig"); + +/// Connection task for the worker pool +const ConnectionTask = struct { + handle: posix.socket_t, + address: net.Address, +}; + +/// Simple bounded work queue for connection tasks +const WorkQueue = struct { + items: [QUEUE_SIZE]?ConnectionTask, + head: usize, + tail: usize, + count: usize, + mutex: std.Thread.Mutex, + not_empty: std.Thread.Condition, + not_full: std.Thread.Condition, + + const QUEUE_SIZE = 64; + + fn init() WorkQueue { + return .{ + .items = [_]?ConnectionTask{null} ** QUEUE_SIZE, + .head = 0, + .tail = 0, + .count = 0, + .mutex = .{}, + .not_empty = .{}, + .not_full = .{}, + }; + } + + /// Push a task, returns false if queue is full + fn tryPush(self: *WorkQueue, task: ConnectionTask) 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; + } + + /// Pop a task, blocks until available or timeout + fn pop(self: *WorkQueue, running: *std.atomic.Value(bool)) ?ConnectionTask { + self.mutex.lock(); + defer self.mutex.unlock(); + + while (self.count == 0) { + if (!running.load(.acquire)) { + return null; + } + // Wait with timeout to periodically check running flag + 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; + self.not_full.signal(); + return task; + } + + fn wakeAll(self: *WorkQueue) void { + self.mutex.lock(); + defer self.mutex.unlock(); + self.not_empty.broadcast(); + } +}; + +pub const TcpServer = struct { + listener: net.Server, + allocator: Allocator, + handler: *Handler, + running: std.atomic.Value(bool), + active_connections: std.atomic.Value(u32), + max_connections: u32, + connection_timeout_ms: u32, + work_queue: WorkQueue, + workers: []std.Thread, + num_workers: u32, + + /// Default maximum concurrent TCP connections + pub const DEFAULT_MAX_CONNECTIONS: u32 = 100; + + /// Default connection read/write timeout (30 seconds) + pub const DEFAULT_CONNECTION_TIMEOUT_MS: u32 = 30000; + + /// 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{ + ListenFailed, + } || net.Address.ListenError; + + /// Initialize the TCP server + pub fn init(bind_addr: net.Address, handler: *Handler, allocator: Allocator) InitError!TcpServer { + return initWithConfig(bind_addr, handler, allocator, .{}); + } + + pub const Config = struct { + max_connections: u32 = DEFAULT_MAX_CONNECTIONS, + connection_timeout_ms: u32 = DEFAULT_CONNECTION_TIMEOUT_MS, + num_workers: u32 = DEFAULT_NUM_WORKERS, + }; + + /// Initialize the TCP server with custom configuration + pub fn initWithConfig(bind_addr: net.Address, handler: *Handler, allocator: Allocator, config: Config) InitError!TcpServer { + const listener = bind_addr.listen(.{ + .reuse_address = true, + }) catch { + return error.ListenFailed; + }; + + return TcpServer{ + .listener = listener, + .allocator = allocator, + .handler = handler, + .running = std.atomic.Value(bool).init(false), + .active_connections = std.atomic.Value(u32).init(0), + .max_connections = config.max_connections, + .connection_timeout_ms = config.connection_timeout_ms, + .work_queue = WorkQueue.init(), + .workers = &[_]std.Thread{}, + .num_workers = config.num_workers, + }; + } + + /// Start the server loop + pub fn run(self: *TcpServer) !void { + self.running.store(true, .release); + + // Start worker threads + self.workers = self.allocator.alloc(std.Thread, self.num_workers) catch |err| { + std.log.err("TCP: 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("TCP: failed to start worker thread: {}", .{err}); + return error.ThreadSpawnFailed; + }; + started += 1; + } + + std.log.info("TCP: started {} worker threads", .{self.num_workers}); + + while (self.running.load(.acquire)) { + // Use poll with timeout to allow checking running flag + var fds = [1]posix.pollfd{ + .{ + .fd = self.listener.stream.handle, + .events = posix.POLL.IN, + .revents = 0, + }, + }; + + const poll_result = posix.poll(&fds, 100) catch |err| { + std.log.warn("TCP poll error: {}", .{err}); + continue; + }; + + // Timeout - check running flag and continue + if (poll_result == 0) continue; + + // No connection pending + if (fds[0].revents & posix.POLL.IN == 0) continue; + + const conn = self.listener.accept() catch |err| { + std.log.warn("TCP accept error: {}", .{err}); + continue; + }; + + // Reserve a connection slot atomically (prevents TOCTOU race) + // fetchAdd returns the OLD value, so new count is old + 1 + const old_count = self.active_connections.fetchAdd(1, .acq_rel); + if (old_count >= self.max_connections) { + // Over limit - rollback reservation and reject + _ = self.active_connections.fetchSub(1, .release); + std.log.warn("TCP: max connections ({}) reached, rejecting new connection", .{self.max_connections}); + posix.close(conn.stream.handle); + continue; + } + + // Submit to worker pool + const task = ConnectionTask{ + .handle = conn.stream.handle, + .address = conn.address, + }; + + if (!self.work_queue.tryPush(task)) { + // Queue full - rollback reservation and reject + _ = self.active_connections.fetchSub(1, .release); + std.log.warn("TCP: work queue full, rejecting connection", .{}); + posix.close(conn.stream.handle); + } + } + + // 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: *TcpServer) void { + while (self.running.load(.acquire)) { + if (self.work_queue.pop(&self.running)) |task| { + self.handleConnectionTask(task); + } + } + } + + /// Maximum DNS message size for TCP (RFC 1035: 2-byte length prefix allows up to 65535) + const MAX_DNS_MESSAGE_SIZE: usize = 65535; + + fn handleConnectionTask(self: *TcpServer, task: ConnectionTask) void { + // Count was already incremented in accept loop when slot was reserved + defer _ = self.active_connections.fetchSub(1, .release); + defer posix.close(task.handle); + + // Set read/write timeout on the connection to prevent slow clients from blocking + const timeout = posix.timeval{ + .sec = @intCast(self.connection_timeout_ms / 1000), + .usec = @intCast((self.connection_timeout_ms % 1000) * 1000), + }; + posix.setsockopt(task.handle, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch {}; + posix.setsockopt(task.handle, posix.SOL.SOCKET, posix.SO.SNDTIMEO, std.mem.asBytes(&timeout)) catch {}; + + var buffer: [MAX_DNS_MESSAGE_SIZE]u8 = undefined; + + while (self.running.load(.acquire)) { + // Read 2-byte length prefix using posix + var len_buf: [2]u8 = undefined; + if (!readExact(task.handle, &len_buf)) break; + + const length = std.mem.readInt(u16, &len_buf, .big); + // Validate: must be at least a DNS header, and within buffer capacity + if (length < types.DNS_HEADER_SIZE) { + std.log.debug("TCP: message too small ({} bytes), dropping connection", .{length}); + break; + } + if (length > MAX_DNS_MESSAGE_SIZE) { + std.log.warn("TCP: message too large ({} bytes), dropping connection", .{length}); + break; + } + + // Read DNS message using posix + if (!readExact(task.handle, buffer[0..length])) break; + + // Handle the query + const response = self.handler.handle( + buffer[0..length], + task.address, + self.allocator, + ) orelse continue; + defer self.allocator.free(response); + + // Validate response fits in TCP DNS message (u16 length prefix) + if (response.len > MAX_DNS_MESSAGE_SIZE) { + std.log.warn("TCP: response too large ({} bytes), dropping", .{response.len}); + continue; + } + + // Write length-prefixed response + var resp_len: [2]u8 = undefined; + std.mem.writeInt(u16, &resp_len, @intCast(response.len), .big); + + if (!writeAll(task.handle, &resp_len)) break; + if (!writeAll(task.handle, response)) break; + } + } + + /// Stop the server + pub fn stop(self: *TcpServer) void { + self.running.store(false, .release); + self.work_queue.wakeAll(); + } + + /// Wait for active connections to finish (with timeout) + pub fn waitForConnections(self: *TcpServer, timeout_ms: u64) void { + const start = std.time.milliTimestamp(); + while (self.active_connections.load(.acquire) > 0) { + const elapsed: u64 = @intCast(std.time.milliTimestamp() - start); + if (elapsed >= timeout_ms) { + std.log.warn("TCP: {} connections still active after timeout", .{self.active_connections.load(.acquire)}); + break; + } + std.posix.nanosleep(0, 10 * std.time.ns_per_ms); + } + } + + /// Clean up resources + pub fn deinit(self: *TcpServer) void { + self.listener.deinit(); + } +}; + +/// Read exactly the requested number of bytes using posix +fn readExact(handle: posix.socket_t, buf: []u8) bool { + var total_read: usize = 0; + while (total_read < buf.len) { + const n = posix.read(handle, buf[total_read..]) catch return false; + if (n == 0) return false; // EOF + total_read += n; + } + return true; +} + +/// Write all bytes using posix +fn writeAll(handle: posix.socket_t, buf: []const u8) bool { + var total_written: usize = 0; + while (total_written < buf.len) { + const n = posix.write(handle, buf[total_written..]) catch return false; + if (n == 0) return false; + total_written += n; + } + return true; +} + +test "TcpServer init" { + const testing = std.testing; + const allocator = testing.allocator; + + const Handler = TcpServer.Handler; + var handler = Handler{ + .context = undefined, + .handleFn = struct { + fn handle(_: *anyopaque, _: []const u8, _: std.net.Address, _: Allocator) ?[]const u8 { + return null; + } + }.handle, + }; + + const addr = net.Address.initIp4(.{ 127, 0, 0, 1 }, 0); + var server = TcpServer.init(addr, &handler, allocator) catch |err| { + std.debug.print("TcpServer init failed: {}\n", .{err}); + return error.TestFailed; + }; + defer server.deinit(); +} diff --git a/src/server/udp.zig b/src/server/udp.zig new file mode 100644 index 0000000..a7f6b3e --- /dev/null +++ b/src/server/udp.zig @@ -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); +} diff --git a/src/storage/db.zig b/src/storage/db.zig new file mode 100644 index 0000000..1e9d895 --- /dev/null +++ b/src/storage/db.zig @@ -0,0 +1,305 @@ +const std = @import("std"); +const c = @cImport({ + @cInclude("sqlite3.h"); +}); +const Allocator = std.mem.Allocator; + +/// SQLite Database wrapper +pub const Database = struct { + conn: *c.sqlite3, + allocator: Allocator, + /// Mutex to serialize transaction access (SQLite doesn't support nested transactions) + tx_mutex: std.Thread.Mutex = .{}, + + pub const Error = error{ + OpenFailed, + ExecuteFailed, + PrepareFailed, + BindFailed, + StepFailed, + ColumnError, + OutOfMemory, + }; + + /// Open a database file + pub fn open(path: []const u8, allocator: Allocator) Error!Database { + var conn: ?*c.sqlite3 = null; + + // Create null-terminated path + const path_z = allocator.dupeZ(u8, path) catch return error.OutOfMemory; + defer allocator.free(path_z); + + const result = c.sqlite3_open(path_z.ptr, &conn); + if (result != c.SQLITE_OK) { + if (conn) |cn| { + _ = c.sqlite3_close(cn); + } + return error.OpenFailed; + } + + // Enable WAL mode for better concurrency + var stmt: ?*c.sqlite3_stmt = null; + _ = c.sqlite3_prepare_v2(conn, "PRAGMA journal_mode=WAL", -1, &stmt, null); + if (stmt) |s| { + _ = c.sqlite3_step(s); + _ = c.sqlite3_finalize(s); + } + + // Enable foreign keys + _ = c.sqlite3_prepare_v2(conn, "PRAGMA foreign_keys=ON", -1, &stmt, null); + if (stmt) |s| { + _ = c.sqlite3_step(s); + _ = c.sqlite3_finalize(s); + } + + // Set busy timeout to handle concurrent write contention + // SQLite will retry for up to 5 seconds before returning SQLITE_BUSY + _ = c.sqlite3_prepare_v2(conn, "PRAGMA busy_timeout=5000", -1, &stmt, null); + if (stmt) |s| { + _ = c.sqlite3_step(s); + _ = c.sqlite3_finalize(s); + } + + return Database{ + .conn = conn.?, + .allocator = allocator, + }; + } + + /// Close the database + pub fn close(self: *Database) void { + _ = c.sqlite3_close(self.conn); + } + + /// Execute a SQL statement + pub fn exec(self: *Database, sql: []const u8) Error!void { + const sql_z = self.allocator.dupeZ(u8, sql) catch return error.OutOfMemory; + defer self.allocator.free(sql_z); + + var err_msg: [*c]u8 = null; + const result = c.sqlite3_exec(self.conn, sql_z.ptr, null, null, &err_msg); + + if (result != c.SQLITE_OK) { + if (err_msg) |msg| { + std.log.err("SQLite error: {s}", .{msg}); + c.sqlite3_free(msg); + } + return error.ExecuteFailed; + } + } + + /// Execute multiple SQL statements (for schema creation) + pub fn execMulti(self: *Database, sql: []const u8) Error!void { + var iter = std.mem.splitSequence(u8, sql, ";"); + while (iter.next()) |stmt_sql| { + const trimmed = std.mem.trim(u8, stmt_sql, " \t\n\r"); + if (trimmed.len == 0) continue; + + // Add back semicolon + const full_stmt = std.fmt.allocPrint(self.allocator, "{s};", .{trimmed}) catch return error.OutOfMemory; + defer self.allocator.free(full_stmt); + + self.exec(full_stmt) catch |err| { + std.log.err("Failed to execute: {s}", .{trimmed}); + return err; + }; + } + } + + /// Prepare a SQL statement + pub fn prepare(self: *Database, sql: []const u8) Error!Statement { + const sql_z = self.allocator.dupeZ(u8, sql) catch return error.OutOfMemory; + defer self.allocator.free(sql_z); + + var stmt: ?*c.sqlite3_stmt = null; + const result = c.sqlite3_prepare_v2(self.conn, sql_z.ptr, @intCast(sql_z.len), &stmt, null); + + if (result != c.SQLITE_OK or stmt == null) { + const err = c.sqlite3_errmsg(self.conn); + if (err) |msg| { + std.log.err("SQLite prepare error: {s}", .{msg}); + } + return error.PrepareFailed; + } + + return Statement{ + .stmt = stmt.?, + .allocator = self.allocator, + }; + } + + /// Get the last insert rowid + pub fn lastInsertRowId(self: *Database) i64 { + return c.sqlite3_last_insert_rowid(self.conn); + } + + /// Get the number of rows changed by the last statement + pub fn changes(self: *Database) i32 { + return c.sqlite3_changes(self.conn); + } + + /// Begin a transaction with RAII semantics - auto-rollback on scope exit if not committed + pub fn begin(self: *Database) Error!Transaction { + self.tx_mutex.lock(); + errdefer self.tx_mutex.unlock(); + try self.exec("BEGIN IMMEDIATE"); + return Transaction{ .db = self }; + } + + /// Legacy: Begin a transaction (acquires tx_mutex to prevent concurrent transactions) + /// Prefer using begin() which returns a Transaction with RAII semantics + pub fn beginTransaction(self: *Database) Error!void { + self.tx_mutex.lock(); + errdefer self.tx_mutex.unlock(); + return self.exec("BEGIN IMMEDIATE"); + } + + /// Legacy: Commit a transaction (releases tx_mutex) + /// Prefer using Transaction.commit() for RAII safety + pub fn commit(self: *Database) Error!void { + defer self.tx_mutex.unlock(); + return self.exec("COMMIT"); + } + + /// Legacy: Rollback a transaction (releases tx_mutex) + /// Prefer using Transaction.deinit() for RAII safety + pub fn rollback(self: *Database) Error!void { + defer self.tx_mutex.unlock(); + return self.exec("ROLLBACK"); + } +}; + +/// RAII Transaction - automatically rolls back if not committed +pub const Transaction = struct { + db: *Database, + committed: bool = false, + + /// Commit the transaction + pub fn commit(self: *Transaction) Database.Error!void { + try self.db.exec("COMMIT"); + self.committed = true; + } + + /// Cleanup: rollback if not committed, always release mutex + pub fn deinit(self: *Transaction) void { + if (!self.committed) { + self.db.exec("ROLLBACK") catch |err| { + std.log.err("Transaction rollback failed: {}", .{err}); + }; + } + self.db.tx_mutex.unlock(); + } +}; + +/// Prepared statement wrapper +pub const Statement = struct { + stmt: *c.sqlite3_stmt, + allocator: Allocator, + + /// Bind an integer value + pub fn bindInt(self: *Statement, index: usize, value: i64) Database.Error!void { + const result = c.sqlite3_bind_int64(self.stmt, @intCast(index), value); + if (result != c.SQLITE_OK) { + return error.BindFailed; + } + } + + /// Bind an unsigned integer value (clamps to i64 max for SQLite compatibility) + pub fn bindUint64(self: *Statement, index: usize, value: u64) Database.Error!void { + const clamped: i64 = @intCast(@min(value, @as(u64, std.math.maxInt(i64)))); + return self.bindInt(index, clamped); + } + + /// Bind a text value + pub fn bindText(self: *Statement, index: usize, value: []const u8) Database.Error!void { + const value_z = self.allocator.dupeZ(u8, value) catch return error.OutOfMemory; + defer self.allocator.free(value_z); + + const result = c.sqlite3_bind_text( + self.stmt, + @intCast(index), + value_z.ptr, + @intCast(value_z.len), // dupeZ's .len excludes the null terminator + c.SQLITE_TRANSIENT, + ); + + if (result != c.SQLITE_OK) { + return error.BindFailed; + } + } + + /// Bind a null value + pub fn bindNull(self: *Statement, index: usize) Database.Error!void { + const result = c.sqlite3_bind_null(self.stmt, @intCast(index)); + if (result != c.SQLITE_OK) { + return error.BindFailed; + } + } + + /// Execute a step (returns true if row available) + pub fn step(self: *Statement) Database.Error!bool { + const result = c.sqlite3_step(self.stmt); + if (result == c.SQLITE_ROW) { + return true; + } else if (result == c.SQLITE_DONE) { + return false; + } else { + return error.StepFailed; + } + } + + /// Get an integer column value + pub fn getInt(self: *Statement, col: usize) i64 { + return c.sqlite3_column_int64(self.stmt, @intCast(col)); + } + + /// Get a text column value + pub fn getText(self: *Statement, col: usize) ?[]const u8 { + const ptr = c.sqlite3_column_text(self.stmt, @intCast(col)); + if (ptr == null) return null; + + const len = c.sqlite3_column_bytes(self.stmt, @intCast(col)); + if (len <= 0) return ""; + + return ptr[0..@intCast(len)]; + } + + /// Check if column is null + pub fn isNull(self: *Statement, col: usize) bool { + return c.sqlite3_column_type(self.stmt, @intCast(col)) == c.SQLITE_NULL; + } + + /// Reset the statement for reuse + pub fn reset(self: *Statement) void { + _ = c.sqlite3_reset(self.stmt); + _ = c.sqlite3_clear_bindings(self.stmt); + } + + /// Finalize (close) the statement + pub fn finalize(self: *Statement) void { + _ = c.sqlite3_finalize(self.stmt); + } +}; + +test "Database open and close" { + const testing = std.testing; + const allocator = testing.allocator; + + var db = try Database.open(":memory:", allocator); + defer db.close(); + + try db.exec("CREATE TABLE test (id INTEGER PRIMARY KEY, name TEXT)"); + try db.exec("INSERT INTO test (name) VALUES ('hello')"); + + var stmt = try db.prepare("SELECT id, name FROM test WHERE id = 1"); + defer stmt.finalize(); + + const has_row = try stmt.step(); + try testing.expect(has_row); + + const id = stmt.getInt(0); + try testing.expectEqual(@as(i64, 1), id); + + const name = stmt.getText(1); + try testing.expectEqualStrings("hello", name.?); +} diff --git a/src/storage/logger.zig b/src/storage/logger.zig new file mode 100644 index 0000000..b09bfcc --- /dev/null +++ b/src/storage/logger.zig @@ -0,0 +1,388 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const Database = @import("db.zig").Database; +const handler = @import("../server/handler.zig"); +const schema = @import("schema.zig"); + +/// Query log entry for batching +pub const LogEntry = struct { + timestamp: i64, + domain: []const u8, + client_ip: []const u8, + qtype: u16, + // New status tracking (Gap 1) + status: schema.QueryStatus, + // Attribution tracking (Gap 2) + list_id: ?i64, + rule_id: ?i64, + // Reply type tracking (Gap 3) + reply_type: schema.ReplyType, + // Protocol tracking (Gap 8) + protocol: schema.ClientProtocol, + response_time_us: u64, + upstream: ?[]const u8, + reason: ?[]const u8, + dnssec_validated: bool, + + // Legacy compatibility + pub fn isDenied(self: LogEntry) bool { + return self.status.isDenied(); + } +}; + +/// Subscriber callback for real-time query notifications +pub const Subscriber = struct { + context: *anyopaque, + notifyFn: *const fn (*anyopaque, LogEntry) void, + + pub fn notify(self: Subscriber, entry: LogEntry) void { + self.notifyFn(self.context, entry); + } +}; + +/// Batched query logger - batches entries for efficient SQLite writes +pub const QueryLogger = struct { + buffer: std.ArrayListUnmanaged(LogEntry), + db: *Database, + allocator: Allocator, + mutex: std.Thread.Mutex, + last_flush: i64, + batch_size: usize, + flush_interval_ms: i64, + subscriber: ?Subscriber, + dropped_entries: std.atomic.Value(u64), + + const BATCH_SIZE: usize = 100; + const FLUSH_INTERVAL_MS: i64 = 100; + + pub fn init(db: *Database, allocator: Allocator) QueryLogger { + return QueryLogger{ + .buffer = std.ArrayListUnmanaged(LogEntry){}, + .db = db, + .allocator = allocator, + .mutex = std.Thread.Mutex{}, + .last_flush = std.time.milliTimestamp(), + .batch_size = BATCH_SIZE, + .flush_interval_ms = FLUSH_INTERVAL_MS, + .subscriber = null, + .dropped_entries = std.atomic.Value(u64).init(0), + }; + } + + /// Get count of dropped log entries (for monitoring) + pub fn getDroppedEntries(self: *QueryLogger) u64 { + return self.dropped_entries.load(.monotonic); + } + + /// Set a subscriber for real-time query notifications + pub fn setSubscriber(self: *QueryLogger, subscriber: Subscriber) void { + self.subscriber = subscriber; + } + + pub fn deinit(self: *QueryLogger) void { + // Best effort flush during shutdown - log but don't fail if it errors + self.flush() catch |err| { + std.log.warn("QueryLogger: flush failed during deinit: {}", .{err}); + }; + + // Free any remaining entries + for (self.buffer.items) |entry| { + self.allocator.free(entry.domain); + self.allocator.free(entry.client_ip); + if (entry.upstream) |u| self.allocator.free(u); + if (entry.reason) |r| self.allocator.free(r); + } + self.buffer.deinit(self.allocator); + } + + /// Log a query (non-blocking) + pub fn log(self: *QueryLogger, entry: handler.QueryLogEntry) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + // Copy strings since the original may be freed + const domain_copy = self.allocator.dupe(u8, entry.domain) catch { + _ = self.dropped_entries.fetchAdd(1, .monotonic); + return; + }; + errdefer self.allocator.free(domain_copy); + + const client_ip_copy = self.allocator.dupe(u8, entry.client_ip) catch { + _ = self.dropped_entries.fetchAdd(1, .monotonic); + self.allocator.free(domain_copy); + return; + }; + errdefer self.allocator.free(client_ip_copy); + + const upstream_copy: ?[]const u8 = if (entry.upstream) |u| + self.allocator.dupe(u8, u) catch null + else + null; + + const reason_copy: ?[]const u8 = if (entry.reason) |r| + self.allocator.dupe(u8, r) catch null + else + null; + + const log_entry = LogEntry{ + .timestamp = entry.timestamp, + .domain = domain_copy, + .client_ip = client_ip_copy, + .qtype = @intFromEnum(entry.qtype), + .status = entry.status, + .list_id = entry.list_id, + .rule_id = entry.rule_id, + .reply_type = entry.reply_type, + .protocol = entry.protocol, + .response_time_us = entry.response_time_us, + .upstream = upstream_copy, + .reason = reason_copy, + .dnssec_validated = entry.dnssec_validated, + }; + + // Notify real-time subscriber (for SSE) + if (self.subscriber) |sub| { + sub.notify(log_entry); + } + + self.buffer.append(self.allocator, log_entry) catch { + std.log.warn("QueryLogger: failed to append log entry", .{}); + self.allocator.free(domain_copy); + self.allocator.free(client_ip_copy); + if (upstream_copy) |u| self.allocator.free(u); + if (reason_copy) |r| self.allocator.free(r); + return; + }; + + const now = std.time.milliTimestamp(); + if (self.buffer.items.len >= self.batch_size or + now - self.last_flush >= self.flush_interval_ms) + { + self.flushLocked(); + } + } + + /// Flush the buffer to the database + pub fn flush(self: *QueryLogger) !void { + self.mutex.lock(); + defer self.mutex.unlock(); + self.flushLocked(); + } + + fn flushLocked(self: *QueryLogger) void { + if (self.buffer.items.len == 0) return; + + var tx = self.db.begin() catch |err| { + std.log.warn("QueryLogger: failed to begin transaction: {}", .{err}); + return; + }; + defer tx.deinit(); + + var insert_count: usize = 0; + for (self.buffer.items) |entry| { + self.insertEntry(entry) catch |err| { + std.log.debug("QueryLogger: failed to insert entry: {}", .{err}); + continue; + }; + insert_count += 1; + } + + tx.commit() catch |err| { + std.log.warn("QueryLogger: commit failed, rolling back {} entries: {}", .{ insert_count, err }); + return; + }; + + // Free and clear buffer + for (self.buffer.items) |entry| { + self.allocator.free(entry.domain); + self.allocator.free(entry.client_ip); + if (entry.upstream) |u| self.allocator.free(u); + if (entry.reason) |r| self.allocator.free(r); + } + self.buffer.clearRetainingCapacity(); + self.last_flush = std.time.milliTimestamp(); + } + + fn insertEntry(self: *QueryLogger, entry: LogEntry) !void { + // Get or create domain ID + const domain_id = try self.getOrCreateDomainId(entry.domain); + + // Get or create client ID + const client_id = try self.getOrCreateClientId(entry.client_ip); + + // Insert log entry with all new columns + var stmt = try self.db.prepare( + \\INSERT INTO query_log (timestamp, domain_id, client_id, qtype, denied, status, list_id, rule_id, reply_type, protocol, response_time_us, upstream, reason, dnssec_validated) + \\VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14) + ); + defer stmt.finalize(); + + try stmt.bindInt(1, entry.timestamp); + try stmt.bindInt(2, domain_id); + try stmt.bindInt(3, client_id); + try stmt.bindInt(4, entry.qtype); + try stmt.bindInt(5, if (entry.isDenied()) 1 else 0); // Legacy denied column + try stmt.bindInt(6, @intFromEnum(entry.status)); // New status column + if (entry.list_id) |lid| { + try stmt.bindInt(7, lid); + } else { + try stmt.bindNull(7); + } + if (entry.rule_id) |rid| { + try stmt.bindInt(8, rid); + } else { + try stmt.bindNull(8); + } + try stmt.bindInt(9, @intFromEnum(entry.reply_type)); + try stmt.bindInt(10, @intFromEnum(entry.protocol)); + try stmt.bindUint64(11, entry.response_time_us); + + if (entry.upstream) |u| { + try stmt.bindText(12, u); + } else { + try stmt.bindNull(12); + } + + if (entry.reason) |r| { + try stmt.bindText(13, r); + } else { + try stmt.bindNull(13); + } + + try stmt.bindInt(14, if (entry.dnssec_validated) 1 else 0); + + _ = try stmt.step(); + } + + fn getOrCreateDomainId(self: *QueryLogger, domain: []const u8) !i64 { + // Try to get existing + var select_stmt = try self.db.prepare("SELECT id FROM domains WHERE domain = ?1"); + defer select_stmt.finalize(); + + try select_stmt.bindText(1, domain); + + if (try select_stmt.step()) { + return select_stmt.getInt(0); + } + + // Insert new + var insert_stmt = try self.db.prepare("INSERT INTO domains (domain) VALUES (?1)"); + defer insert_stmt.finalize(); + + try insert_stmt.bindText(1, domain); + _ = try insert_stmt.step(); + + return self.db.lastInsertRowId(); + } + + fn getOrCreateClientId(self: *QueryLogger, client_ip: []const u8) !i64 { + // Try to get existing + var select_stmt = try self.db.prepare("SELECT id FROM clients WHERE ip = ?1"); + defer select_stmt.finalize(); + + try select_stmt.bindText(1, client_ip); + + if (try select_stmt.step()) { + const client_id = select_stmt.getInt(0); + + // Update last_seen + var update_stmt = try self.db.prepare( + \\UPDATE clients SET last_seen = strftime('%s', 'now') WHERE id = ?1 + ); + defer update_stmt.finalize(); + + try update_stmt.bindInt(1, client_id); + _ = try update_stmt.step(); + + return client_id; + } + + // Insert new + var insert_stmt = try self.db.prepare("INSERT INTO clients (ip) VALUES (?1)"); + defer insert_stmt.finalize(); + + try insert_stmt.bindText(1, client_ip); + _ = try insert_stmt.step(); + + return self.db.lastInsertRowId(); + } + + /// Convert to handler-compatible Logger interface + pub fn toHandlerLogger(self: *QueryLogger) handler.Logger { + return handler.Logger{ + .context = self, + .logFn = logWrapper, + }; + } + + fn logWrapper(ctx: *anyopaque, entry: handler.QueryLogEntry) void { + const self: *QueryLogger = @ptrCast(@alignCast(ctx)); + self.log(entry); + } +}; + +/// Clean up old query logs based on retention policy +/// retention_seconds: duration in seconds (use config.logging.retentionSeconds()) +pub fn cleanupOldLogs(db: *Database, retention_seconds: i64) !void { + const cutoff = std.time.timestamp() - retention_seconds; + + var stmt = try db.prepare("DELETE FROM query_log WHERE timestamp < ?1"); + defer stmt.finalize(); + + try stmt.bindInt(1, cutoff); + _ = try stmt.step(); + + const deleted = db.changes(); + if (deleted > 0) { + std.log.info("Cleaned up {} old query log entries", .{deleted}); + } +} + +test "QueryLogger basic test" { + const testing = std.testing; + const allocator = testing.allocator; + const types = @import("../dns/types.zig"); + + var db = try Database.open(":memory:", allocator); + defer db.close(); + + try schema.createSchema(&db); + + var logger = QueryLogger.init(&db, allocator); + defer logger.deinit(); + + // Log a query with new fields + logger.log(.{ + .timestamp = std.time.timestamp(), + .domain = "example.com", + .client_ip = "192.168.1.1", + .qtype = types.QType.A, + .status = .forwarded, + .list_id = null, + .rule_id = null, + .reply_type = .ip, + .protocol = .udp, + .response_time_us = 1234, + .upstream = "cloudflare", + .reason = null, + .dnssec_validated = true, + }); + + // Flush + try logger.flush(); + + // Verify it was logged + var stmt = try db.prepare("SELECT COUNT(*) FROM query_log"); + defer stmt.finalize(); + + _ = try stmt.step(); + const count = stmt.getInt(0); + try testing.expectEqual(@as(i64, 1), count); + + // Verify status was logged correctly + var status_stmt = try db.prepare("SELECT status FROM query_log LIMIT 1"); + defer status_stmt.finalize(); + _ = try status_stmt.step(); + const status = status_stmt.getInt(0); + try testing.expectEqual(@as(i64, 1), status); // forwarded = 1 +} diff --git a/src/storage/schema.zig b/src/storage/schema.zig new file mode 100644 index 0000000..05a8edb --- /dev/null +++ b/src/storage/schema.zig @@ -0,0 +1,539 @@ +const std = @import("std"); +const Database = @import("db.zig").Database; + +pub const SCHEMA_VERSION: i64 = 1; + +// ============================================================================= +// Query Status - tracks exactly why a query was allowed/denied +// ============================================================================= + +pub const QueryStatus = enum(u8) { + unknown = 0, // Status not determined + forwarded = 1, // Sent to upstream, got response + cached = 2, // Answered from cache + denied_denylist = 3, // Blocked by denylist (gravity) + denied_rule = 4, // Blocked by user rule + denied_upstream = 5, // Upstream returned block (e.g., safe search rewrite) + allowed_rule = 6, // Explicitly allowed by rule, overriding denylist + upstream_error = 7, // Upstream failed, SERVFAIL, etc. + + pub fn toString(self: QueryStatus) []const u8 { + return switch (self) { + .unknown => "unknown", + .forwarded => "forwarded", + .cached => "cached", + .denied_denylist => "denied_denylist", + .denied_rule => "denied_rule", + .denied_upstream => "denied_upstream", + .allowed_rule => "allowed_rule", + .upstream_error => "upstream_error", + }; + } + + pub fn fromInt(value: i64) QueryStatus { + return switch (value) { + 1 => .forwarded, + 2 => .cached, + 3 => .denied_denylist, + 4 => .denied_rule, + 5 => .denied_upstream, + 6 => .allowed_rule, + 7 => .upstream_error, + else => .unknown, + }; + } + + pub fn isDenied(self: QueryStatus) bool { + return switch (self) { + .denied_denylist, .denied_rule, .denied_upstream => true, + else => false, + }; + } +}; + +// ============================================================================= +// Reply Type - tracks DNS response type +// ============================================================================= + +pub const ReplyType = enum(u8) { + unknown = 0, + nodata = 1, // Empty response (no records) + nxdomain = 2, // Domain doesn't exist + cname = 3, // CNAME response + ip = 4, // A/AAAA response + servfail = 5, // Server failure + refused = 6, // Query refused + + pub fn toString(self: ReplyType) []const u8 { + return switch (self) { + .unknown => "unknown", + .nodata => "nodata", + .nxdomain => "nxdomain", + .cname => "cname", + .ip => "ip", + .servfail => "servfail", + .refused => "refused", + }; + } + + pub fn fromInt(value: i64) ReplyType { + return switch (value) { + 1 => .nodata, + 2 => .nxdomain, + 3 => .cname, + 4 => .ip, + 5 => .servfail, + 6 => .refused, + else => .unknown, + }; + } +}; + +// ============================================================================= +// Source Type - denylist vs allowlist +// ============================================================================= + +pub const SourceType = enum(u8) { + denylist = 0, + allowlist = 1, + + pub fn toString(self: SourceType) []const u8 { + return switch (self) { + .denylist => "denylist", + .allowlist => "allowlist", + }; + } + + pub fn fromInt(value: i64) SourceType { + return switch (value) { + 1 => .allowlist, + else => .denylist, + }; + } +}; + +// ============================================================================= +// Client Protocol - UDP vs TCP +// ============================================================================= + +pub const ClientProtocol = enum(u8) { + udp = 0, + tcp = 1, + + pub fn toString(self: ClientProtocol) []const u8 { + return switch (self) { + .udp => "udp", + .tcp => "tcp", + }; + } + + pub fn fromInt(value: i64) ClientProtocol { + return switch (value) { + 1 => .tcp, + else => .udp, + }; + } +}; + +/// Denylist fetch status (follows Pi-hole conventions) +pub const DenylistStatus = enum(u8) { + pending = 0, // Never fetched / unknown + updated = 1, // Successfully fetched new data + unchanged = 2, // Content unchanged (same hash) + cached = 3, // Fetch failed, using cached data + failed = 4, // Fetch failed, no data available + + pub fn toString(self: DenylistStatus) []const u8 { + return switch (self) { + .pending => "pending", + .updated => "updated", + .unchanged => "unchanged", + .cached => "cached", + .failed => "failed", + }; + } + + pub fn fromInt(value: i64) DenylistStatus { + return switch (value) { + 1 => .updated, + 2 => .unchanged, + 3 => .cached, + 4 => .failed, + else => .pending, + }; + } +}; + +/// Create the initial database schema +pub fn createSchema(db: *Database) !void { + try db.execMulti( + \\-- Schema version tracking + \\CREATE TABLE IF NOT EXISTS schema_version ( + \\ version INTEGER PRIMARY KEY + \\); + \\ + \\-- Groups for client classification + \\CREATE TABLE IF NOT EXISTS groups ( + \\ id INTEGER PRIMARY KEY, + \\ name TEXT NOT NULL UNIQUE, + \\ description TEXT, + \\ created_at INTEGER DEFAULT (strftime('%s', 'now')), + \\ date_modified INTEGER + \\); + \\ + \\-- Insert default group + \\INSERT OR IGNORE INTO groups (id, name, description) VALUES (0, 'default', 'Default group for all clients'); + \\ + \\-- Clients (devices) + \\CREATE TABLE IF NOT EXISTS clients ( + \\ id INTEGER PRIMARY KEY, + \\ ip TEXT NOT NULL UNIQUE, + \\ name TEXT, + \\ group_id INTEGER DEFAULT 0 REFERENCES groups(id), + \\ first_seen INTEGER DEFAULT (strftime('%s', 'now')), + \\ last_seen INTEGER DEFAULT (strftime('%s', 'now')), + \\ date_modified INTEGER + \\); + \\ + \\CREATE INDEX IF NOT EXISTS idx_clients_ip ON clients(ip); + \\CREATE INDEX IF NOT EXISTS idx_clients_group ON clients(group_id); + \\ + \\-- Denylist sources (type: 0=denylist, 1=allowlist) + \\CREATE TABLE IF NOT EXISTS denylist_sources ( + \\ id INTEGER PRIMARY KEY, + \\ url TEXT NOT NULL UNIQUE, + \\ comment TEXT, + \\ category TEXT DEFAULT 'ads', + \\ enabled INTEGER DEFAULT 1, + \\ status INTEGER DEFAULT 0, + \\ last_updated INTEGER, + \\ domain_count INTEGER DEFAULT 0, + \\ invalid_domains INTEGER DEFAULT 0, + \\ content_hash TEXT, + \\ update_interval INTEGER DEFAULT 86400, + \\ type INTEGER DEFAULT 0, + \\ created_at INTEGER DEFAULT (strftime('%s', 'now')), + \\ date_modified INTEGER + \\); + \\ + \\-- Domains from denylists + \\CREATE TABLE IF NOT EXISTS denylist_domains ( + \\ id INTEGER PRIMARY KEY, + \\ domain TEXT NOT NULL, + \\ source_id INTEGER NOT NULL REFERENCES denylist_sources(id) ON DELETE CASCADE, + \\ UNIQUE(domain, source_id) + \\); + \\ + \\CREATE INDEX IF NOT EXISTS idx_denylist_domains_domain ON denylist_domains(domain); + \\CREATE INDEX IF NOT EXISTS idx_denylist_domains_source ON denylist_domains(source_id); + \\ + \\-- Which groups use which denylist sources + \\CREATE TABLE IF NOT EXISTS group_sources ( + \\ group_id INTEGER NOT NULL REFERENCES groups(id) ON DELETE CASCADE, + \\ source_id INTEGER NOT NULL REFERENCES denylist_sources(id) ON DELETE CASCADE, + \\ PRIMARY KEY (group_id, source_id) + \\); + \\ + \\-- Custom allow/deny rules (per-group) + \\CREATE TABLE IF NOT EXISTS rules ( + \\ id INTEGER PRIMARY KEY, + \\ domain TEXT NOT NULL, + \\ group_id INTEGER REFERENCES groups(id) ON DELETE CASCADE, + \\ action TEXT NOT NULL CHECK(action IN ('allow', 'deny')), + \\ comment TEXT, + \\ created_at INTEGER DEFAULT (strftime('%s', 'now')), + \\ date_modified INTEGER, + \\ UNIQUE(domain, group_id) + \\); + \\ + \\CREATE INDEX IF NOT EXISTS idx_rules_domain ON rules(domain); + \\CREATE INDEX IF NOT EXISTS idx_rules_group ON rules(group_id); + \\ + \\-- String interning for query log (domains) + \\CREATE TABLE IF NOT EXISTS domains ( + \\ id INTEGER PRIMARY KEY, + \\ domain TEXT NOT NULL UNIQUE + \\); + \\ + \\CREATE INDEX IF NOT EXISTS idx_domains_domain ON domains(domain); + \\ + \\-- Query log (status: see QueryStatus enum, protocol: 0=udp, 1=tcp) + \\CREATE TABLE IF NOT EXISTS query_log ( + \\ id INTEGER PRIMARY KEY, + \\ timestamp INTEGER NOT NULL, + \\ domain_id INTEGER NOT NULL REFERENCES domains(id), + \\ client_id INTEGER NOT NULL REFERENCES clients(id), + \\ qtype INTEGER, + \\ denied INTEGER NOT NULL DEFAULT 0, + \\ status INTEGER DEFAULT 0, + \\ list_id INTEGER REFERENCES denylist_sources(id), + \\ rule_id INTEGER REFERENCES rules(id), + \\ reply_type INTEGER DEFAULT 0, + \\ protocol INTEGER DEFAULT 0, + \\ response_time_us INTEGER, + \\ upstream TEXT, + \\ reason TEXT, + \\ dnssec_validated INTEGER NOT NULL DEFAULT 0 + \\); + \\ + \\CREATE INDEX IF NOT EXISTS idx_query_log_timestamp ON query_log(timestamp); + \\CREATE INDEX IF NOT EXISTS idx_query_log_client ON query_log(client_id); + \\CREATE INDEX IF NOT EXISTS idx_query_log_domain ON query_log(domain_id); + \\CREATE INDEX IF NOT EXISTS idx_query_log_denied ON query_log(denied); + \\CREATE INDEX IF NOT EXISTS idx_query_log_status ON query_log(status); + \\ + \\-- Client-group many-to-many junction table + \\CREATE TABLE IF NOT EXISTS client_groups ( + \\ client_id INTEGER NOT NULL REFERENCES clients(id) ON DELETE CASCADE, + \\ group_id INTEGER NOT NULL REFERENCES groups(id) ON DELETE CASCADE, + \\ PRIMARY KEY (client_id, group_id) + \\); + \\ + \\-- Allowlist domains (for allowlist sources) + \\CREATE TABLE IF NOT EXISTS allowlist_domains ( + \\ id INTEGER PRIMARY KEY, + \\ domain TEXT NOT NULL, + \\ source_id INTEGER NOT NULL REFERENCES denylist_sources(id) ON DELETE CASCADE, + \\ UNIQUE(domain, source_id) + \\); + \\CREATE INDEX IF NOT EXISTS idx_allowlist_domains_domain ON allowlist_domains(domain); + \\CREATE INDEX IF NOT EXISTS idx_allowlist_domains_source ON allowlist_domains(source_id); + \\ + \\-- Settings + \\CREATE TABLE IF NOT EXISTS settings ( + \\ key TEXT PRIMARY KEY, + \\ value TEXT NOT NULL + \\); + \\ + \\-- Default settings + \\INSERT OR IGNORE INTO settings (key, value) VALUES ('blocking_response', 'zero'); + \\INSERT OR IGNORE INTO settings (key, value) VALUES ('safe_search_enabled', 'true'); + \\INSERT OR IGNORE INTO settings (key, value) VALUES ('log_retention', '30 days'); + ); + + try db.exec("INSERT OR REPLACE INTO schema_version (version) VALUES (1)"); +} + +/// Get current schema version +pub fn getSchemaVersion(db: *Database) !i64 { + // Check if schema_version table exists + var stmt = db.prepare( + \\SELECT name FROM sqlite_master + \\WHERE type='table' AND name='schema_version' + ) catch return 0; + defer stmt.finalize(); + + if (!(try stmt.step())) { + return 0; + } + + // Get version + var version_stmt = try db.prepare("SELECT version FROM schema_version LIMIT 1"); + defer version_stmt.finalize(); + + if (try version_stmt.step()) { + return version_stmt.getInt(0); + } + + return 0; +} + +pub fn migrate(db: *Database) !void { + const current_version = try getSchemaVersion(db); + if (current_version == 0) { + try createSchema(db); + } +} + +/// Denylist categories +pub const DenylistCategory = enum { + ads, + malware, + adult, + telemetry, + gambling, + social, + + pub fn toString(self: DenylistCategory) []const u8 { + return switch (self) { + .ads => "ads", + .malware => "malware", + .adult => "adult", + .telemetry => "telemetry", + .gambling => "gambling", + .social => "social", + }; + } +}; + +/// Pre-configured denylist source +const DefaultDenylist = struct { + url: []const u8, + comment: ?[]const u8, + category: DenylistCategory, + enabled_by_default: bool, +}; + +/// Default denylist sources from PLAN.md +pub const DEFAULT_DENYLISTS = [_]DefaultDenylist{ + // Ads & Trackers + .{ + .url = "https://raw.githubusercontent.com/StevenBlack/hosts/master/hosts", + .comment = "StevenBlack Unified", + .category = .ads, + .enabled_by_default = true, + }, + .{ + .url = "https://big.oisd.nl/domainswild", + .comment = "OISD Big (large list)", + .category = .ads, + .enabled_by_default = false, + }, + .{ + .url = "https://adguardteam.github.io/AdGuardSDNSFilter/Filters/filter.txt", + .comment = "AdGuard DNS Filter", + .category = .ads, + .enabled_by_default = true, + }, + .{ + .url = "https://pgl.yoyo.org/adservers/serverlist.php?hostformat=hosts&showintro=0", + .comment = "Peter Lowe's Ad Servers", + .category = .ads, + .enabled_by_default = true, + }, + + // Malware & Phishing + .{ + .url = "https://urlhaus.abuse.ch/downloads/hostfile/", + .comment = "URLhaus Malware", + .category = .malware, + .enabled_by_default = true, + }, + .{ + .url = "https://threatfox.abuse.ch/downloads/hostfile/", + .comment = "ThreatFox IOCs", + .category = .malware, + .enabled_by_default = true, + }, + .{ + .url = "https://malware-filter.gitlab.io/malware-filter/phishing-filter-hosts.txt", + .comment = "Phishing Filter", + .category = .malware, + .enabled_by_default = true, + }, + + // Adult Content + .{ + .url = "https://nsfw.oisd.nl/domainswild", + .comment = "OISD NSFW", + .category = .adult, + .enabled_by_default = false, + }, + .{ + .url = "https://raw.githubusercontent.com/StevenBlack/hosts/master/alternates/porn/hosts", + .comment = null, + .category = .adult, + .enabled_by_default = false, + }, + + // Native Telemetry + .{ + .url = "https://raw.githubusercontent.com/nickspaargaren/no-google/master/categories/proxies", + .comment = null, + .category = .telemetry, + .enabled_by_default = false, + }, + .{ + .url = "https://raw.githubusercontent.com/nickspaargaren/no-google/master/categories/analytics", + .comment = null, + .category = .telemetry, + .enabled_by_default = false, + }, + + // Gambling + .{ + .url = "https://raw.githubusercontent.com/StevenBlack/hosts/master/alternates/gambling/hosts", + .comment = null, + .category = .gambling, + .enabled_by_default = false, + }, + + // Social Media + .{ + .url = "https://raw.githubusercontent.com/StevenBlack/hosts/master/alternates/social/hosts", + .comment = null, + .category = .social, + .enabled_by_default = false, + }, +}; + +/// Insert default denylist sources +pub fn insertDefaultDenylists(db: *Database) !void { + for (DEFAULT_DENYLISTS) |list| { + var stmt = try db.prepare( + \\INSERT OR IGNORE INTO denylist_sources (url, comment, category, enabled) + \\VALUES (?1, ?2, ?3, ?4) + ); + defer stmt.finalize(); + + try stmt.bindText(1, list.url); + if (list.comment) |comment| { + try stmt.bindText(2, comment); + } else { + try stmt.bindNull(2); + } + try stmt.bindText(3, list.category.toString()); + try stmt.bindInt(4, if (list.enabled_by_default) 1 else 0); + _ = try stmt.step(); + } + + // Associate enabled denylists with default group + try db.exec( + \\INSERT OR IGNORE INTO group_sources (group_id, source_id) + \\SELECT 0, id FROM denylist_sources WHERE enabled = 1 + ); +} + +test "schema creation" { + const testing = std.testing; + const allocator = testing.allocator; + + var db = try Database.open(":memory:", allocator); + defer db.close(); + + try createSchema(&db); + + // Verify tables exist + var stmt = try db.prepare( + \\SELECT name FROM sqlite_master + \\WHERE type='table' ORDER BY name + ); + defer stmt.finalize(); + + var tables = std.ArrayListUnmanaged([]const u8){}; + defer tables.deinit(allocator); + + while (try stmt.step()) { + if (stmt.getText(0)) |name| { + try tables.append(allocator, try allocator.dupe(u8, name)); + } + } + defer { + for (tables.items) |t| allocator.free(t); + } + + try testing.expect(tables.items.len > 0); +} + +test "migration from fresh db" { + const testing = std.testing; + const allocator = testing.allocator; + + var db = try Database.open(":memory:", allocator); + defer db.close(); + + try migrate(&db); + + const version = try getSchemaVersion(&db); + try testing.expectEqual(SCHEMA_VERSION, version); +} diff --git a/src/upstream/connection_pool.zig b/src/upstream/connection_pool.zig new file mode 100644 index 0000000..23e5ca1 --- /dev/null +++ b/src/upstream/connection_pool.zig @@ -0,0 +1,340 @@ +const std = @import("std"); +const net = std.net; +const posix = std.posix; +const tls = std.crypto.tls; +const Certificate = std.crypto.Certificate; +const Allocator = std.mem.Allocator; + +/// Persistent TLS connection for DNS-over-TLS +/// Heap-allocates all TLS state to allow connection reuse across queries. +/// This avoids expensive TLS handshakes for repeated queries to the same server. +pub const PersistentDotConnection = struct { + allocator: Allocator, + host: []const u8, + port: u16, + state: ?*ConnectionState, + timeout_ms: u32, + + const ConnectionState = struct { + stream: net.Stream, + // Heap-allocated buffers for the TLS client + stream_read_buffer: *[tls.Client.min_buffer_len]u8, + stream_write_buffer: *[tls.Client.min_buffer_len]u8, + tls_read_buffer: *[tls.Client.min_buffer_len]u8, + tls_write_buffer: *[tls.Client.min_buffer_len]u8, + // The stream reader/writer use pointers to buffers above + stream_reader: net.Stream.Reader, + stream_writer: net.Stream.Writer, + // TLS client + tls_client: tls.Client, + // CA bundle must stay alive for the connection duration + ca_bundle: Certificate.Bundle, + + fn deinit(self: *ConnectionState, allocator: Allocator) void { + self.ca_bundle.deinit(allocator); + self.stream.close(); + allocator.destroy(self.stream_read_buffer); + allocator.destroy(self.stream_write_buffer); + allocator.destroy(self.tls_read_buffer); + allocator.destroy(self.tls_write_buffer); + } + }; + + pub const Error = error{ + InvalidHost, + ConnectionFailed, + TlsHandshakeFailed, + CertificateLoadFailed, + SendFailed, + ReceiveFailed, + InvalidResponse, + Timeout, + OutOfMemory, + }; + + pub fn init(host: []const u8, port: u16, allocator: Allocator) PersistentDotConnection { + return .{ + .allocator = allocator, + .host = host, + .port = port, + .state = null, + .timeout_ms = 5000, + }; + } + + pub fn deinit(self: *PersistentDotConnection) void { + self.close(); + } + + pub fn setTimeout(self: *PersistentDotConnection, timeout_ms: u32) void { + self.timeout_ms = timeout_ms; + } + + /// Close the persistent connection + pub fn close(self: *PersistentDotConnection) void { + if (self.state) |state| { + state.deinit(self.allocator); + self.allocator.destroy(state); + self.state = null; + } + } + + /// Ensure we have a valid connection, creating one if needed + fn ensureConnected(self: *PersistentDotConnection) Error!*ConnectionState { + if (self.state) |state| { + // Check if connection is still alive using poll + // POLLIN without data but no POLLHUP/POLLERR means connection is good + var fds = [1]posix.pollfd{.{ + .fd = state.stream.handle, + .events = posix.POLL.IN, + .revents = 0, + }}; + + // Non-blocking poll (timeout=0) + const poll_result = posix.poll(&fds, 0) catch { + self.close(); + return self.connect(); + }; + + // If poll returns 0, no events - connection is idle and good + if (poll_result == 0) { + return state; + } + + // Check for error conditions + if (fds[0].revents & (posix.POLL.ERR | posix.POLL.HUP | posix.POLL.NVAL) != 0) { + // Connection is dead + self.close(); + return self.connect(); + } + + // Connection has data or is ready - that's fine, return it + return state; + } + + // Create new connection + return self.connect(); + } + + /// Establish a new TLS connection + fn connect(self: *PersistentDotConnection) Error!*ConnectionState { + // Allocate buffers on heap + const stream_read_buffer = self.allocator.create([tls.Client.min_buffer_len]u8) catch return error.OutOfMemory; + errdefer self.allocator.destroy(stream_read_buffer); + + const stream_write_buffer = self.allocator.create([tls.Client.min_buffer_len]u8) catch return error.OutOfMemory; + errdefer self.allocator.destroy(stream_write_buffer); + + const tls_read_buffer = self.allocator.create([tls.Client.min_buffer_len]u8) catch return error.OutOfMemory; + errdefer self.allocator.destroy(tls_read_buffer); + + const tls_write_buffer = self.allocator.create([tls.Client.min_buffer_len]u8) catch return error.OutOfMemory; + errdefer self.allocator.destroy(tls_write_buffer); + + // Load system CA certificates + var ca_bundle: Certificate.Bundle = .{}; + ca_bundle.rescan(self.allocator) catch |err| { + std.log.err("Failed to load system CA certificates: {}", .{err}); + return error.CertificateLoadFailed; + }; + errdefer ca_bundle.deinit(self.allocator); + + // Connect to server via TCP + const stream = net.tcpConnectToHost(self.allocator, self.host, self.port) catch |err| { + std.log.warn("DoT connection to {s}:{d} failed: {}", .{ self.host, self.port, err }); + return error.ConnectionFailed; + }; + errdefer stream.close(); + + // Apply socket timeouts + const timeout = posix.timeval{ + .sec = @intCast(self.timeout_ms / 1000), + .usec = @intCast((self.timeout_ms % 1000) * 1000), + }; + posix.setsockopt(stream.handle, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch {}; + posix.setsockopt(stream.handle, posix.SOL.SOCKET, posix.SO.SNDTIMEO, std.mem.asBytes(&timeout)) catch {}; + + // Allocate connection state + const state = self.allocator.create(ConnectionState) catch return error.OutOfMemory; + errdefer self.allocator.destroy(state); + + // Initialize stream reader/writer with heap buffers + state.stream = stream; + state.stream_read_buffer = stream_read_buffer; + state.stream_write_buffer = stream_write_buffer; + state.tls_read_buffer = tls_read_buffer; + state.tls_write_buffer = tls_write_buffer; + state.ca_bundle = ca_bundle; + + state.stream_reader = net.Stream.Reader.init(stream, stream_read_buffer); + state.stream_writer = net.Stream.Writer.init(stream, stream_write_buffer); + + // Initialize TLS client with handshake + state.tls_client = tls.Client.init(state.stream_reader.interface(), &state.stream_writer.interface, .{ + .host = .{ .explicit = self.host }, + .ca = .{ .bundle = state.ca_bundle }, + .write_buffer = tls_write_buffer, + .read_buffer = tls_read_buffer, + }) catch |err| { + std.log.warn("DoT TLS handshake with {s} failed: {}", .{ self.host, err }); + // Clean up manually since errdefer won't handle partial state + ca_bundle.deinit(self.allocator); + stream.close(); + self.allocator.destroy(stream_read_buffer); + self.allocator.destroy(stream_write_buffer); + self.allocator.destroy(tls_read_buffer); + self.allocator.destroy(tls_write_buffer); + self.allocator.destroy(state); + return error.TlsHandshakeFailed; + }; + + self.state = state; + return state; + } + + /// Send a DNS query and receive the response + /// Reuses existing connection if available, creates new one if needed + pub fn query(self: *PersistentDotConnection, dns_packet: []const u8) Error![]const u8 { + const state = self.ensureConnected() catch |err| { + return err; + }; + + // DNS-over-TLS uses 2-byte length prefix (big-endian) + var len_buf: [2]u8 = undefined; + std.mem.writeInt(u16, &len_buf, @intCast(dns_packet.len), .big); + + // Send length prefix + DNS packet via TLS + state.tls_client.writer.writeAll(&len_buf) catch { + self.close(); // Connection failed, close it + return error.SendFailed; + }; + state.tls_client.writer.writeAll(dns_packet) catch { + self.close(); + return error.SendFailed; + }; + state.tls_client.writer.flush() catch { + self.close(); + return error.SendFailed; + }; + + // Read response length via TLS + var resp_len_buf: [2]u8 = undefined; + state.tls_client.reader.readSliceAll(&resp_len_buf) catch { + self.close(); + return error.ReceiveFailed; + }; + + const resp_len = std.mem.readInt(u16, &resp_len_buf, .big); + + // Validate response length + if (resp_len < 12 or resp_len > 65535) { + self.close(); + return error.InvalidResponse; + } + + // Read response + const response = self.allocator.alloc(u8, resp_len) catch return error.OutOfMemory; + errdefer self.allocator.free(response); + + state.tls_client.reader.readSliceAll(response) catch { + self.allocator.free(response); + self.close(); + return error.ReceiveFailed; + }; + + return response; + } +}; + +/// Connection pool for managing persistent connections to upstream DNS servers. +/// Maintains one connection per host:port combination. +pub const DotConnectionPool = struct { + connections: std.StringHashMapUnmanaged(*PersistentDotConnection), + mutex: std.Thread.Mutex, + allocator: Allocator, + + pub fn init(allocator: Allocator) DotConnectionPool { + return .{ + .connections = .{}, + .mutex = .{}, + .allocator = allocator, + }; + } + + pub fn deinit(self: *DotConnectionPool) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + var iter = self.connections.iterator(); + while (iter.next()) |entry| { + entry.value_ptr.*.deinit(); + self.allocator.destroy(entry.value_ptr.*); + self.allocator.free(entry.key_ptr.*); + } + self.connections.deinit(self.allocator); + } + + /// Get or create a persistent connection for the given host:port + pub fn getConnection(self: *DotConnectionPool, host: []const u8, port: u16) !*PersistentDotConnection { + self.mutex.lock(); + defer self.mutex.unlock(); + + var key_buf: [280]u8 = undefined; + const key = std.fmt.bufPrint(&key_buf, "{s}:{d}", .{ host, port }) catch return error.InvalidHost; + + if (self.connections.get(key)) |conn| { + return conn; + } + + // Create new persistent connection + const conn = try self.allocator.create(PersistentDotConnection); + conn.* = PersistentDotConnection.init(host, port, self.allocator); + + const key_copy = try self.allocator.dupe(u8, key); + errdefer self.allocator.free(key_copy); + + try self.connections.put(self.allocator, key_copy, conn); + return conn; + } + + /// Close a specific connection (e.g., after repeated failures) + pub fn closeConnection(self: *DotConnectionPool, host: []const u8, port: u16) void { + self.mutex.lock(); + defer self.mutex.unlock(); + + var key_buf: [280]u8 = undefined; + const key = std.fmt.bufPrint(&key_buf, "{s}:{d}", .{ host, port }) catch return; + + if (self.connections.get(key)) |conn| { + conn.close(); + } + } +}; + +test "PersistentDotConnection basic" { + const testing = std.testing; + const allocator = testing.allocator; + + var conn = PersistentDotConnection.init("cloudflare-dns.com", 853, allocator); + defer conn.deinit(); + + // Just verify initialization works - actual connection test would need network + try testing.expectEqual(@as(u16, 853), conn.port); + try testing.expectEqualStrings("cloudflare-dns.com", conn.host); +} + +test "DotConnectionPool basic" { + const testing = std.testing; + const allocator = testing.allocator; + + var pool = DotConnectionPool.init(allocator); + defer pool.deinit(); + + // Get a connection (won't actually connect without network) + const conn = try pool.getConnection("cloudflare-dns.com", 853); + try testing.expectEqual(@as(u16, 853), conn.port); + + // Getting same host:port should return same connection + const conn2 = try pool.getConnection("cloudflare-dns.com", 853); + try testing.expectEqual(conn, conn2); +} diff --git a/src/upstream/doh.zig b/src/upstream/doh.zig new file mode 100644 index 0000000..6fc4190 --- /dev/null +++ b/src/upstream/doh.zig @@ -0,0 +1,234 @@ +const std = @import("std"); +const http = std.http; +const Allocator = std.mem.Allocator; +const Writer = std.Io.Writer; + +/// Number of parallel connection slots per URL +const NUM_SLOTS = 8; + +/// DNS-over-HTTPS connection pool using slot-based sharding. +/// Maintains NUM_SLOTS independent HTTP clients per URL, with round-robin assignment. +/// Each slot has its own mutex, allowing up to NUM_SLOTS parallel queries. +pub const DohConnectionPool = struct { + /// Per-URL pools + url_pools: std.StringHashMapUnmanaged(*UrlPool), + pools_mutex: std.Thread.Mutex, + allocator: Allocator, + + /// Pool for a single URL with NUM_SLOTS parallel slots + const UrlPool = struct { + slots: [NUM_SLOTS]Slot, + next: std.atomic.Value(usize), + url: []const u8, + + const Slot = struct { + mutex: std.Thread.Mutex, + client: ?http.Client, + }; + + fn init(url: []const u8) UrlPool { + var pool = UrlPool{ + .slots = undefined, + .next = std.atomic.Value(usize).init(0), + .url = url, + }; + for (&pool.slots) |*slot| { + slot.* = Slot{ + .mutex = .{}, + .client = null, + }; + } + return pool; + } + + fn deinit(self: *UrlPool) void { + for (&self.slots) |*slot| { + if (slot.client) |*client| { + client.deinit(); + } + } + } + + fn query(self: *UrlPool, dns_packet: []const u8, allocator: Allocator) ![]const u8 { + // Round-robin slot selection + const idx = self.next.fetchAdd(1, .monotonic) % NUM_SLOTS; + const slot = &self.slots[idx]; + + slot.mutex.lock(); + defer slot.mutex.unlock(); + + // Lazily create client for this slot + if (slot.client == null) { + slot.client = http.Client{ .allocator = allocator }; + } + + var response_writer = Writer.Allocating.init(allocator); + errdefer response_writer.deinit(); + + const result = slot.client.?.fetch(.{ + .location = .{ .url = self.url }, + .method = .POST, + .payload = dns_packet, + .extra_headers = &[_]http.Header{ + .{ .name = "Content-Type", .value = "application/dns-message" }, + .{ .name = "Accept", .value = "application/dns-message" }, + }, + .response_writer = &response_writer.writer, + }) catch |err| { + std.log.warn("DoH request failed for {s}: {}", .{ self.url, err }); + // Reset client on failure - might be stale connection + if (slot.client) |*client| { + client.deinit(); + slot.client = null; + } + return error.ConnectionFailed; + }; + + if (result.status != .ok) { + std.log.warn("DoH HTTP {d} from {s}", .{ @intFromEnum(result.status), self.url }); + return error.HttpError; + } + + const data = response_writer.written(); + if (data.len < 12) { + std.log.warn("DoH response too short ({d} bytes) from {s}", .{ data.len, self.url }); + return error.InvalidResponse; + } + + return response_writer.toOwnedSlice() catch return error.OutOfMemory; + } + }; + + pub const Error = error{ + ConnectionFailed, + InvalidResponse, + OutOfMemory, + HttpError, + }; + + pub fn init(allocator: Allocator) DohConnectionPool { + return .{ + .url_pools = .{}, + .pools_mutex = .{}, + .allocator = allocator, + }; + } + + pub fn deinit(self: *DohConnectionPool) void { + self.pools_mutex.lock(); + defer self.pools_mutex.unlock(); + + var iter = self.url_pools.iterator(); + while (iter.next()) |entry| { + entry.value_ptr.*.deinit(); + self.allocator.destroy(entry.value_ptr.*); + self.allocator.free(entry.key_ptr.*); + } + self.url_pools.deinit(self.allocator); + } + + /// Query a DoH server, using pooled connections for parallelism + pub fn query(self: *DohConnectionPool, url: []const u8, dns_packet: []const u8) Error![]const u8 { + const pool = self.getOrCreatePool(url) catch return error.OutOfMemory; + return pool.query(dns_packet, self.allocator); + } + + fn getOrCreatePool(self: *DohConnectionPool, url: []const u8) !*UrlPool { + self.pools_mutex.lock(); + defer self.pools_mutex.unlock(); + + if (self.url_pools.get(url)) |pool| { + return pool; + } + + // Create new pool for this URL + const url_copy = try self.allocator.dupe(u8, url); + errdefer self.allocator.free(url_copy); + + const pool = try self.allocator.create(UrlPool); + pool.* = UrlPool.init(url_copy); + + try self.url_pools.put(self.allocator, url_copy, pool); + return pool; + } +}; + +// Legacy single-query client (no pooling) - kept for compatibility +pub const DohClient = struct { + url: []const u8, + allocator: Allocator, + + pub const Error = error{ + InvalidUrl, + ConnectionFailed, + InvalidResponse, + OutOfMemory, + HttpError, + }; + + pub fn init(url: []const u8, allocator: Allocator) Error!DohClient { + if (!std.mem.startsWith(u8, url, "https://")) { + return error.InvalidUrl; + } + return DohClient{ + .url = url, + .allocator = allocator, + }; + } + + pub fn query(self: *DohClient, dns_packet: []const u8) Error![]const u8 { + var client = http.Client{ .allocator = self.allocator }; + defer client.deinit(); + + var response_writer = Writer.Allocating.init(self.allocator); + errdefer response_writer.deinit(); + + const result = client.fetch(.{ + .location = .{ .url = self.url }, + .method = .POST, + .payload = dns_packet, + .extra_headers = &[_]http.Header{ + .{ .name = "Content-Type", .value = "application/dns-message" }, + .{ .name = "Accept", .value = "application/dns-message" }, + }, + .response_writer = &response_writer.writer, + }) catch |err| { + std.log.warn("DoH request failed for {s}: {}", .{ self.url, err }); + return error.ConnectionFailed; + }; + + if (result.status != .ok) { + std.log.warn("DoH HTTP {d} from {s}", .{ @intFromEnum(result.status), self.url }); + return error.HttpError; + } + + const data = response_writer.written(); + if (data.len < 12) { + std.log.warn("DoH response too short ({d} bytes) from {s}", .{ data.len, self.url }); + return error.InvalidResponse; + } + + return response_writer.toOwnedSlice() catch return error.OutOfMemory; + } + + pub fn setTimeout(self: *DohClient, timeout_secs: u32) void { + _ = self; + _ = timeout_secs; + } +}; + +test "DohConnectionPool basic" { + const testing = std.testing; + const allocator = testing.allocator; + + var pool = DohConnectionPool.init(allocator); + defer pool.deinit(); + + // Just test pool creation - actual queries need network + const url_pool = try pool.getOrCreatePool("https://cloudflare-dns.com/dns-query"); + try testing.expectEqualStrings("https://cloudflare-dns.com/dns-query", url_pool.url); + + // Same URL should return same pool + const url_pool2 = try pool.getOrCreatePool("https://cloudflare-dns.com/dns-query"); + try testing.expectEqual(url_pool, url_pool2); +} diff --git a/src/upstream/dot.zig b/src/upstream/dot.zig new file mode 100644 index 0000000..e6a5691 --- /dev/null +++ b/src/upstream/dot.zig @@ -0,0 +1,189 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const net = std.net; +const posix = std.posix; +const tls = std.crypto.tls; +const Io = std.Io; +const Certificate = std.crypto.Certificate; + +/// DNS-over-TLS client +/// Sends DNS queries over TLS on port 853 using 2-byte length prefix format (RFC 7858) +pub const DotClient = struct { + host: []const u8, + port: u16, + allocator: Allocator, + timeout_ms: u32, + + pub const DEFAULT_PORT: u16 = 853; + + pub const Error = error{ + InvalidHost, + ConnectionFailed, + TlsHandshakeFailed, + CertificateLoadFailed, + SendFailed, + ReceiveFailed, + InvalidResponse, + Timeout, + OutOfMemory, + }; + + /// Initialize a DoT client + /// host: hostname like "cloudflare-dns.com" or "1.1.1.1" + /// port: typically 853 + pub fn init(host: []const u8, port: u16, allocator: Allocator) DotClient { + return DotClient{ + .host = host, + .port = port, + .allocator = allocator, + .timeout_ms = 5000, // 5 second default timeout + }; + } + + /// Parse a tls:// URL and create a DoT client + /// e.g., "tls://cloudflare-dns.com" or "tls://1.1.1.1:853" + pub fn fromUrl(url: []const u8, allocator: Allocator) Error!DotClient { + // Strip tls:// prefix + const prefix = "tls://"; + if (!std.mem.startsWith(u8, url, prefix)) { + return error.InvalidHost; + } + + 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 error.InvalidHost; + return DotClient.init(host, port, allocator); + } + + return DotClient.init(host_port, DEFAULT_PORT, allocator); + } + + /// Query the DoT server with a DNS packet + /// Returns the DNS response packet (caller owns the memory) + pub fn query(self: *DotClient, dns_packet: []const u8) Error![]const u8 { + // Load system CA certificates for TLS verification + var ca_bundle: Certificate.Bundle = .{}; + ca_bundle.rescan(self.allocator) catch |err| { + std.log.err("Failed to load system CA certificates: {}", .{err}); + return error.CertificateLoadFailed; + }; + defer ca_bundle.deinit(self.allocator); + + // Connect to server via TCP + const stream = net.tcpConnectToHost(self.allocator, self.host, self.port) catch |err| { + std.log.warn("DoT connection to {s}:{d} failed: {}", .{ self.host, self.port, err }); + return error.ConnectionFailed; + }; + defer stream.close(); + + // Apply socket timeouts + const timeout = posix.timeval{ + .sec = @intCast(self.timeout_ms / 1000), + .usec = @intCast((self.timeout_ms % 1000) * 1000), + }; + posix.setsockopt(stream.handle, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch |err| { + std.log.warn("DoT: failed to set receive timeout: {}", .{err}); + }; + posix.setsockopt(stream.handle, posix.SOL.SOCKET, posix.SO.SNDTIMEO, std.mem.asBytes(&timeout)) catch |err| { + std.log.warn("DoT: failed to set send timeout: {}", .{err}); + }; + + // Buffers for the underlying TCP stream reader/writer + var stream_read_buffer: [tls.Client.min_buffer_len]u8 = undefined; + var stream_write_buffer: [tls.Client.min_buffer_len]u8 = undefined; + + // Create stream reader/writer + var stream_reader = net.Stream.Reader.init(stream, &stream_read_buffer); + var stream_writer = net.Stream.Writer.init(stream, &stream_write_buffer); + + // Buffers for TLS layer + var tls_read_buffer: [tls.Client.min_buffer_len]u8 = undefined; + var tls_write_buffer: [tls.Client.min_buffer_len]u8 = undefined; + + // Initialize TLS client with handshake using system CA bundle + var tls_client = tls.Client.init(stream_reader.interface(), &stream_writer.interface, .{ + .host = .{ .explicit = self.host }, + .ca = .{ .bundle = ca_bundle }, + .write_buffer = &tls_write_buffer, + .read_buffer = &tls_read_buffer, + }) catch |err| { + std.log.warn("DoT TLS handshake with {s} failed: {}", .{ self.host, err }); + return error.TlsHandshakeFailed; + }; + + // DNS-over-TLS uses 2-byte length prefix (big-endian) + var len_buf: [2]u8 = undefined; + std.mem.writeInt(u16, &len_buf, @intCast(dns_packet.len), .big); + + // Send length prefix + DNS packet via TLS + tls_client.writer.writeAll(&len_buf) catch return error.SendFailed; + tls_client.writer.writeAll(dns_packet) catch return error.SendFailed; + tls_client.writer.flush() catch return error.SendFailed; + + // Read response length via TLS + var resp_len_buf: [2]u8 = undefined; + tls_client.reader.readSliceAll(&resp_len_buf) catch return error.ReceiveFailed; + + const resp_len = std.mem.readInt(u16, &resp_len_buf, .big); + + // Validate response length + if (resp_len < 12 or resp_len > 65535) return error.InvalidResponse; + + // Read response + const response = self.allocator.alloc(u8, resp_len) catch return error.OutOfMemory; + errdefer self.allocator.free(response); + + tls_client.reader.readSliceAll(response) catch { + self.allocator.free(response); + return error.ReceiveFailed; + }; + + return response; + } + + /// Set the timeout in milliseconds + pub fn setTimeout(self: *DotClient, timeout_ms: u32) void { + self.timeout_ms = timeout_ms; + } +}; + +/// Query a DoT server (convenience function) +pub fn queryDot(host: []const u8, port: u16, dns_packet: []const u8, allocator: Allocator) DotClient.Error![]const u8 { + var client = DotClient.init(host, port, allocator); + return client.query(dns_packet); +} + +/// Query a DoT server from a URL (e.g., "tls://cloudflare-dns.com") +pub fn queryDotUrl(url: []const u8, dns_packet: []const u8, allocator: Allocator) DotClient.Error![]const u8 { + var client = try DotClient.fromUrl(url, allocator); + return client.query(dns_packet); +} + +test "DotClient URL parsing" { + const testing = std.testing; + const allocator = testing.allocator; + + // Valid TLS URLs + const client1 = try DotClient.fromUrl("tls://cloudflare-dns.com", allocator); + try testing.expectEqualStrings("cloudflare-dns.com", client1.host); + try testing.expectEqual(DotClient.DEFAULT_PORT, client1.port); + + const client2 = try DotClient.fromUrl("tls://1.1.1.1:853", allocator); + try testing.expectEqualStrings("1.1.1.1", client2.host); + try testing.expectEqual(@as(u16, 853), client2.port); + + const client3 = try DotClient.fromUrl("tls://dns.google:8853", allocator); + try testing.expectEqualStrings("dns.google", client3.host); + try testing.expectEqual(@as(u16, 8853), client3.port); + + // Invalid URLs should fail + const invalid_result = DotClient.fromUrl("https://example.com", allocator); + try testing.expectError(error.InvalidHost, invalid_result); + + const invalid_result2 = DotClient.fromUrl("not-a-url", allocator); + try testing.expectError(error.InvalidHost, invalid_result2); +} diff --git a/src/upstream/pool.zig b/src/upstream/pool.zig new file mode 100644 index 0000000..f7a024d --- /dev/null +++ b/src/upstream/pool.zig @@ -0,0 +1,413 @@ +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); +} diff --git a/src/util.zig b/src/util.zig new file mode 100644 index 0000000..240bd59 --- /dev/null +++ b/src/util.zig @@ -0,0 +1,24 @@ +const std = @import("std"); + +/// Convert a string to lowercase using a pre-allocated buffer. +/// Returns null if the buffer is too small. +pub fn toLower(s: []const u8, buf: []u8) ?[]const u8 { + if (s.len > buf.len) return null; + for (s, 0..) |c, i| { + buf[i] = std.ascii.toLower(c); + } + return buf[0..s.len]; +} + +test "toLower" { + const testing = std.testing; + var buf: [64]u8 = undefined; + + try testing.expectEqualStrings("hello", toLower("HELLO", &buf).?); + try testing.expectEqualStrings("hello", toLower("hello", &buf).?); + try testing.expectEqualStrings("hello123", toLower("HeLLo123", &buf).?); + + // Buffer too small + var small_buf: [2]u8 = undefined; + try testing.expect(toLower("hello", &small_buf) == null); +} diff --git a/src/web/api/clients.zig b/src/web/api/clients.zig new file mode 100644 index 0000000..e764559 --- /dev/null +++ b/src/web/api/clients.zig @@ -0,0 +1,85 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +pub fn list(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare( + \\SELECT c.id, c.ip, c.name, c.group_id, g.name, c.first_seen, c.last_seen + \\FROM clients c + \\LEFT JOIN groups g ON c.group_id = g.id + \\ORDER BY c.last_seen DESC + ) catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginArray() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + w.beginObject() catch break; + w.writeIntField("id", stmt.getInt(0)) catch break; + w.writeStringField("ip", stmt.getText(1) orelse "unknown") catch break; + w.writeOptionalStringField("name", stmt.getText(2)) catch break; + w.writeIntField("group_id", stmt.getInt(3)) catch break; + w.writeStringField("group_name", stmt.getText(4) orelse "default") catch break; + w.writeIntField("first_seen", stmt.getInt(5)) catch break; + w.writeIntField("last_seen", stmt.getInt(6)) catch break; + w.endObject() catch break; + } + + w.endArray() catch return; + + response.sendJson(handle, w.items()); +} + +const UpdateRequest = struct { + name: ?[]const u8 = null, + group_id: ?i64 = null, +}; + +pub fn update(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, id: i64, body: []const u8) void { + const parsed = std.json.parseFromSlice(UpdateRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + if (parsed.value.name) |name| { + var stmt = db.prepare("UPDATE clients SET name = ? WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + stmt.bindText(1, name) catch return; + stmt.bindInt(2, id) catch return; + _ = stmt.step() catch |err| { + std.log.warn("Failed to update client name: {}", .{err}); + }; + } + + if (parsed.value.group_id) |group_id| { + var stmt = db.prepare("UPDATE clients SET group_id = ? WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + stmt.bindInt(1, group_id) catch return; + stmt.bindInt(2, id) catch return; + _ = stmt.step() catch |err| { + std.log.warn("Failed to update client group: {}", .{err}); + }; + } + + response.sendOk(handle); +} diff --git a/src/web/api/denylists.zig b/src/web/api/denylists.zig new file mode 100644 index 0000000..bc5b207 --- /dev/null +++ b/src/web/api/denylists.zig @@ -0,0 +1,229 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const schema = @import("../../storage/schema.zig"); +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +/// Maximum lengths for validation +const MAX_COMMENT_LENGTH = 256; +const MAX_URL_LENGTH = 2048; + +/// Validate a denylist URL +fn isValidUrl(url: []const u8) bool { + if (url.len == 0 or url.len > MAX_URL_LENGTH) return false; + + // Must start with http:// or https:// + if (!std.mem.startsWith(u8, url, "http://") and !std.mem.startsWith(u8, url, "https://")) { + return false; + } + + // Check for invalid characters (basic URL validation) + for (url) |c| { + // URL-safe characters + if (!std.ascii.isAlphanumeric(c) and + c != ':' and c != '/' and c != '.' and c != '-' and c != '_' and + c != '?' and c != '&' and c != '=' and c != '%' and c != '+' and + c != '#' and c != '@' and c != '!' and c != '$' and c != '\'' and + c != '(' and c != ')' and c != '*' and c != ',' and c != ';') + { + return false; + } + } + + return true; +} + +pub fn list(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare( + \\SELECT id, url, comment, enabled, status, domain_count FROM denylist_sources ORDER BY id + ) catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginArray() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + const status_int = stmt.getInt(4); + const status = schema.DenylistStatus.fromInt(status_int); + + w.beginObject() catch break; + w.writeIntField("id", stmt.getInt(0)) catch break; + w.writeStringField("url", stmt.getText(1) orelse "") catch break; + if (stmt.getText(2)) |comment| { + w.writeStringField("comment", comment) catch break; + } + w.writeBoolField("enabled", stmt.getInt(3) == 1) catch break; + w.writeIntField("status", status_int) catch break; + w.writeStringField("status_text", status.toString()) catch break; + w.writeIntField("domain_count", stmt.getInt(5)) catch break; + w.endObject() catch break; + } + + w.endArray() catch return; + + response.sendJson(handle, w.items()); +} + +const AddRequest = struct { + url: []const u8, + comment: ?[]const u8 = null, +}; + +pub fn add(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, body: []const u8) void { + const parsed = std.json.parseFromSlice(AddRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + // Validate denylist URL + if (!isValidUrl(parsed.value.url)) { + response.sendBadRequest(handle, "Invalid URL (must be http:// or https://, max 2048 chars)"); + return; + } + + // Validate comment length if provided + if (parsed.value.comment) |comment| { + if (comment.len > MAX_COMMENT_LENGTH) { + response.sendBadRequest(handle, "Comment too long (max 256 chars)"); + return; + } + } + + var stmt = db.prepare("INSERT INTO denylist_sources (url, comment, enabled, domain_count) VALUES (?, ?, 1, 0)") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindText(1, parsed.value.url) catch return; + if (parsed.value.comment) |comment| { + stmt.bindText(2, comment) catch return; + } else { + stmt.bindNull(2) catch return; + } + + _ = stmt.step() catch { + response.sendServerError(handle, "Insert failed (URL may already exist)"); + return; + }; + + const id = db.lastInsertRowId(); + + // Link to default group (0) so it's active immediately + linkToDefaultGroup: { + var group_stmt = db.prepare("INSERT INTO group_sources (group_id, source_id) VALUES (0, ?)") catch { + std.log.warn("Failed to link denylist {d} to default group", .{id}); + break :linkToDefaultGroup; + }; + defer group_stmt.finalize(); + group_stmt.bindInt(1, id) catch break :linkToDefaultGroup; + _ = group_stmt.step() catch |err| { + std.log.warn("Failed to insert group_sources for denylist {d}: {}", .{ id, err }); + }; + } + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.endObject() catch return; + + // Send 201 Created + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 201 Created\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + response.CORS_API, + w.items().len, + }) catch return; + + _ = posix.write(handle, header) catch |err| { + std.log.debug("Failed to write response header: {}", .{err}); + return; + }; + _ = posix.write(handle, w.items()) catch |err| { + std.log.debug("Failed to write response body: {}", .{err}); + }; +} + +pub fn delete(db: *Database, handle: posix.socket_t, id: i64) void { + // Delete associated domains first + var del_domains = db.prepare("DELETE FROM denylist_domains WHERE source_id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer del_domains.finalize(); + del_domains.bindInt(1, id) catch return; + _ = del_domains.step() catch |err| { + std.log.warn("Failed to delete denylist domains for source {d}: {}", .{ id, err }); + // Continue - we still want to try deleting the source + }; + + // Delete from group_sources + var del_groups = db.prepare("DELETE FROM group_sources WHERE source_id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer del_groups.finalize(); + del_groups.bindInt(1, id) catch return; + _ = del_groups.step() catch |err| { + std.log.warn("Failed to delete group_sources for source {d}: {}", .{ id, err }); + // Continue - we still want to try deleting the source + }; + + // Delete the source + var stmt = db.prepare("DELETE FROM denylist_sources WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + stmt.bindInt(1, id) catch return; + _ = stmt.step() catch { + response.sendServerError(handle, "Delete failed"); + return; + }; + + response.sendOk(handle); +} + +const ToggleRequest = struct { + enabled: bool, +}; + +pub fn toggle(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, id: i64, body: []const u8) void { + const parsed = std.json.parseFromSlice(ToggleRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + var stmt = db.prepare("UPDATE denylist_sources SET enabled = ? WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindInt(1, if (parsed.value.enabled) 1 else 0) catch return; + stmt.bindInt(2, id) catch return; + + _ = stmt.step() catch { + response.sendServerError(handle, "Update failed"); + return; + }; + + response.sendOk(handle); +} diff --git a/src/web/api/groups.zig b/src/web/api/groups.zig new file mode 100644 index 0000000..5859f36 --- /dev/null +++ b/src/web/api/groups.zig @@ -0,0 +1,183 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +/// Maximum length for group/client names +const MAX_NAME_LENGTH = 64; + +/// Validate a name (group or client name) +fn isValidName(name: []const u8) bool { + // Empty or too long + if (name.len == 0 or name.len > MAX_NAME_LENGTH) return false; + + // Must start with alphanumeric + if (!std.ascii.isAlphanumeric(name[0])) return false; + + // Can only contain alphanumeric, spaces, hyphens, and underscores + for (name) |c| { + if (!std.ascii.isAlphanumeric(c) and c != ' ' and c != '-' and c != '_') return false; + } + + return true; +} + +pub fn list(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare("SELECT id, name, description FROM groups ORDER BY id") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginArray() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + w.beginObject() catch break; + w.writeIntField("id", stmt.getInt(0)) catch break; + w.writeStringField("name", stmt.getText(1) orelse "unknown") catch break; + w.writeOptionalStringField("description", stmt.getText(2)) catch break; + w.endObject() catch break; + } + + w.endArray() catch return; + + response.sendJson(handle, w.items()); +} + +const GroupRequest = struct { + name: []const u8, + description: ?[]const u8 = null, +}; + +pub fn add(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, body: []const u8) void { + const parsed = std.json.parseFromSlice(GroupRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + // Validate group name + if (!isValidName(parsed.value.name)) { + response.sendBadRequest(handle, "Invalid group name (max 64 chars, alphanumeric start)"); + return; + } + + var stmt = db.prepare("INSERT INTO groups (name, description) VALUES (?, ?)") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindText(1, parsed.value.name) catch return; + if (parsed.value.description) |desc| { + stmt.bindText(2, desc) catch return; + } else { + stmt.bindNull(2) catch return; + } + + _ = stmt.step() catch { + response.sendServerError(handle, "Insert failed"); + return; + }; + + const id = db.lastInsertRowId(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.endObject() catch return; + + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 201 Created\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + response.CORS_API, + w.items().len, + }) catch return; + + _ = posix.write(handle, header) catch |err| { + std.log.debug("Failed to write response header: {}", .{err}); + return; + }; + _ = posix.write(handle, w.items()) catch |err| { + std.log.debug("Failed to write response body: {}", .{err}); + }; +} + +pub fn update(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, id: i64, body: []const u8) void { + if (id == 0) { + response.sendBadRequest(handle, "Cannot modify default group"); + return; + } + + const parsed = std.json.parseFromSlice(GroupRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + var stmt = db.prepare("UPDATE groups SET name = ?, description = ? WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindText(1, parsed.value.name) catch return; + if (parsed.value.description) |desc| { + stmt.bindText(2, desc) catch return; + } else { + stmt.bindNull(2) catch return; + } + stmt.bindInt(3, id) catch return; + + _ = stmt.step() catch { + response.sendServerError(handle, "Update failed"); + return; + }; + + response.sendOk(handle); +} + +pub fn delete(db: *Database, handle: posix.socket_t, id: i64) void { + if (id == 0) { + response.sendBadRequest(handle, "Cannot delete default group"); + return; + } + + // Move clients in this group to default group first + var move_stmt = db.prepare("UPDATE clients SET group_id = 0 WHERE group_id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer move_stmt.finalize(); + move_stmt.bindInt(1, id) catch return; + _ = move_stmt.step() catch |err| { + std.log.warn("Failed to move clients to default group: {}", .{err}); + }; + + var stmt = db.prepare("DELETE FROM groups WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindInt(1, id) catch return; + _ = stmt.step() catch { + response.sendServerError(handle, "Delete failed"); + return; + }; + + response.sendOk(handle); +} diff --git a/src/web/api/queries.zig b/src/web/api/queries.zig new file mode 100644 index 0000000..23222f8 --- /dev/null +++ b/src/web/api/queries.zig @@ -0,0 +1,47 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +/// Maximum number of queries to return in the API response +const DEFAULT_QUERY_LIMIT = 100; + +pub fn sendQueries(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + const query = std.fmt.comptimePrint( + \\SELECT ql.timestamp, d.domain, COALESCE(c.name, c.ip), ql.qtype, ql.denied, ql.dnssec_validated + \\FROM query_log ql + \\JOIN domains d ON ql.domain_id = d.id + \\JOIN clients c ON ql.client_id = c.id + \\ORDER BY ql.timestamp DESC + \\LIMIT {d} + , .{DEFAULT_QUERY_LIMIT}); + var stmt = db.prepare(query) catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginArray() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + w.beginObject() catch break; + w.writeIntField("timestamp", stmt.getInt(0)) catch break; + w.writeStringField("domain", stmt.getText(1) orelse "unknown") catch break; + w.writeStringField("client", stmt.getText(2) orelse "unknown") catch break; + w.writeIntField("qtype", stmt.getInt(3)) catch break; + w.writeBoolField("denied", stmt.getInt(4) == 1) catch break; + w.writeBoolField("dnssec", stmt.getInt(5) == 1) catch break; + w.endObject() catch break; + } + + w.endArray() catch return; + + response.sendJson(handle, w.items()); +} diff --git a/src/web/api/rules.zig b/src/web/api/rules.zig new file mode 100644 index 0000000..56c733e --- /dev/null +++ b/src/web/api/rules.zig @@ -0,0 +1,180 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +/// Maximum length for a DNS domain name (RFC 1035) +const MAX_DOMAIN_LENGTH = 253; + +/// Validate a domain name according to RFC 1035 rules +fn isValidDomain(domain: []const u8) bool { + // Empty or too long + if (domain.len == 0 or domain.len > MAX_DOMAIN_LENGTH) return false; + + // Wildcard domains like *.example.com are allowed + var check_domain = domain; + if (std.mem.startsWith(u8, domain, "*.")) { + check_domain = domain[2..]; + if (check_domain.len == 0) return false; + } + + // Check each label + var labels = std.mem.splitScalar(u8, check_domain, '.'); + var label_count: usize = 0; + while (labels.next()) |label| { + // Label cannot be empty (except trailing dot creates empty last label) + if (label.len == 0) { + // Allow trailing dot + if (labels.peek() == null) continue; + return false; + } + // Label cannot exceed 63 characters + if (label.len > 63) return false; + + // Label must start with alphanumeric + if (!std.ascii.isAlphanumeric(label[0])) return false; + + // Label must end with alphanumeric + if (!std.ascii.isAlphanumeric(label[label.len - 1])) return false; + + // Label can only contain alphanumeric and hyphens + for (label) |c| { + if (!std.ascii.isAlphanumeric(c) and c != '-') return false; + } + + label_count += 1; + } + + // Must have at least one label + return label_count > 0; +} + +pub fn list(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare( + \\SELECT r.id, r.domain, r.group_id, g.name, r.action, r.comment + \\FROM rules r + \\LEFT JOIN groups g ON r.group_id = g.id + \\ORDER BY r.id DESC + ) catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginArray() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + w.beginObject() catch break; + w.writeIntField("id", stmt.getInt(0)) catch break; + w.writeStringField("domain", stmt.getText(1) orelse "") catch break; + w.writeIntField("group_id", stmt.getInt(2)) catch break; + w.writeStringField("group_name", stmt.getText(3) orelse "all") catch break; + w.writeStringField("action", stmt.getText(4) orelse "deny") catch break; + w.writeOptionalStringField("comment", stmt.getText(5)) catch break; + w.endObject() catch break; + } + + w.endArray() catch return; + + response.sendJson(handle, w.items()); +} + +const AddRequest = struct { + domain: []const u8, + action: []const u8, + group_id: ?i64 = null, + comment: ?[]const u8 = null, +}; + +pub fn add(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, body: []const u8) void { + const parsed = std.json.parseFromSlice(AddRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + // Validate action + if (!std.mem.eql(u8, parsed.value.action, "allow") and !std.mem.eql(u8, parsed.value.action, "deny")) { + response.sendBadRequest(handle, "Action must be 'allow' or 'deny'"); + return; + } + + // Validate domain + if (!isValidDomain(parsed.value.domain)) { + response.sendBadRequest(handle, "Invalid domain name"); + return; + } + + var stmt = db.prepare("INSERT INTO rules (domain, group_id, action, comment) VALUES (?, ?, ?, ?)") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindText(1, parsed.value.domain) catch return; + if (parsed.value.group_id) |gid| { + stmt.bindInt(2, gid) catch return; + } else { + stmt.bindNull(2) catch return; + } + stmt.bindText(3, parsed.value.action) catch return; + if (parsed.value.comment) |c| { + stmt.bindText(4, c) catch return; + } else { + stmt.bindNull(4) catch return; + } + + _ = stmt.step() catch { + response.sendServerError(handle, "Insert failed (duplicate?)"); + return; + }; + + const id = db.lastInsertRowId(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.endObject() catch return; + + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 201 Created\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + response.CORS_API, + w.items().len, + }) catch return; + + _ = posix.write(handle, header) catch |err| { + std.log.debug("Failed to write response header: {}", .{err}); + return; + }; + _ = posix.write(handle, w.items()) catch |err| { + std.log.debug("Failed to write response body: {}", .{err}); + }; +} + +pub fn delete(db: *Database, handle: posix.socket_t, id: i64) void { + var stmt = db.prepare("DELETE FROM rules WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindInt(1, id) catch return; + _ = stmt.step() catch { + response.sendServerError(handle, "Delete failed"); + return; + }; + + response.sendOk(handle); +} diff --git a/src/web/api/settings.zig b/src/web/api/settings.zig new file mode 100644 index 0000000..f80e0ee --- /dev/null +++ b/src/web/api/settings.zig @@ -0,0 +1,71 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +pub fn list(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare("SELECT key, value FROM settings") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginObject() catch return; + + while (true) { + const has_row = stmt.step() catch break; + if (!has_row) break; + + const key = stmt.getText(0) orelse continue; + const value = stmt.getText(1) orelse ""; + + w.writeStringField(key, value) catch break; + } + + w.endObject() catch return; + + response.sendJson(handle, w.items()); +} + +const UpdateRequest = struct { + blocking_response: ?[]const u8 = null, + safe_search_enabled: ?[]const u8 = null, + log_retention: ?[]const u8 = null, +}; + +pub fn update(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator, body: []const u8) void { + const parsed = std.json.parseFromSlice(UpdateRequest, allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + if (parsed.value.blocking_response) |val| { + updateSetting(db, "blocking_response", val); + } + if (parsed.value.safe_search_enabled) |val| { + updateSetting(db, "safe_search_enabled", val); + } + if (parsed.value.log_retention) |val| { + updateSetting(db, "log_retention", val); + } + + response.sendOk(handle); +} + +fn updateSetting(db: *Database, key: []const u8, value: []const u8) void { + var stmt = db.prepare("INSERT OR REPLACE INTO settings (key, value) VALUES (?, ?)") catch return; + defer stmt.finalize(); + + stmt.bindText(1, key) catch return; + stmt.bindText(2, value) catch return; + _ = stmt.step() catch |err| { + std.log.warn("Failed to update setting '{s}': {}", .{ key, err }); + }; +} diff --git a/src/web/api/stats.zig b/src/web/api/stats.zig new file mode 100644 index 0000000..d1c0009 --- /dev/null +++ b/src/web/api/stats.zig @@ -0,0 +1,49 @@ +const std = @import("std"); +const posix = std.posix; +const Database = @import("../../storage/db.zig").Database; +const json = @import("../json.zig"); +const response = @import("../response.zig"); + +pub fn sendStats(db: *Database, handle: posix.socket_t, allocator: std.mem.Allocator) void { + var stmt = db.prepare("SELECT COUNT(*) FROM query_log") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + const total_queries: i64 = blk: { + const has_row = stmt.step() catch break :blk 0; + break :blk if (has_row) stmt.getInt(0) else 0; + }; + + var denied_stmt = db.prepare("SELECT COUNT(*) FROM query_log WHERE denied = 1") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer denied_stmt.finalize(); + + const denied_queries: i64 = blk: { + const has_row = denied_stmt.step() catch break :blk 0; + break :blk if (has_row) denied_stmt.getInt(0) else 0; + }; + + const deny_percentage: f64 = if (total_queries > 0) + @as(f64, @floatFromInt(denied_queries)) / @as(f64, @floatFromInt(total_queries)) * 100.0 + else + 0.0; + + var w = json.JsonWriter.init(allocator); + defer w.deinit(); + + w.beginObject() catch return; + w.writeIntField("total_queries", total_queries) catch return; + w.writeIntField("denied_queries", denied_queries) catch return; + // Write percentage as string to control formatting + w.writeKey("deny_percentage") catch return; + var pct_buf: [16]u8 = undefined; + const pct_str = std.fmt.bufPrint(&pct_buf, "{d:.1}", .{deny_percentage}) catch return; + w.writeRaw(pct_str) catch return; + w.endObject() catch return; + + response.sendJson(handle, w.items()); +} diff --git a/src/web/auth.zig b/src/web/auth.zig new file mode 100644 index 0000000..f0b2f15 --- /dev/null +++ b/src/web/auth.zig @@ -0,0 +1,141 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; +const argon2 = std.crypto.pwhash.argon2; + +/// HTTP Basic Authentication handler using Argon2id +/// +/// Uses OWASP-recommended Argon2id parameters for secure password hashing. +/// PHC format string includes salt and params, so verification is self-contained. +pub const BasicAuth = struct { + username: []const u8, + password_hash: [128]u8, // PHC format string + hash_len: usize, + enabled: bool, + allocator: Allocator, + + /// Create auth with plaintext password (hashes it with Argon2id) + pub fn init(username: []const u8, password: []const u8, allocator: Allocator) BasicAuth { + var hash_buf: [128]u8 = undefined; + const hash = argon2.strHash(password, .{ + .allocator = allocator, + .params = argon2.Params.owasp_2id, + }, &hash_buf) catch { + return BasicAuth.disabled(allocator); + }; + + return BasicAuth{ + .username = username, + .password_hash = hash_buf, + .hash_len = hash.len, + .enabled = true, + .allocator = allocator, + }; + } + + /// Create disabled auth (no authentication required) + pub fn disabled(allocator: Allocator) BasicAuth { + return BasicAuth{ + .username = "", + .password_hash = undefined, + .hash_len = 0, + .enabled = false, + .allocator = allocator, + }; + } + + /// Check if authentication is required + pub fn isEnabled(self: BasicAuth) bool { + return self.enabled; + } + + /// Validate credentials from Authorization header + /// Header format: "Basic base64(username:password)" + pub fn validate(self: BasicAuth, auth_header: ?[]const u8) bool { + if (!self.enabled) return true; + + const header = auth_header orelse return false; + + // Must start with "Basic " + if (!std.mem.startsWith(u8, header, "Basic ")) return false; + + const encoded = header[6..]; + + // Decode base64 + var decode_buf: [256]u8 = undefined; + const decoded_len = std.base64.standard.Decoder.calcSizeForSlice(encoded) catch return false; + if (decoded_len > decode_buf.len) return false; + + std.base64.standard.Decoder.decode(&decode_buf, encoded) catch return false; + const decoded = decode_buf[0..decoded_len]; + + // Find colon separator + const colon_idx = std.mem.indexOfScalar(u8, decoded, ':') orelse return false; + + const username = decoded[0..colon_idx]; + const password = decoded[colon_idx + 1 ..]; + + // Check username + if (!std.mem.eql(u8, username, self.username)) return false; + + // Verify password with Argon2id + argon2.strVerify(self.password_hash[0..self.hash_len], password, .{ + .allocator = self.allocator, + }) catch return false; + + return true; + } + + /// Generate WWW-Authenticate header for 401 response + pub fn getWwwAuthenticateHeader() []const u8 { + return "WWW-Authenticate: Basic realm=\"nxdns\"\r\n"; + } +}; + +/// Extract Authorization header from HTTP request +pub fn extractAuthHeader(request: []const u8) ?[]const u8 { + var lines = std.mem.splitSequence(u8, request, "\r\n"); + + while (lines.next()) |line| { + if (std.ascii.startsWithIgnoreCase(line, "authorization:")) { + const value = std.mem.trim(u8, line[14..], " "); + return value; + } + } + + return null; +} + +test "BasicAuth validation" { + const testing = std.testing; + + var auth = BasicAuth.init("admin", "secret123", testing.allocator); + + // Valid credentials (base64("admin:secret123") = "YWRtaW46c2VjcmV0MTIz") + try testing.expect(auth.validate("Basic YWRtaW46c2VjcmV0MTIz")); + + // Invalid credentials + try testing.expect(!auth.validate("Basic YWRtaW46d3Jvbmc=")); // admin:wrong + try testing.expect(!auth.validate("Basic dXNlcjpzZWNyZXQxMjM=")); // user:secret123 + try testing.expect(!auth.validate(null)); + try testing.expect(!auth.validate("Bearer token")); +} + +test "BasicAuth disabled" { + const testing = std.testing; + + const auth = BasicAuth.disabled(testing.allocator); + + // Should always return true when disabled + try testing.expect(auth.validate(null)); + try testing.expect(auth.validate("Basic anything")); +} + +test "extractAuthHeader" { + const testing = std.testing; + + const request = "GET / HTTP/1.1\r\nHost: localhost\r\nAuthorization: Basic YWRtaW46c2VjcmV0\r\n\r\n"; + const header = extractAuthHeader(request); + + try testing.expect(header != null); + try testing.expectEqualStrings("Basic YWRtaW46c2VjcmV0", header.?); +} diff --git a/src/web/index.html b/src/web/index.html new file mode 100644 index 0000000..b122789 --- /dev/null +++ b/src/web/index.html @@ -0,0 +1,537 @@ + + + + + + nxdns - DNS Sinkhole + + + +
+
+

nxdns

+ +
+ +
+
+
Total Queries
+
-
+
+
+
Denied Queries
+
-
+
+
+
Deny Rate
+
-
+
+
+ + +
+
Recent Queries
+ + + + + + + + + + + + + + +
TimestampDomainClientTypeStatusDNSSEC
Loading...
+
+ + +
+
+
Denylists
+ +
+
+ + + +
+ + + + + + + + + + + + +
URLEntriesStatusActions
Loading...
+
+ + +
+
Custom Rules
+
+ + + + + +
+ + + + + + + + + + + + + +
DomainActionGroupCommentActions
Loading...
+
+ + +
+
Client Groups
+
+ + + +
+ + + + + + + + + + + + +
IDNameDescriptionActions
Loading...
+
+ + +
+
Discovered Clients
+ + + + + + + + + + + + + + +
IP AddressNameGroupFirst SeenLast SeenActions
Loading...
+
+ + +
+
Settings
+
+ + +
+
+ + +
+
+ + +
+
+
+ + + + diff --git a/src/web/json.zig b/src/web/json.zig new file mode 100644 index 0000000..fb386f0 --- /dev/null +++ b/src/web/json.zig @@ -0,0 +1,179 @@ +const std = @import("std"); +const Allocator = std.mem.Allocator; + +/// JSON writer that builds JSON incrementally into a buffer +pub const JsonWriter = struct { + buffer: std.ArrayListUnmanaged(u8), + allocator: Allocator, + first_in_container: bool, + + pub fn init(allocator: Allocator) JsonWriter { + return .{ + .buffer = .{}, + .allocator = allocator, + .first_in_container = true, + }; + } + + pub fn deinit(self: *JsonWriter) void { + self.buffer.deinit(self.allocator); + } + + pub fn toOwnedSlice(self: *JsonWriter) []u8 { + return self.buffer.toOwnedSlice(self.allocator) catch &[_]u8{}; + } + + pub fn items(self: *JsonWriter) []const u8 { + return self.buffer.items; + } + + pub fn beginArray(self: *JsonWriter) !void { + try self.buffer.append(self.allocator, '['); + self.first_in_container = true; + } + + pub fn endArray(self: *JsonWriter) !void { + try self.buffer.append(self.allocator, ']'); + self.first_in_container = false; + } + + pub fn beginObject(self: *JsonWriter) !void { + if (!self.first_in_container) { + try self.buffer.append(self.allocator, ','); + } + try self.buffer.append(self.allocator, '{'); + self.first_in_container = true; + } + + pub fn endObject(self: *JsonWriter) !void { + try self.buffer.append(self.allocator, '}'); + self.first_in_container = false; + } + + pub fn writeKey(self: *JsonWriter, key: []const u8) !void { + if (!self.first_in_container) { + try self.buffer.append(self.allocator, ','); + } + try self.buffer.append(self.allocator, '"'); + try self.buffer.appendSlice(self.allocator, key); + try self.buffer.appendSlice(self.allocator, "\":"); + self.first_in_container = true; + } + + pub fn writeString(self: *JsonWriter, value: []const u8) !void { + try self.buffer.append(self.allocator, '"'); + // Escape special characters + for (value) |c| { + switch (c) { + '"' => try self.buffer.appendSlice(self.allocator, "\\\""), + '\\' => try self.buffer.appendSlice(self.allocator, "\\\\"), + '\n' => try self.buffer.appendSlice(self.allocator, "\\n"), + '\r' => try self.buffer.appendSlice(self.allocator, "\\r"), + '\t' => try self.buffer.appendSlice(self.allocator, "\\t"), + else => try self.buffer.append(self.allocator, c), + } + } + try self.buffer.append(self.allocator, '"'); + self.first_in_container = false; + } + + pub fn writeInt(self: *JsonWriter, value: i64) !void { + var buf: [32]u8 = undefined; + const str = std.fmt.bufPrint(&buf, "{d}", .{value}) catch return error.OutOfMemory; + try self.buffer.appendSlice(self.allocator, str); + self.first_in_container = false; + } + + pub fn writeBool(self: *JsonWriter, value: bool) !void { + try self.buffer.appendSlice(self.allocator, if (value) "true" else "false"); + self.first_in_container = false; + } + + pub fn writeNull(self: *JsonWriter) !void { + try self.buffer.appendSlice(self.allocator, "null"); + self.first_in_container = false; + } + + /// Write raw text (no escaping, for pre-formatted numbers) + pub fn writeRaw(self: *JsonWriter, value: []const u8) !void { + try self.buffer.appendSlice(self.allocator, value); + self.first_in_container = false; + } + + /// Write a key-value pair with string value + pub fn writeStringField(self: *JsonWriter, key: []const u8, value: []const u8) !void { + try self.writeKey(key); + try self.writeString(value); + } + + /// Write a key-value pair with optional string value + pub fn writeOptionalStringField(self: *JsonWriter, key: []const u8, value: ?[]const u8) !void { + try self.writeKey(key); + if (value) |v| { + try self.writeString(v); + } else { + try self.writeNull(); + } + } + + /// Write a key-value pair with integer value + pub fn writeIntField(self: *JsonWriter, key: []const u8, value: i64) !void { + try self.writeKey(key); + try self.writeInt(value); + } + + /// Write a key-value pair with boolean value + pub fn writeBoolField(self: *JsonWriter, key: []const u8, value: bool) !void { + try self.writeKey(key); + try self.writeBool(value); + } +}; + +test "JsonWriter basic" { + const testing = std.testing; + const allocator = testing.allocator; + + var w = JsonWriter.init(allocator); + defer w.deinit(); + + try w.beginObject(); + try w.writeStringField("name", "test"); + try w.writeIntField("count", 42); + try w.writeBoolField("active", true); + try w.endObject(); + + try testing.expectEqualStrings("{\"name\":\"test\",\"count\":42,\"active\":true}", w.items()); +} + +test "JsonWriter array" { + const testing = std.testing; + const allocator = testing.allocator; + + var w = JsonWriter.init(allocator); + defer w.deinit(); + + try w.beginArray(); + try w.beginObject(); + try w.writeIntField("id", 1); + try w.endObject(); + try w.beginObject(); + try w.writeIntField("id", 2); + try w.endObject(); + try w.endArray(); + + try testing.expectEqualStrings("[{\"id\":1},{\"id\":2}]", w.items()); +} + +test "JsonWriter escaping" { + const testing = std.testing; + const allocator = testing.allocator; + + var w = JsonWriter.init(allocator); + defer w.deinit(); + + try w.beginObject(); + try w.writeStringField("text", "hello \"world\"\nline2"); + try w.endObject(); + + try testing.expectEqualStrings("{\"text\":\"hello \\\"world\\\"\\nline2\"}", w.items()); +} diff --git a/src/web/response.zig b/src/web/response.zig new file mode 100644 index 0000000..c0d3b42 --- /dev/null +++ b/src/web/response.zig @@ -0,0 +1,117 @@ +const std = @import("std"); +const posix = std.posix; + +pub const CORS_HEADERS = "Access-Control-Allow-Origin: *\r\nAccess-Control-Allow-Methods: GET, POST, PUT, DELETE, OPTIONS\r\nAccess-Control-Allow-Headers: Content-Type, Authorization\r\n"; +pub const CORS_API = "Access-Control-Allow-Origin: *\r\n"; + +/// Write data to socket, logging on failure +fn writeAll(handle: posix.socket_t, data: []const u8) void { + _ = posix.write(handle, data) catch |err| { + std.log.debug("HTTP write failed: {}", .{err}); + }; +} + +/// Send a JSON response +pub fn sendJson(handle: posix.socket_t, json: []const u8) void { + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 200 OK\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + CORS_API, + json.len, + }) catch return; + + writeAll(handle, header); + writeAll(handle, json); +} + +/// Send a JSON response with 201 Created status +pub fn sendJsonCreated(handle: posix.socket_t, json: []const u8) void { + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 201 Created\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + CORS_API, + json.len, + }) catch return; + + writeAll(handle, header); + writeAll(handle, json); +} + +/// Send a success response with a message +pub fn sendSuccess(handle: posix.socket_t, message: []const u8) void { + var buf: [256]u8 = undefined; + const body = std.fmt.bufPrint(&buf, "{{\"success\":true,\"message\":\"{s}\"}}", .{message}) catch return; + sendJson(handle, body); +} + +/// Send a simple success response +pub fn sendOk(handle: posix.socket_t) void { + const body = "{\"success\":true}"; + sendJson(handle, body); +} + +/// Send an error response +pub fn sendError(handle: posix.socket_t, status: u16, message: []const u8) void { + const status_text = switch (status) { + 400 => "Bad Request", + 401 => "Unauthorized", + 404 => "Not Found", + 405 => "Method Not Allowed", + 500 => "Internal Server Error", + else => "Error", + }; + + var body_buf: [256]u8 = undefined; + const body = std.fmt.bufPrint(&body_buf, "{{\"error\":\"{s}\"}}", .{message}) catch return; + + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 {d} {s}\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + status, + status_text, + CORS_API, + body.len, + }) catch return; + + writeAll(handle, header); + writeAll(handle, body); +} + +pub fn sendBadRequest(handle: posix.socket_t, message: []const u8) void { + sendError(handle, 400, message); +} + +pub fn sendUnauthorized(handle: posix.socket_t) void { + const body = "{\"error\":\"Unauthorized\"}"; + var buf: [256]u8 = undefined; + const response = std.fmt.bufPrint(&buf, "HTTP/1.1 401 Unauthorized\r\n{s}WWW-Authenticate: Basic realm=\"nxdns\"\r\nContent-Type: application/json\r\nContent-Length: {d}\r\n\r\n{s}", .{ + CORS_API, + body.len, + body, + }) catch return; + writeAll(handle, response); +} + +pub fn sendNotFound(handle: posix.socket_t) void { + sendError(handle, 404, "Not found"); +} + +pub fn sendMethodNotAllowed(handle: posix.socket_t) void { + sendError(handle, 405, "Method not allowed"); +} + +pub fn sendServerError(handle: posix.socket_t, message: []const u8) void { + sendError(handle, 500, message); +} + +/// Send HTML content +pub fn sendHtml(handle: posix.socket_t, html: []const u8) void { + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 200 OK\r\nContent-Type: text/html\r\nContent-Length: {d}\r\n\r\n", .{html.len}) catch return; + + writeAll(handle, header); + writeAll(handle, html); +} + +/// Send OPTIONS response for CORS preflight +pub fn sendOptions(handle: posix.socket_t) void { + const response = "HTTP/1.1 204 No Content\r\n" ++ CORS_HEADERS ++ "Content-Length: 0\r\n\r\n"; + writeAll(handle, response); +} diff --git a/src/web/server.zig b/src/web/server.zig new file mode 100644 index 0000000..7e18857 --- /dev/null +++ b/src/web/server.zig @@ -0,0 +1,813 @@ +const std = @import("std"); +const net = std.net; +const posix = std.posix; +const Allocator = std.mem.Allocator; +const Database = @import("../storage/db.zig").Database; +const DenylistFetcher = @import("../filter/fetcher.zig").DenylistFetcher; +const events = @import("../events.zig"); +const auth = @import("auth.zig"); +const logger_mod = @import("../storage/logger.zig"); +const response = @import("response.zig"); +const json = @import("json.zig"); + +// API handlers +const stats_api = @import("api/stats.zig"); +const queries_api = @import("api/queries.zig"); +const denylists_api = @import("api/denylists.zig"); +const groups_api = @import("api/groups.zig"); +const clients_api = @import("api/clients.zig"); +const rules_api = @import("api/rules.zig"); +const settings_api = @import("api/settings.zig"); + +pub const WebServer = struct { + listener: net.Server, + db: *Database, + allocator: Allocator, + running: bool, + basic_auth: auth.BasicAuth, + active_connections: std.atomic.Value(u32), + + // SSE (Server-Sent Events) client tracking + sse_clients: std.ArrayListUnmanaged(posix.socket_t), + sse_mutex: std.Thread.Mutex, + + /// Maximum number of concurrent SSE clients to prevent resource exhaustion + const MAX_SSE_CLIENTS: usize = 100; + + /// Maximum number of concurrent HTTP connections + const MAX_CONNECTIONS: u32 = 100; + + /// Poll timeout for SSE keep-alive (milliseconds) + const SSE_POLL_TIMEOUT_MS: i32 = 1000; + + /// Poll timeout for accept loop (milliseconds) + const ACCEPT_POLL_TIMEOUT_MS: i32 = 100; + + /// HTTP request buffer size (limits header size) + const HTTP_REQUEST_BUFFER_SIZE: usize = 8192; + + /// Maximum HTTP body size (1MB) + const MAX_BODY_SIZE: usize = 1024 * 1024; + + pub fn init(bind_addr: net.Address, db: *Database, allocator: Allocator) !WebServer { + const listener = try bind_addr.listen(.{ + .reuse_address = true, + }); + + return WebServer{ + .listener = listener, + .db = db, + .allocator = allocator, + .running = false, + .basic_auth = auth.BasicAuth.disabled(allocator), + .active_connections = std.atomic.Value(u32).init(0), + .sse_clients = std.ArrayListUnmanaged(posix.socket_t){}, + .sse_mutex = std.Thread.Mutex{}, + }; + } + + /// Create a logger subscriber for real-time SSE notifications + pub fn toLoggerSubscriber(self: *WebServer) logger_mod.Subscriber { + return logger_mod.Subscriber{ + .context = self, + .notifyFn = broadcastQueryWrapper, + }; + } + + fn broadcastQueryWrapper(ctx: *anyopaque, entry: logger_mod.LogEntry) void { + const self: *WebServer = @ptrCast(@alignCast(ctx)); + self.broadcastQuery(entry); + } + + /// Broadcast a query log entry to all SSE clients + pub fn broadcastQuery(self: *WebServer, entry: logger_mod.LogEntry) void { + // Look up client name from database (fall back to IP if not found) + var client_name_buf: [128]u8 = undefined; + var client_name: []const u8 = entry.client_ip; + + var stmt = self.db.prepare("SELECT name FROM clients WHERE ip = ?") catch null; + if (stmt) |*s| { + defer s.finalize(); + s.bindText(1, entry.client_ip) catch {}; + if (s.step() catch false) { + if (s.getText(0)) |name| { + // Copy name to local buffer before statement is finalized + const len = @min(name.len, client_name_buf.len); + @memcpy(client_name_buf[0..len], name[0..len]); + client_name = client_name_buf[0..len]; + } + } + } + + var buf: [1024]u8 = undefined; + const json_data = std.fmt.bufPrint(&buf, "data: {{\"timestamp\":{d},\"domain\":\"{s}\",\"client\":\"{s}\",\"qtype\":{d},\"denied\":{s},\"status\":{d},\"reply_type\":{d},\"protocol\":{d},\"dnssec\":{s}}}\n\n", .{ + entry.timestamp, + entry.domain, + client_name, + entry.qtype, + if (entry.isDenied()) "true" else "false", + @intFromEnum(entry.status), + @intFromEnum(entry.reply_type), + @intFromEnum(entry.protocol), + if (entry.dnssec_validated) "true" else "false", + }) catch return; + + self.sse_mutex.lock(); + defer self.sse_mutex.unlock(); + + // Reverse iteration: safe with swapRemove (removed element is already processed) + var i: usize = self.sse_clients.items.len; + while (i > 0) { + i -= 1; + const written = posix.write(self.sse_clients.items[i], json_data) catch { + _ = self.sse_clients.swapRemove(i); + continue; + }; + if (written == 0) { + _ = self.sse_clients.swapRemove(i); + } + } + } + + /// Handle SSE connection - keeps connection open for streaming + fn handleSSE(self: *WebServer, handle: posix.socket_t) void { + // Check client limit before accepting + self.sse_mutex.lock(); + if (self.sse_clients.items.len >= MAX_SSE_CLIENTS) { + self.sse_mutex.unlock(); + std.log.warn("SSE client limit reached ({d}), rejecting connection", .{MAX_SSE_CLIENTS}); + response.sendError(handle, 503, "Too many SSE clients"); + return; + } + self.sse_mutex.unlock(); + + const headers = "HTTP/1.1 200 OK\r\n" ++ + "Content-Type: text/event-stream\r\n" ++ + "Cache-Control: no-cache\r\n" ++ + "Connection: keep-alive\r\n" ++ + "Access-Control-Allow-Origin: *\r\n\r\n"; + _ = posix.write(handle, headers) catch |err| { + std.log.debug("SSE connection write failed (client likely disconnected): {}", .{err}); + return; + }; + + self.sse_mutex.lock(); + // Re-check after acquiring lock (could have changed) + if (self.sse_clients.items.len >= MAX_SSE_CLIENTS) { + self.sse_mutex.unlock(); + return; + } + self.sse_clients.append(self.allocator, handle) catch { + self.sse_mutex.unlock(); + return; + }; + self.sse_mutex.unlock(); + + // Keep connection alive - block until client disconnects + var buf: [1]u8 = undefined; + while (self.running) { + var fds = [1]posix.pollfd{ + .{ + .fd = handle, + .events = posix.POLL.IN, + .revents = 0, + }, + }; + const poll_result = posix.poll(&fds, SSE_POLL_TIMEOUT_MS) catch break; + + if (poll_result > 0) { + if (fds[0].revents & (posix.POLL.HUP | posix.POLL.ERR) != 0) break; + if (fds[0].revents & posix.POLL.IN != 0) { + const n = posix.read(handle, &buf) catch break; + if (n == 0) break; + } + } + } + + // Unregister client + self.sse_mutex.lock(); + defer self.sse_mutex.unlock(); + + for (self.sse_clients.items, 0..) |client, i| { + if (client == handle) { + _ = self.sse_clients.swapRemove(i); + break; + } + } + } + + /// Set HTTP Basic Auth credentials + pub fn setAuth(self: *WebServer, username: []const u8, password: []const u8) void { + self.basic_auth = auth.BasicAuth.init(username, password, self.allocator); + } + + pub fn run(self: *WebServer) !void { + self.running = true; + + while (self.running) { + var fds = [1]posix.pollfd{ + .{ + .fd = self.listener.stream.handle, + .events = posix.POLL.IN, + .revents = 0, + }, + }; + + const poll_result = posix.poll(&fds, ACCEPT_POLL_TIMEOUT_MS) catch |err| { + std.log.warn("HTTP poll error: {}", .{err}); + continue; + }; + + if (poll_result == 0) continue; + if (fds[0].revents & posix.POLL.IN == 0) continue; + + const conn = self.listener.accept() catch |err| { + std.log.warn("HTTP accept error: {}", .{err}); + continue; + }; + + // Check connection limit to prevent resource exhaustion + if (self.active_connections.load(.acquire) >= MAX_CONNECTIONS) { + response.sendError(conn.stream.handle, 503, "Server overloaded"); + posix.close(conn.stream.handle); + continue; + } + + const thread = std.Thread.spawn(.{}, handleConnection, .{ self, conn }) catch |err| { + std.log.warn("Failed to spawn thread: {}", .{err}); + posix.close(conn.stream.handle); + continue; + }; + thread.detach(); + } + } + + /// Read timeout for HTTP requests (5 seconds) - prevents slowloris attacks + const REQUEST_TIMEOUT_SECS = 5; + + fn handleConnection(self: *WebServer, conn: net.Server.Connection) void { + _ = self.active_connections.fetchAdd(1, .monotonic); + defer _ = self.active_connections.fetchSub(1, .monotonic); + defer posix.close(conn.stream.handle); + + // Set read timeout to prevent slowloris attacks + const timeout = posix.timeval{ .sec = REQUEST_TIMEOUT_SECS, .usec = 0 }; + posix.setsockopt(conn.stream.handle, posix.SOL.SOCKET, posix.SO.RCVTIMEO, std.mem.asBytes(&timeout)) catch |err| { + std.log.warn("Failed to set HTTP socket timeout: {}", .{err}); + }; + + var buffer: [HTTP_REQUEST_BUFFER_SIZE]u8 = undefined; + const n = posix.read(conn.stream.handle, &buffer) catch return; + if (n == 0) return; + + const request = buffer[0..n]; + self.handleRequest(conn.stream.handle, request); + } + + fn handleRequest(self: *WebServer, handle: posix.socket_t, request: []const u8) void { + var lines = std.mem.splitSequence(u8, request, "\r\n"); + const first_line = lines.first(); + + var parts = std.mem.splitScalar(u8, first_line, ' '); + const method = parts.first(); + const path = parts.next() orelse "/"; + + // Find body (after \r\n\r\n) + const body = if (std.mem.indexOf(u8, request, "\r\n\r\n")) |idx| + request[idx + 4 ..] + else + ""; + + // Handle OPTIONS (CORS preflight) + if (std.mem.eql(u8, method, "OPTIONS")) { + response.sendOptions(handle); + return; + } + + // Check authentication for API routes + if (self.basic_auth.isEnabled() and std.mem.startsWith(u8, path, "/api/")) { + const auth_header = auth.extractAuthHeader(request); + if (!self.basic_auth.validate(auth_header)) { + response.sendUnauthorized(handle); + return; + } + } + + // Route requests + if (std.mem.eql(u8, path, "/") or std.mem.eql(u8, path, "/index.html")) { + self.sendHtml(handle); + } else if (std.mem.startsWith(u8, path, "/api/")) { + self.handleApi(handle, method, path, body); + } else { + response.sendNotFound(handle); + } + } + + fn handleApi(self: *WebServer, handle: posix.socket_t, method: []const u8, path: []const u8, body: []const u8) void { + // Stats + if (std.mem.eql(u8, path, "/api/stats")) { + stats_api.sendStats(self.db, handle, self.allocator); + return; + } + + // Queries + if (std.mem.eql(u8, path, "/api/queries")) { + queries_api.sendQueries(self.db, handle, self.allocator); + return; + } + if (std.mem.eql(u8, path, "/api/queries/live")) { + self.handleSSE(handle); + return; + } + + // Denylists + if (std.mem.eql(u8, path, "/api/denylists")) { + if (std.mem.eql(u8, method, "GET")) { + denylists_api.list(self.db, handle, self.allocator); + } else if (std.mem.eql(u8, method, "POST")) { + self.addDenylist(handle, body); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + if (std.mem.startsWith(u8, path, "/api/denylists/")) { + const suffix = path[15..]; + // Check for /api/denylists/update endpoint + if (std.mem.eql(u8, suffix, "update")) { + if (std.mem.eql(u8, method, "POST")) { + self.triggerDenylistUpdate(handle); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + // Check for /api/denylists/:id/update endpoint + if (std.mem.endsWith(u8, suffix, "/update")) { + const id_str = suffix[0 .. suffix.len - 7]; // Remove "/update" + const id = std.fmt.parseInt(i64, id_str, 10) catch { + response.sendBadRequest(handle, "Invalid denylist ID"); + return; + }; + if (std.mem.eql(u8, method, "POST")) { + self.updateSingleDenylist(handle, id); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + // Otherwise it's /api/denylists/:id + const id = std.fmt.parseInt(i64, suffix, 10) catch { + response.sendBadRequest(handle, "Invalid denylist ID"); + return; + }; + if (std.mem.eql(u8, method, "DELETE")) { + denylists_api.delete(self.db, handle, id); + events.signalDenylistReload(); + } else if (std.mem.eql(u8, method, "PUT")) { + denylists_api.toggle(self.db, handle, self.allocator, id, body); + events.signalDenylistReload(); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + + // Groups + if (std.mem.eql(u8, path, "/api/groups")) { + if (std.mem.eql(u8, method, "GET")) { + groups_api.list(self.db, handle, self.allocator); + } else if (std.mem.eql(u8, method, "POST")) { + groups_api.add(self.db, handle, self.allocator, body); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + if (std.mem.startsWith(u8, path, "/api/groups/")) { + const id = std.fmt.parseInt(i64, path[12..], 10) catch { + response.sendBadRequest(handle, "Invalid group ID"); + return; + }; + if (std.mem.eql(u8, method, "PUT")) { + groups_api.update(self.db, handle, self.allocator, id, body); + } else if (std.mem.eql(u8, method, "DELETE")) { + groups_api.delete(self.db, handle, id); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + + // Clients + if (std.mem.eql(u8, path, "/api/clients")) { + if (std.mem.eql(u8, method, "GET")) { + clients_api.list(self.db, handle, self.allocator); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + if (std.mem.startsWith(u8, path, "/api/clients/")) { + const id = std.fmt.parseInt(i64, path[13..], 10) catch { + response.sendBadRequest(handle, "Invalid client ID"); + return; + }; + if (std.mem.eql(u8, method, "PUT")) { + clients_api.update(self.db, handle, self.allocator, id, body); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + + // Rules + if (std.mem.eql(u8, path, "/api/rules")) { + if (std.mem.eql(u8, method, "GET")) { + rules_api.list(self.db, handle, self.allocator); + } else if (std.mem.eql(u8, method, "POST")) { + rules_api.add(self.db, handle, self.allocator, body); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + if (std.mem.startsWith(u8, path, "/api/rules/")) { + const id = std.fmt.parseInt(i64, path[11..], 10) catch { + response.sendBadRequest(handle, "Invalid rule ID"); + return; + }; + if (std.mem.eql(u8, method, "DELETE")) { + rules_api.delete(self.db, handle, id); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + + // Settings + if (std.mem.eql(u8, path, "/api/settings")) { + if (std.mem.eql(u8, method, "GET")) { + settings_api.list(self.db, handle, self.allocator); + } else if (std.mem.eql(u8, method, "PUT")) { + settings_api.update(self.db, handle, self.allocator, body); + } else { + response.sendMethodNotAllowed(handle); + } + return; + } + + // Not found + response.sendNotFound(handle); + } + + fn triggerDenylistUpdate(self: *WebServer, handle: posix.socket_t) void { + // Heap-allocate the fetcher so it survives after this function returns + const fetcher = self.allocator.create(DenylistFetcher) catch { + response.sendServerError(handle, "Memory allocation failed"); + return; + }; + fetcher.* = DenylistFetcher.init(self.db, self.allocator); + + const ThreadContext = struct { + fetcher: *DenylistFetcher, + allocator: std.mem.Allocator, + + fn run(ctx: *@This()) void { + const alloc = ctx.allocator; + const dlf = ctx.fetcher; + defer alloc.destroy(ctx); + defer alloc.destroy(dlf); + + dlf.updateAll() catch |err| { + std.log.warn("Denylist update failed: {}", .{err}); + }; + } + }; + + const ctx = self.allocator.create(ThreadContext) catch { + self.allocator.destroy(fetcher); + response.sendServerError(handle, "Memory allocation failed"); + return; + }; + ctx.* = .{ .fetcher = fetcher, .allocator = self.allocator }; + + const thread = std.Thread.spawn(.{}, ThreadContext.run, .{ctx}) catch { + self.allocator.destroy(ctx); + self.allocator.destroy(fetcher); + response.sendServerError(handle, "Failed to start update"); + return; + }; + thread.detach(); + + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeBoolField("success", true) catch return; + w.writeStringField("message", "Update started in background") catch return; + w.endObject() catch return; + + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 202 Accepted\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + response.CORS_API, + w.items().len, + }) catch return; + + _ = posix.write(handle, header) catch |err| { + std.log.debug("Failed to write HTTP header: {}", .{err}); + return; + }; + _ = posix.write(handle, w.items()) catch |err| { + std.log.debug("Failed to write HTTP body: {}", .{err}); + }; + } + + const AddDenylistRequest = struct { + url: []const u8, + comment: ?[]const u8 = null, + }; + + /// Add a denylist and automatically fetch its domains + fn addDenylist(self: *WebServer, handle: posix.socket_t, body: []const u8) void { + const parsed = std.json.parseFromSlice(AddDenylistRequest, self.allocator, body, .{ + .ignore_unknown_fields = true, + }) catch { + response.sendBadRequest(handle, "Invalid JSON"); + return; + }; + defer parsed.deinit(); + + const url = parsed.value.url; + + // Basic URL validation + if (url.len == 0 or url.len > 2048) { + response.sendBadRequest(handle, "Invalid URL length"); + return; + } + if (!std.mem.startsWith(u8, url, "http://") and !std.mem.startsWith(u8, url, "https://")) { + response.sendBadRequest(handle, "URL must start with http:// or https://"); + return; + } + + // Insert into database + var stmt = self.db.prepare("INSERT INTO denylist_sources (url, comment, enabled, domain_count) VALUES (?, ?, 1, 0)") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindText(1, url) catch return; + if (parsed.value.comment) |comment| { + stmt.bindText(2, comment) catch return; + } else { + stmt.bindNull(2) catch return; + } + + _ = stmt.step() catch { + response.sendServerError(handle, "Insert failed (URL may already exist)"); + return; + }; + + const id = self.db.lastInsertRowId(); + + // Link to default group + linkGroup: { + var group_stmt = self.db.prepare("INSERT INTO group_sources (group_id, source_id) VALUES (0, ?)") catch break :linkGroup; + defer group_stmt.finalize(); + group_stmt.bindInt(1, id) catch break :linkGroup; + _ = group_stmt.step() catch {}; + } + + // Spawn background thread to fetch domains (don't block the API) + const FetchContext = struct { + fetcher: *DenylistFetcher, + url: []const u8, + id: i64, + allocator: std.mem.Allocator, + + fn run(ctx: *@This()) void { + const alloc = ctx.allocator; + const dlf = ctx.fetcher; + const dl_url = ctx.url; + const dl_id = ctx.id; + defer alloc.destroy(ctx); + defer alloc.destroy(dlf); + defer alloc.free(dl_url); + + const count = dlf.fetchAndUpdate(dl_id, dl_url) catch |err| { + std.log.warn("Background fetch failed for denylist {d}: {}", .{ dl_id, err }); + return; + }; + std.log.info("Denylist {d} fetched: {d} domains", .{ dl_id, count }); + events.signalDenylistReload(); + } + }; + + const fetcher = self.allocator.create(DenylistFetcher) catch { + // Denylist added, but can't start background fetch + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.writeBoolField("fetching", false) catch return; + w.endObject() catch return; + response.sendJsonCreated(handle, w.items()); + return; + }; + fetcher.* = DenylistFetcher.init(self.db, self.allocator); + + const url_copy = self.allocator.dupe(u8, url) catch { + self.allocator.destroy(fetcher); + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.writeBoolField("fetching", false) catch return; + w.endObject() catch return; + response.sendJsonCreated(handle, w.items()); + return; + }; + + const ctx = self.allocator.create(FetchContext) catch { + self.allocator.free(url_copy); + self.allocator.destroy(fetcher); + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.writeBoolField("fetching", false) catch return; + w.endObject() catch return; + response.sendJsonCreated(handle, w.items()); + return; + }; + ctx.* = .{ .fetcher = fetcher, .url = url_copy, .id = id, .allocator = self.allocator }; + + const thread = std.Thread.spawn(.{}, FetchContext.run, .{ctx}) catch { + self.allocator.destroy(ctx); + self.allocator.free(url_copy); + self.allocator.destroy(fetcher); + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.writeBoolField("fetching", false) catch return; + w.endObject() catch return; + response.sendJsonCreated(handle, w.items()); + return; + }; + thread.detach(); + + // Return immediately - fetch happens in background + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeIntField("id", id) catch return; + w.writeBoolField("success", true) catch return; + w.writeBoolField("fetching", true) catch return; + w.endObject() catch return; + + response.sendJsonCreated(handle, w.items()); + } + + fn updateSingleDenylist(self: *WebServer, handle: posix.socket_t, id: i64) void { + var stmt = self.db.prepare("SELECT url FROM denylist_sources WHERE id = ?") catch { + response.sendServerError(handle, "Database error"); + return; + }; + defer stmt.finalize(); + + stmt.bindInt(1, id) catch return; + + const has_row = stmt.step() catch { + response.sendServerError(handle, "Database error"); + return; + }; + if (!has_row) { + response.sendError(handle, 404, "Denylist not found"); + return; + } + + const url = stmt.getText(0) orelse { + response.sendServerError(handle, "No URL for denylist"); + return; + }; + + // Spawn background thread to fetch (don't block the API) + const FetchContext = struct { + fetcher: *DenylistFetcher, + url: []const u8, + id: i64, + allocator: std.mem.Allocator, + + fn run(ctx: *@This()) void { + const alloc = ctx.allocator; + const dlf = ctx.fetcher; + const dl_url = ctx.url; + const dl_id = ctx.id; + defer alloc.destroy(ctx); + defer alloc.destroy(dlf); + defer alloc.free(dl_url); + + const count = dlf.fetchAndUpdate(dl_id, dl_url) catch |err| { + std.log.warn("Background update failed for denylist {d}: {}", .{ dl_id, err }); + return; + }; + std.log.info("Denylist {d} updated: {d} domains", .{ dl_id, count }); + events.signalDenylistReload(); + } + }; + + const fetcher = self.allocator.create(DenylistFetcher) catch { + response.sendServerError(handle, "Memory allocation failed"); + return; + }; + fetcher.* = DenylistFetcher.init(self.db, self.allocator); + + const url_copy = self.allocator.dupe(u8, url) catch { + self.allocator.destroy(fetcher); + response.sendServerError(handle, "Memory allocation failed"); + return; + }; + + const ctx = self.allocator.create(FetchContext) catch { + self.allocator.free(url_copy); + self.allocator.destroy(fetcher); + response.sendServerError(handle, "Memory allocation failed"); + return; + }; + ctx.* = .{ .fetcher = fetcher, .url = url_copy, .id = id, .allocator = self.allocator }; + + const thread = std.Thread.spawn(.{}, FetchContext.run, .{ctx}) catch { + self.allocator.destroy(ctx); + self.allocator.free(url_copy); + self.allocator.destroy(fetcher); + response.sendServerError(handle, "Failed to start update"); + return; + }; + thread.detach(); + + // Return 202 Accepted - update happens in background + var w = json.JsonWriter.init(self.allocator); + defer w.deinit(); + w.beginObject() catch return; + w.writeBoolField("success", true) catch return; + w.writeStringField("message", "Update started") catch return; + w.endObject() catch return; + + var header_buf: [256]u8 = undefined; + const header = std.fmt.bufPrint(&header_buf, "HTTP/1.1 202 Accepted\r\n{s}Content-Type: application/json\r\nContent-Length: {d}\r\n\r\n", .{ + response.CORS_API, + w.items().len, + }) catch return; + + _ = posix.write(handle, header) catch return; + _ = posix.write(handle, w.items()) catch {}; + } + + fn sendHtml(self: *WebServer, handle: posix.socket_t) void { + _ = self; + const html = @embedFile("index.html"); + response.sendHtml(handle, html); + } + + pub fn stop(self: *WebServer) void { + self.running = false; + } + + /// Wait for active connections to finish (with timeout) + pub fn waitForConnections(self: *WebServer, timeout_ms: u64) void { + const start = std.time.milliTimestamp(); + while (self.active_connections.load(.acquire) > 0) { + const elapsed: u64 = @intCast(std.time.milliTimestamp() - start); + if (elapsed >= timeout_ms) { + std.log.warn("Web: {} connections still active after timeout", .{self.active_connections.load(.acquire)}); + break; + } + std.posix.nanosleep(0, 10 * std.time.ns_per_ms); + } + } + + pub fn deinit(self: *WebServer) void { + self.sse_clients.deinit(self.allocator); + self.listener.deinit(); + } +}; + +test "WebServer init" { + const testing = std.testing; + const allocator = testing.allocator; + + var db = try Database.open(":memory:", allocator); + defer db.close(); + + const addr = net.Address.initIp4(.{ 127, 0, 0, 1 }, 0); + var server = WebServer.init(addr, &db, allocator) catch |err| { + std.debug.print("WebServer init failed: {}\n", .{err}); + return error.TestFailed; + }; + defer server.deinit(); +} diff --git a/tests/dns/protocol_tests.zig b/tests/dns/protocol_tests.zig new file mode 100644 index 0000000..c0f6cfb --- /dev/null +++ b/tests/dns/protocol_tests.zig @@ -0,0 +1,466 @@ +const std = @import("std"); +const testing = std.testing; + +const packet = @import("packet"); +const Header = @import("header").Header; +const Name = @import("name").Name; +const Question = @import("question").Question; +const types = @import("types"); + +// ============================================================================ +// Real DNS packet captures for testing +// ============================================================================ + +/// Standard A query for google.com +/// Captured from: dig google.com A +const GOOGLE_A_QUERY = [_]u8{ + // Header + 0xAB, 0xCD, // ID: 0xABCD + 0x01, 0x00, // Flags: standard query, RD=1 + 0x00, 0x01, // QDCOUNT: 1 + 0x00, 0x00, // ANCOUNT: 0 + 0x00, 0x00, // NSCOUNT: 0 + 0x00, 0x00, // ARCOUNT: 0 + // Question: google.com A IN + 0x06, 'g', 'o', 'o', 'g', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, // null terminator + 0x00, 0x01, // QTYPE: A (1) + 0x00, 0x01, // QCLASS: IN (1) +}; + +/// AAAA query for example.org +const EXAMPLE_AAAA_QUERY = [_]u8{ + // Header + 0x12, 0x34, // ID + 0x01, 0x00, // Flags: standard query, RD=1 + 0x00, 0x01, // QDCOUNT: 1 + 0x00, 0x00, // ANCOUNT: 0 + 0x00, 0x00, // NSCOUNT: 0 + 0x00, 0x00, // ARCOUNT: 0 + // Question: example.org AAAA IN + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'o', 'r', 'g', + 0x00, + 0x00, 0x1C, // QTYPE: AAAA (28) + 0x00, 0x01, // QCLASS: IN +}; + +/// Response with A record +const SIMPLE_A_RESPONSE = [_]u8{ + // Header + 0xAB, 0xCD, // ID + 0x81, 0x80, // Flags: response, RD=1, RA=1 + 0x00, 0x01, // QDCOUNT: 1 + 0x00, 0x01, // ANCOUNT: 1 + 0x00, 0x00, // NSCOUNT: 0 + 0x00, 0x00, // ARCOUNT: 0 + // Question: google.com A IN (with compression) + 0x06, 'g', 'o', 'o', 'g', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, + 0x00, 0x01, // QTYPE: A + 0x00, 0x01, // QCLASS: IN + // Answer: A record using compression pointer + 0xC0, 0x0C, // Name pointer to offset 12 (google.com) + 0x00, 0x01, // TYPE: A + 0x00, 0x01, // CLASS: IN + 0x00, 0x00, 0x01, 0x2C, // TTL: 300 seconds + 0x00, 0x04, // RDLENGTH: 4 + 0xD8, 0x3A, 0xD3, 0x8E, // RDATA: 216.58.211.142 +}; + +/// Response with CNAME chain +const CNAME_RESPONSE = [_]u8{ + // Header + 0x55, 0x66, // ID + 0x81, 0x80, // Flags: response, RD=1, RA=1 + 0x00, 0x01, // QDCOUNT: 1 + 0x00, 0x02, // ANCOUNT: 2 (CNAME + A) + 0x00, 0x00, // NSCOUNT: 0 + 0x00, 0x00, // ARCOUNT: 0 + // Question: www.example.com A IN + 0x03, 'w', 'w', 'w', + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, + 0x00, 0x01, // QTYPE: A + 0x00, 0x01, // QCLASS: IN + // Answer 1: CNAME www.example.com -> example.com + 0xC0, 0x0C, // Name pointer to www.example.com + 0x00, 0x05, // TYPE: CNAME + 0x00, 0x01, // CLASS: IN + 0x00, 0x00, 0x0E, 0x10, // TTL: 3600 + 0x00, 0x02, // RDLENGTH: 2 (compression pointer) + 0xC0, 0x10, // RDATA: pointer to example.com + // Answer 2: A record for example.com + 0xC0, 0x10, // Name pointer to example.com + 0x00, 0x01, // TYPE: A + 0x00, 0x01, // CLASS: IN + 0x00, 0x00, 0x01, 0x2C, // TTL: 300 + 0x00, 0x04, // RDLENGTH: 4 + 0x5D, 0xB8, 0xD8, 0x22, // RDATA: 93.184.216.34 +}; + +/// Malformed packet - truncated header +const MALFORMED_TRUNCATED = [_]u8{ + 0x12, 0x34, // Only 2 bytes, header needs 12 +}; + +/// Malformed packet - invalid compression pointer (loop) +const MALFORMED_COMPRESSION_LOOP = [_]u8{ + // Header + 0x00, 0x01, + 0x01, 0x00, + 0x00, 0x01, + 0x00, 0x00, + 0x00, 0x00, + 0x00, 0x00, + // Question with self-referencing pointer + 0xC0, 0x0C, // Points to itself + 0x00, 0x01, + 0x00, 0x01, +}; + +// ============================================================================ +// Header Tests +// ============================================================================ + +test "Header.parse - standard query" { + const h = try Header.parse(&GOOGLE_A_QUERY); + + try testing.expectEqual(@as(u16, 0xABCD), h.id); + try testing.expect(!h.qr); // Query, not response + try testing.expectEqual(types.OpCode.Query, h.opcode); + try testing.expect(!h.aa); // Not authoritative + try testing.expect(!h.tc); // Not truncated + try testing.expect(h.rd); // Recursion desired + try testing.expect(!h.ra); // Recursion not available (query) + try testing.expectEqual(types.RCode.NoError, h.rcode); + try testing.expectEqual(@as(u16, 1), h.qdcount); + try testing.expectEqual(@as(u16, 0), h.ancount); + try testing.expectEqual(@as(u16, 0), h.nscount); + try testing.expectEqual(@as(u16, 0), h.arcount); +} + +test "Header.parse - standard response" { + const h = try Header.parse(&SIMPLE_A_RESPONSE); + + try testing.expectEqual(@as(u16, 0xABCD), h.id); + try testing.expect(h.qr); // Response + try testing.expect(h.rd); // RD copied from query + try testing.expect(h.ra); // Recursion available + try testing.expectEqual(@as(u16, 1), h.qdcount); + try testing.expectEqual(@as(u16, 1), h.ancount); +} + +test "Header.parse - buffer too small" { + const result = Header.parse(&MALFORMED_TRUNCATED); + try testing.expectError(error.BufferTooSmall, result); +} + +test "Header.encode - roundtrip" { + const original = try Header.parse(&GOOGLE_A_QUERY); + + var buf: [12]u8 = undefined; + original.encode(&buf); + + const decoded = try Header.parse(&buf); + try testing.expectEqual(original.id, decoded.id); + try testing.expectEqual(original.qr, decoded.qr); + try testing.expectEqual(original.opcode, decoded.opcode); + try testing.expectEqual(original.rd, decoded.rd); + try testing.expectEqual(original.qdcount, decoded.qdcount); +} + +// ============================================================================ +// Name Tests +// ============================================================================ + +test "Name.parse - simple domain" { + const allocator = testing.allocator; + const buffer = [_]u8{ 0x06, 'g', 'o', 'o', 'g', 'l', 'e', 0x03, 'c', 'o', 'm', 0x00 }; + + const result = try Name.parse(&buffer, &buffer, allocator); + defer result.name.deinit(); + + try testing.expectEqual(@as(usize, 12), result.bytes_read); + + var str_buf: [256]u8 = undefined; + const str = result.name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("google.com", str); +} + +test "Name.parse - root domain" { + const allocator = testing.allocator; + const buffer = [_]u8{0x00}; // Just null byte = root + + const result = try Name.parse(&buffer, &buffer, allocator); + defer result.name.deinit(); + + try testing.expectEqual(@as(usize, 1), result.bytes_read); + + var str_buf: [256]u8 = undefined; + const str = result.name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("", str); +} + +test "Name.parse - compression pointer" { + const allocator = testing.allocator; + + // Full packet with question followed by answer using compression + const result = try Name.parse(SIMPLE_A_RESPONSE[27..], &SIMPLE_A_RESPONSE, allocator); + defer result.name.deinit(); + + var str_buf: [256]u8 = undefined; + const str = result.name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("google.com", str); +} + +test "Name.parse - max label length (63)" { + const allocator = testing.allocator; + var buffer: [67]u8 = undefined; + buffer[0] = 63; // Label length = max + for (1..64) |i| { + buffer[i] = 'a'; + } + buffer[64] = 0x03; + buffer[65] = 'c'; + buffer[66] = 'o'; + // Would need more bytes for full domain, but testing max label + + // This should work (63 is max label length) + const result = Name.parse(buffer[0..67], buffer[0..67], allocator); + if (result) |r| { + r.name.deinit(); + } else |_| { + // May fail due to incomplete buffer, which is fine + } +} + +test "Name.parse - label too long (64+)" { + const allocator = testing.allocator; + var buffer: [68]u8 = undefined; + buffer[0] = 64; // Label length > 63 is invalid (and not a pointer) + @memset(buffer[1..65], 'a'); + buffer[65] = 0x00; + + const result = Name.parse(&buffer, &buffer, allocator); + try testing.expectError(error.InvalidLabel, result); +} + +// ============================================================================ +// Question Tests +// ============================================================================ + +test "Question.parse - A record query" { + const allocator = testing.allocator; + + const result = try Question.parse(GOOGLE_A_QUERY[12..], &GOOGLE_A_QUERY, allocator); + defer result.question.deinit(); + + try testing.expectEqual(types.QType.A, result.question.qtype); + try testing.expectEqual(types.QClass.IN, result.question.qclass); + + var str_buf: [256]u8 = undefined; + const domain = result.question.name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("google.com", domain); +} + +test "Question.parse - AAAA record query" { + const allocator = testing.allocator; + + const result = try Question.parse(EXAMPLE_AAAA_QUERY[12..], &EXAMPLE_AAAA_QUERY, allocator); + defer result.question.deinit(); + + try testing.expectEqual(types.QType.AAAA, result.question.qtype); + try testing.expectEqual(types.QClass.IN, result.question.qclass); + + var str_buf: [256]u8 = undefined; + const domain = result.question.name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("example.org", domain); +} + +// ============================================================================ +// Packet Tests +// ============================================================================ + +test "Packet.parse - simple query" { + const allocator = testing.allocator; + + var pkt = try packet.Packet.parse(&GOOGLE_A_QUERY, allocator); + defer pkt.deinit(); + + try testing.expectEqual(@as(u16, 0xABCD), pkt.header.id); + try testing.expect(!pkt.header.qr); + try testing.expectEqual(@as(usize, 1), pkt.questions.len); + try testing.expectEqual(@as(usize, 0), pkt.answers.len); + + var str_buf: [256]u8 = undefined; + const domain = pkt.questions[0].name.toStringBuf(&str_buf).?; + try testing.expectEqualStrings("google.com", domain); + try testing.expectEqual(types.QType.A, pkt.questions[0].qtype); +} + +test "Packet.parse - response with A record" { + const allocator = testing.allocator; + + var pkt = try packet.Packet.parse(&SIMPLE_A_RESPONSE, allocator); + defer pkt.deinit(); + + try testing.expect(pkt.header.qr); + try testing.expectEqual(@as(usize, 1), pkt.questions.len); + try testing.expectEqual(@as(usize, 1), pkt.answers.len); + + // Check the A record + const answer = pkt.answers[0]; + try testing.expectEqual(types.QType.A, answer.rtype); + try testing.expectEqual(@as(u32, 300), answer.ttl); + + // Verify IP address + const ip = answer.getA().?; + try testing.expectEqual([4]u8{ 0xD8, 0x3A, 0xD3, 0x8E }, ip); +} + +test "Packet.parse - CNAME chain" { + const allocator = testing.allocator; + + var pkt = try packet.Packet.parse(&CNAME_RESPONSE, allocator); + defer pkt.deinit(); + + try testing.expectEqual(@as(usize, 2), pkt.answers.len); + + // First answer should be CNAME + try testing.expectEqual(types.QType.CNAME, pkt.answers[0].rtype); + + // Second answer should be A record + try testing.expectEqual(types.QType.A, pkt.answers[1].rtype); + + const ip = pkt.answers[1].getA().?; + try testing.expectEqual([4]u8{ 0x5D, 0xB8, 0xD8, 0x22 }, ip); +} + +test "Packet.parse - truncated packet" { + const allocator = testing.allocator; + + const result = packet.Packet.parse(&MALFORMED_TRUNCATED, allocator); + try testing.expectError(error.HeaderParseError, result); +} + +test "Packet.encode - roundtrip" { + const allocator = testing.allocator; + + var original = try packet.Packet.parse(&GOOGLE_A_QUERY, allocator); + defer original.deinit(); + + var buf: [512]u8 = undefined; + const encoded_len = try original.encode(&buf); + + var decoded = try packet.Packet.parse(buf[0..encoded_len], allocator); + defer decoded.deinit(); + + try testing.expectEqual(original.header.id, decoded.header.id); + try testing.expectEqual(original.questions.len, decoded.questions.len); +} + +test "Packet.createBlockedResponse - creates valid response" { + const allocator = testing.allocator; + + var query = try packet.Packet.parse(&GOOGLE_A_QUERY, allocator); + defer query.deinit(); + + var response = try packet.Packet.createBlockedResponse(&query, allocator); + defer response.deinit(); + + try testing.expect(response.header.qr); // Is response + try testing.expectEqual(query.header.id, response.header.id); + try testing.expectEqual(@as(usize, 1), response.answers.len); + + // Should return 0.0.0.0 for blocked + const ip = response.answers[0].getA().?; + try testing.expectEqual([4]u8{ 0, 0, 0, 0 }, ip); +} + +test "Packet.createNxdomainResponse - creates valid NXDOMAIN" { + const allocator = testing.allocator; + + var query = try packet.Packet.parse(&GOOGLE_A_QUERY, allocator); + defer query.deinit(); + + var response = try packet.Packet.createNxdomainResponse(&query, allocator); + defer response.deinit(); + + try testing.expect(response.header.qr); + try testing.expectEqual(types.RCode.NXDomain, response.header.rcode); + try testing.expectEqual(@as(usize, 0), response.answers.len); +} + +// ============================================================================ +// Edge Cases +// ============================================================================ + +test "DNS max name length (253 chars)" { + const allocator = testing.allocator; + + // Build a name with max length: 63.63.63.63 = 253 chars + labels + var buffer: [512]u8 = undefined; + var pos: usize = 12; // Skip header + + // Add 4 labels of 63 chars each + for (0..4) |_| { + buffer[pos] = 63; + pos += 1; + @memset(buffer[pos .. pos + 63], 'a'); + pos += 63; + } + buffer[pos] = 0; // Null terminator + pos += 1; + + // Add qtype and qclass + buffer[pos] = 0x00; + buffer[pos + 1] = 0x01; + buffer[pos + 2] = 0x00; + buffer[pos + 3] = 0x01; + pos += 4; + + // Set up header + @memset(buffer[0..12], 0); + buffer[2] = 0x01; // RD=1 + buffer[5] = 0x01; // QDCOUNT=1 + + const pkt_result = packet.Packet.parse(buffer[0..pos], allocator); + // This may fail due to name being too long (4*63=252 + dots) + if (pkt_result) |*pkt| { + pkt.deinit(); + } else |_| {} +} + +test "DNS various query types" { + const allocator = testing.allocator; + const query_types = [_]types.QType{ .A, .AAAA, .CNAME, .MX, .NS, .TXT, .SOA, .PTR, .SRV }; + + for (query_types) |qtype| { + var buf: [64]u8 = undefined; + @memset(buf[0..12], 0); + buf[2] = 0x01; // RD + buf[5] = 0x01; // QDCOUNT + + // Simple question: a.b + buf[12] = 0x01; + buf[13] = 'a'; + buf[14] = 0x01; + buf[15] = 'b'; + buf[16] = 0x00; + + const qtype_val = @intFromEnum(qtype); + buf[17] = @intCast((qtype_val >> 8) & 0xFF); + buf[18] = @intCast(qtype_val & 0xFF); + buf[19] = 0x00; + buf[20] = 0x01; // CLASS IN + + var pkt = try packet.Packet.parse(buf[0..21], allocator); + defer pkt.deinit(); + + try testing.expectEqual(qtype, pkt.questions[0].qtype); + } +} diff --git a/tests/integration_tests.zig b/tests/integration_tests.zig new file mode 100644 index 0000000..c410268 --- /dev/null +++ b/tests/integration_tests.zig @@ -0,0 +1,917 @@ +const std = @import("std"); +const testing = std.testing; + +const handler_mod = @import("handler"); +const rate_limiter = @import("rate_limiter"); +const cache = @import("cache"); +const blocklist_mod = @import("blocklist"); +const packet = @import("packet"); +const types = @import("types"); +const Name = @import("name").Name; + +// ============================================================================ +// Test DNS Query/Response Packets +// ============================================================================ + +/// Standard A query for example.com +fn createTestQuery() [29]u8 { + return [_]u8{ + // Header + 0x00, 0x01, // ID + 0x01, 0x00, // Flags: standard query, RD=1 + 0x00, 0x01, // QDCOUNT: 1 + 0x00, 0x00, // ANCOUNT: 0 + 0x00, 0x00, // NSCOUNT: 0 + 0x00, 0x00, // ARCOUNT: 0 + // Question: example.com A IN + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, // null + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + }; +} + +/// Create a query for a specific domain +fn createQueryForDomain(domain: []const u8, buf: *[512]u8) usize { + // Header + buf[0] = 0x00; + buf[1] = 0x02; // ID = 2 + buf[2] = 0x01; + buf[3] = 0x00; // RD=1 + buf[4] = 0x00; + buf[5] = 0x01; // QDCOUNT=1 + buf[6] = 0x00; + buf[7] = 0x00; + buf[8] = 0x00; + buf[9] = 0x00; + buf[10] = 0x00; + buf[11] = 0x00; + + // Question section - encode domain name + var pos: usize = 12; + + // Split domain by dots and encode labels + var iter = std.mem.splitScalar(u8, domain, '.'); + while (iter.next()) |label| { + if (label.len > 63 or label.len == 0) continue; + buf[pos] = @intCast(label.len); + pos += 1; + @memcpy(buf[pos..][0..label.len], label); + pos += label.len; + } + buf[pos] = 0; // Null terminator + pos += 1; + + // QTYPE = A + buf[pos] = 0x00; + buf[pos + 1] = 0x01; + // QCLASS = IN + buf[pos + 2] = 0x00; + buf[pos + 3] = 0x01; + pos += 4; + + return pos; +} + +// ============================================================================ +// Mock Upstream for Testing +// ============================================================================ + +const MockUpstream = struct { + response: ?[]const u8, + allocator: std.mem.Allocator, + call_count: usize = 0, + + pub fn init(response: ?[]const u8, allocator: std.mem.Allocator) MockUpstream { + return .{ + .response = response, + .allocator = allocator, + .call_count = 0, + }; + } + + pub fn toHandlerUpstream(self: *MockUpstream) handler_mod.Upstream { + return .{ + .context = self, + .queryFn = queryWrapper, + }; + } + + fn queryWrapper(ctx: *anyopaque, _: []const u8, allocator: std.mem.Allocator) ?[]const u8 { + const self: *MockUpstream = @ptrCast(@alignCast(ctx)); + self.call_count += 1; + if (self.response) |r| { + return allocator.dupe(u8, r) catch null; + } + return null; + } +}; + +// ============================================================================ +// Handler Tests +// ============================================================================ + +test "Handler - returns SERVFAIL without upstream" { + const allocator = testing.allocator; + + var handler = handler_mod.Handler.init(allocator); + const query = createTestQuery(); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&query, addr, allocator); + + if (response) |r| { + defer allocator.free(r); + // Should be at least header size + try testing.expect(r.len >= types.DNS_HEADER_SIZE); + // Check RCODE is SERVFAIL (2) - in flags byte + const rcode = r[3] & 0x0F; + try testing.expectEqual(@as(u8, 2), rcode); + } +} + +test "Handler - cache integration" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + var handler = handler_mod.Handler.init(allocator); + handler.setCache(dns_cache.toHandlerCache()); + + // Pre-populate cache + const cached_response = [_]u8{ + 0x00, 0x01, // ID (will be overwritten) + 0x81, 0x80, // Flags: response + 0x00, 0x01, // QDCOUNT + 0x00, 0x01, // ANCOUNT + 0x00, 0x00, // NSCOUNT + 0x00, 0x00, // ARCOUNT + // Question + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', 0x00, + 0x00, 0x01, 0x00, 0x01, + // Answer + 0xC0, 0x0C, // Compression pointer + 0x00, 0x01, // TYPE A + 0x00, 0x01, // CLASS IN + 0x00, 0x00, 0x01, 0x2C, // TTL 300 + 0x00, 0x04, // RDLENGTH + 0x01, 0x02, 0x03, 0x04, // IP: 1.2.3.4 + }; + + dns_cache.put("example.com", types.QType.A, &cached_response, 300); + + const query = createTestQuery(); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&query, addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // ID should be updated to match query + try testing.expectEqual(@as(u8, 0x00), r[0]); + try testing.expectEqual(@as(u8, 0x01), r[1]); + } +} + +// ============================================================================ +// Rate Limiter Tests +// ============================================================================ + +test "RateLimiter - allows requests under limit" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 10, + .window_ms = 1000, + }); + defer limiter.deinit(); + + // Should allow first 10 requests + for (0..10) |_| { + try testing.expect(limiter.checkRequest("192.168.1.1")); + } +} + +test "RateLimiter - blocks requests over limit" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 5, + .window_ms = 1000, + }); + defer limiter.deinit(); + + // First 5 should be allowed + for (0..5) |_| { + try testing.expect(limiter.checkRequest("10.0.0.1")); + } + + // 6th should be blocked + try testing.expect(!limiter.checkRequest("10.0.0.1")); +} + +test "RateLimiter - tracks clients independently" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 2, + .window_ms = 1000, + }); + defer limiter.deinit(); + + // Client A uses quota + try testing.expect(limiter.checkRequest("192.168.1.1")); + try testing.expect(limiter.checkRequest("192.168.1.1")); + try testing.expect(!limiter.checkRequest("192.168.1.1")); // Blocked + + // Client B still has quota + try testing.expect(limiter.checkRequest("192.168.1.2")); + try testing.expect(limiter.checkRequest("192.168.1.2")); + try testing.expect(!limiter.checkRequest("192.168.1.2")); // Blocked +} + +test "RateLimiter - can be disabled" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 1, + .enabled = false, + }); + defer limiter.deinit(); + + // All requests allowed when disabled + for (0..100) |_| { + try testing.expect(limiter.checkRequest("any.ip")); + } +} + +test "RateLimiter - statistics tracking" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 3, + .window_ms = 1000, + }); + defer limiter.deinit(); + + _ = limiter.checkRequest("1.1.1.1"); // allowed + _ = limiter.checkRequest("1.1.1.1"); // allowed + _ = limiter.checkRequest("1.1.1.1"); // allowed + _ = limiter.checkRequest("1.1.1.1"); // blocked + _ = limiter.checkRequest("2.2.2.2"); // allowed (different client) + + const stats = limiter.getStats(); + try testing.expectEqual(@as(u64, 5), stats.total_requests); + try testing.expectEqual(@as(u64, 1), stats.rate_limited); +} + +// ============================================================================ +// Cache Tests +// ============================================================================ + +test "DnsCache - stores and retrieves entries" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + const response = "test response data"; + dns_cache.put("test.com", types.QType.A, response, 300); + + const cached = try dns_cache.getCopy("test.com", types.QType.A); + defer if (cached) |c| allocator.free(c); + + try testing.expect(cached != null); + try testing.expectEqualStrings(response, cached.?); +} + +test "DnsCache - returns null for missing entries" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + const cached = try dns_cache.getCopy("nonexistent.com", types.QType.A); + try testing.expect(cached == null); +} + +test "DnsCache - separates entries by qtype" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + dns_cache.put("example.com", types.QType.A, "A record", 300); + dns_cache.put("example.com", types.QType.AAAA, "AAAA record", 300); + + const a_cached = try dns_cache.getCopy("example.com", types.QType.A); + defer if (a_cached) |c| allocator.free(c); + + const aaaa_cached = try dns_cache.getCopy("example.com", types.QType.AAAA); + defer if (aaaa_cached) |c| allocator.free(c); + + try testing.expectEqualStrings("A record", a_cached.?); + try testing.expectEqualStrings("AAAA record", aaaa_cached.?); +} + +test "DnsCache - respects max entries limit" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.initWithConfig(allocator, 3, 60, 86400); + defer dns_cache.deinit(); + + // Add 4 entries, should evict oldest + dns_cache.put("one.com", types.QType.A, "1", 300); + dns_cache.put("two.com", types.QType.A, "2", 300); + dns_cache.put("three.com", types.QType.A, "3", 300); + dns_cache.put("four.com", types.QType.A, "4", 300); + + const stats = dns_cache.getStats(); + try testing.expect(stats.entry_count <= 3); +} + +test "DnsCache - updates existing entries" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + dns_cache.put("update.com", types.QType.A, "original", 300); + dns_cache.put("update.com", types.QType.A, "updated", 300); + + const cached = try dns_cache.getCopy("update.com", types.QType.A); + defer if (cached) |c| allocator.free(c); + + try testing.expectEqualStrings("updated", cached.?); +} + +test "DnsCache - does not cache TTL=0 responses" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + // Try to cache with TTL=0 - should NOT be cached per RFC 2308 + dns_cache.put("nocache.com", types.QType.A, "should not cache", 0); + + const cached = try dns_cache.getCopy("nocache.com", types.QType.A); + try testing.expect(cached == null); +} + +// ============================================================================ +// Handler with Rate Limiting Integration +// ============================================================================ + +test "Handler - respects rate limiter" { + const allocator = testing.allocator; + + var limiter = rate_limiter.RateLimiter.initWithConfig(allocator, .{ + .max_qps = 2, + .window_ms = 1000, + }); + defer limiter.deinit(); + + var handler = handler_mod.Handler.init(allocator); + handler.setRateLimiter(&limiter); + + const query = createTestQuery(); + const addr = std.net.Address.initIp4([4]u8{ 10, 0, 0, 1 }, 12345); + + // First 2 requests succeed (return SERVFAIL due to no upstream) + const r1 = handler.handle(&query, addr, allocator); + const r2 = handler.handle(&query, addr, allocator); + + if (r1) |r| allocator.free(r); + if (r2) |r| allocator.free(r); + + // 3rd request should be rate limited (REFUSED) + const r3 = handler.handle(&query, addr, allocator); + if (r3) |r| { + defer allocator.free(r); + // REFUSED = RCODE 5 + const rcode = r[3] & 0x0F; + try testing.expectEqual(@as(u8, 5), rcode); + } +} + +// ============================================================================ +// Blocklist Integration Tests +// ============================================================================ + +test "Handler with blocklist - blocks matching domains" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + + try blocklist.addBlockedDomain("ads.example.com", 0); + + var handler = handler_mod.Handler.init(allocator); + handler.setBlocklist(blocklist.toHandlerBlocklist()); + + // Query for blocked domain + var query_buf: [512]u8 = undefined; + const query_len = createQueryForDomain("ads.example.com", &query_buf); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(query_buf[0..query_len], addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Should be a valid response with blocked content (0.0.0.0 or NXDOMAIN) + try testing.expect(r.len >= types.DNS_HEADER_SIZE); + // QR bit should be set (response) + try testing.expect((r[2] & 0x80) != 0); + } +} + +test "Handler with blocklist - allows non-blocked domains" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + + try blocklist.addBlockedDomain("blocked.com", 0); + + // Create mock upstream that returns a valid response + const mock_response = [_]u8{ + 0x00, 0x02, 0x81, 0x80, // Header (response) + 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, + // Question: allowed.com A IN + 0x07, 'a', 'l', 'l', 'o', 'w', 'e', 'd', + 0x03, 'c', 'o', 'm', 0x00, + 0x00, 0x01, 0x00, 0x01, + // Answer + 0xC0, 0x0C, 0x00, 0x01, 0x00, 0x01, + 0x00, 0x00, 0x01, 0x2C, // TTL 300 + 0x00, 0x04, 0x08, 0x08, 0x08, 0x08, // IP 8.8.8.8 + }; + + var mock_upstream = MockUpstream.init(&mock_response, allocator); + + var handler = handler_mod.Handler.init(allocator); + handler.setBlocklist(blocklist.toHandlerBlocklist()); + handler.setUpstream(mock_upstream.toHandlerUpstream()); + + // Query for non-blocked domain + var query_buf: [512]u8 = undefined; + const query_len = createQueryForDomain("allowed.com", &query_buf); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(query_buf[0..query_len], addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Should forward to upstream and return its response + try testing.expect(mock_upstream.call_count == 1); + } +} + +test "Handler with blocklist - blocks subdomains" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + + // Block parent domain + try blocklist.addBlockedDomain("doubleclick.net", 0); + + var handler = handler_mod.Handler.init(allocator); + handler.setBlocklist(blocklist.toHandlerBlocklist()); + + // Query for subdomain - should also be blocked + var query_buf: [512]u8 = undefined; + const query_len = createQueryForDomain("ads.doubleclick.net", &query_buf); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(query_buf[0..query_len], addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Should be blocked + try testing.expect(r.len >= types.DNS_HEADER_SIZE); + } +} + +test "Blocklist - allow rules override block rules" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + + try blocklist.addBlockedDomain("example.com", 0); + try blocklist.addAllowRule("allowed.example.com", 0); + + // Subdomain blocked + try testing.expect(blocklist.isBlocked("blocked.example.com", 0)); + // But allowed.example.com is explicitly allowed + try testing.expect(!blocklist.isBlocked("allowed.example.com", 0)); +} + +test "Blocklist - group isolation" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + + // Block domain only for group 1 + try blocklist.addBlockedDomain("group1only.com", 1); + + // Group 1 sees it blocked + try testing.expect(blocklist.isBlocked("group1only.com", 1)); + // Group 0 and 2 do not see it blocked + try testing.expect(!blocklist.isBlocked("group1only.com", 0)); + try testing.expect(!blocklist.isBlocked("group1only.com", 2)); +} + +// ============================================================================ +// Compression Loop Detection Tests +// ============================================================================ + +test "Name parsing - detects compression loop" { + const allocator = testing.allocator; + + // Packet with self-referencing compression pointer + const bad_packet = [_]u8{ + 0x00, 0x01, 0x01, 0x00, + 0x00, 0x01, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, + // Name with compression pointer pointing to itself (offset 12) + 0xC0, 0x0C, + 0x00, 0x01, 0x00, 0x01, + }; + + // Parsing should detect the loop and fail + const result = Name.parse(bad_packet[12..], &bad_packet, allocator); + try testing.expectError(error.CompressionLoop, result); +} + +test "Name parsing - detects indirect compression loop" { + const allocator = testing.allocator; + + // Packet where pointer A -> pointer B -> pointer A + const bad_packet = [_]u8{ + 0x00, 0x01, 0x01, 0x00, + 0x00, 0x01, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, + // Offset 12: pointer to offset 14 + 0xC0, 0x0E, + // Offset 14: pointer to offset 12 + 0xC0, 0x0C, + 0x00, 0x01, 0x00, 0x01, + }; + + const result = Name.parse(bad_packet[12..], &bad_packet, allocator); + try testing.expectError(error.CompressionLoop, result); +} + +// ============================================================================ +// Malformed Packet Handling Tests +// ============================================================================ + +test "Handler - handles truncated packet" { + const allocator = testing.allocator; + var handler = handler_mod.Handler.init(allocator); + + // Only 2 bytes - too short for DNS header + const truncated = [_]u8{ 0x00, 0x01 }; + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&truncated, addr, allocator); + + if (response) |r| { + defer allocator.free(r); + // Should return FORMERR + const rcode = r[3] & 0x0F; + try testing.expectEqual(@as(u8, 1), rcode); // FORMERR + } +} + +test "Handler - handles empty question section" { + const allocator = testing.allocator; + var handler = handler_mod.Handler.init(allocator); + + // Valid header but QDCOUNT = 0 + const no_question = [_]u8{ + 0x00, 0x01, 0x01, 0x00, + 0x00, 0x00, // QDCOUNT = 0 + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + }; + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&no_question, addr, allocator); + + if (response) |r| { + defer allocator.free(r); + const rcode = r[3] & 0x0F; + try testing.expectEqual(@as(u8, 1), rcode); // FORMERR + } +} + +// ============================================================================ +// Full Query Flow Integration Tests +// ============================================================================ + +test "Full flow - query hits cache" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + var mock_upstream = MockUpstream.init(null, allocator); + + var handler = handler_mod.Handler.init(allocator); + handler.setCache(dns_cache.toHandlerCache()); + handler.setUpstream(mock_upstream.toHandlerUpstream()); + + // Pre-populate cache + const cached_response = [_]u8{ + 0x00, 0x01, 0x81, 0x80, + 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', 0x00, + 0x00, 0x01, 0x00, 0x01, + 0xC0, 0x0C, 0x00, 0x01, 0x00, 0x01, + 0x00, 0x00, 0x01, 0x2C, + 0x00, 0x04, 0x01, 0x02, 0x03, 0x04, + }; + dns_cache.put("example.com", types.QType.A, &cached_response, 300); + + const query = createTestQuery(); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&query, addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Upstream should NOT have been called + try testing.expectEqual(@as(usize, 0), mock_upstream.call_count); + } +} + +test "Full flow - cache miss goes to upstream" { + const allocator = testing.allocator; + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + const upstream_response = [_]u8{ + 0x00, 0x01, 0x81, 0x80, + 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', 0x00, + 0x00, 0x01, 0x00, 0x01, + 0xC0, 0x0C, 0x00, 0x01, 0x00, 0x01, + 0x00, 0x00, 0x01, 0x2C, + 0x00, 0x04, 0x08, 0x08, 0x08, 0x08, + }; + + var mock_upstream = MockUpstream.init(&upstream_response, allocator); + + var handler = handler_mod.Handler.init(allocator); + handler.setCache(dns_cache.toHandlerCache()); + handler.setUpstream(mock_upstream.toHandlerUpstream()); + + const query = createTestQuery(); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(&query, addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Upstream should have been called + try testing.expectEqual(@as(usize, 1), mock_upstream.call_count); + + // Result should now be cached + const cached = try dns_cache.getCopy("example.com", types.QType.A); + try testing.expect(cached != null); + if (cached) |c| allocator.free(c); + } +} + +test "Full flow - blocklist takes precedence over cache" { + const allocator = testing.allocator; + + var blocklist = blocklist_mod.Blocklist.init(allocator); + defer blocklist.deinit(); + try blocklist.addBlockedDomain("blocked.com", 0); + + var dns_cache = cache.DnsCache.init(allocator); + defer dns_cache.deinit(); + + // Pre-populate cache with a response for blocked.com + const cached_response = [_]u8{ + 0x00, 0x01, 0x81, 0x80, + 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, + 0x07, 'b', 'l', 'o', 'c', 'k', 'e', 'd', + 0x03, 'c', 'o', 'm', 0x00, + 0x00, 0x01, 0x00, 0x01, + 0xC0, 0x0C, 0x00, 0x01, 0x00, 0x01, + 0x00, 0x00, 0x01, 0x2C, + 0x00, 0x04, 0x08, 0x08, 0x08, 0x08, // 8.8.8.8 + }; + dns_cache.put("blocked.com", types.QType.A, &cached_response, 300); + + var handler = handler_mod.Handler.init(allocator); + handler.setBlocklist(blocklist.toHandlerBlocklist()); + handler.setCache(dns_cache.toHandlerCache()); + + var query_buf: [512]u8 = undefined; + const query_len = createQueryForDomain("blocked.com", &query_buf); + const addr = std.net.Address.initIp4([4]u8{ 127, 0, 0, 1 }, 12345); + + const response = handler.handle(query_buf[0..query_len], addr, allocator); + try testing.expect(response != null); + + if (response) |r| { + defer allocator.free(r); + // Parse response to check answer + var pkt = packet.Packet.parse(r, allocator) catch { + try testing.expect(false); + return; + }; + defer pkt.deinit(); + + // Should be blocked (0.0.0.0), not cached (8.8.8.8) + if (pkt.answers.len > 0) { + const ip = pkt.answers[0].getA(); + if (ip) |addr_bytes| { + try testing.expectEqual([4]u8{ 0, 0, 0, 0 }, addr_bytes); + } + } + } +} + +// ============================================================================ +// Packet Encoding/Decoding Roundtrip Tests +// ============================================================================ + +test "Packet roundtrip - query" { + const allocator = testing.allocator; + + const original_query = createTestQuery(); + var pkt = try packet.Packet.parse(&original_query, allocator); + defer pkt.deinit(); + + var buf: [512]u8 = undefined; + const encoded_len = try pkt.encode(&buf); + + var decoded = try packet.Packet.parse(buf[0..encoded_len], allocator); + defer decoded.deinit(); + + try testing.expectEqual(pkt.header.id, decoded.header.id); + try testing.expectEqual(pkt.header.qr, decoded.header.qr); + try testing.expectEqual(pkt.questions.len, decoded.questions.len); +} + +test "Packet - createBlockedResponse has 0.0.0.0" { + const allocator = testing.allocator; + + const query = createTestQuery(); + var query_pkt = try packet.Packet.parse(&query, allocator); + defer query_pkt.deinit(); + + var response = try packet.Packet.createBlockedResponse(&query_pkt, allocator); + defer response.deinit(); + + try testing.expect(response.header.qr); // Is response + try testing.expectEqual(@as(usize, 1), response.answers.len); + + const ip = response.answers[0].getA(); + try testing.expect(ip != null); + try testing.expectEqual([4]u8{ 0, 0, 0, 0 }, ip.?); +} + +test "Packet - createNxdomainResponse has NXDOMAIN RCODE" { + const allocator = testing.allocator; + + const query = createTestQuery(); + var query_pkt = try packet.Packet.parse(&query, allocator); + defer query_pkt.deinit(); + + var response = try packet.Packet.createNxdomainResponse(&query_pkt, allocator); + defer response.deinit(); + + try testing.expect(response.header.qr); + try testing.expectEqual(types.RCode.NXDomain, response.header.rcode); + try testing.expectEqual(@as(usize, 0), response.answers.len); +} + +// ============================================================================ +// DoH/DoT URL Parsing Tests +// ============================================================================ + +const dot = @import("dot"); +const doh = @import("doh"); + +test "DoT URL parsing - valid URLs" { + const allocator = testing.allocator; + + // Standard TLS URL + const client1 = try dot.DotClient.fromUrl("tls://cloudflare-dns.com", allocator); + try testing.expectEqualStrings("cloudflare-dns.com", client1.host); + try testing.expectEqual(@as(u16, 853), client1.port); + + // With explicit port + const client2 = try dot.DotClient.fromUrl("tls://1.1.1.1:853", allocator); + try testing.expectEqualStrings("1.1.1.1", client2.host); + try testing.expectEqual(@as(u16, 853), client2.port); + + // Custom port + const client3 = try dot.DotClient.fromUrl("tls://dns.google:8853", allocator); + try testing.expectEqualStrings("dns.google", client3.host); + try testing.expectEqual(@as(u16, 8853), client3.port); +} + +test "DoT URL parsing - invalid URLs" { + const allocator = testing.allocator; + + // Wrong scheme + try testing.expectError(error.InvalidHost, dot.DotClient.fromUrl("https://example.com", allocator)); + try testing.expectError(error.InvalidHost, dot.DotClient.fromUrl("not-a-url", allocator)); +} + +test "DoH URL parsing - valid URLs" { + const allocator = testing.allocator; + + // Standard DoH URL + const client1 = try doh.DohClient.init("https://cloudflare-dns.com/dns-query", allocator); + try testing.expectEqualStrings("cloudflare-dns.com", client1.host); + try testing.expectEqualStrings("/dns-query", client1.path); + try testing.expectEqual(@as(u16, 443), client1.port); + + // With custom port + const client2 = try doh.DohClient.init("https://dns.quad9.net:8443/dns-query", allocator); + try testing.expectEqualStrings("dns.quad9.net", client2.host); + try testing.expectEqual(@as(u16, 8443), client2.port); + + // No path (defaults to /dns-query) + const client3 = try doh.DohClient.init("https://example.com", allocator); + try testing.expectEqualStrings("example.com", client3.host); + try testing.expectEqualStrings("/dns-query", client3.path); +} + +test "DoH URL parsing - invalid URLs" { + const allocator = testing.allocator; + + // Wrong scheme + try testing.expectError(error.InvalidUrl, doh.DohClient.init("http://example.com/dns-query", allocator)); + try testing.expectError(error.InvalidUrl, doh.DohClient.init("not-a-url", allocator)); +} + +// ============================================================================ +// Edge Case Tests +// ============================================================================ + +test "Handler - IPv6 client address" { + const allocator = testing.allocator; + var handler = handler_mod.Handler.init(allocator); + + const query = createTestQuery(); + + // Create IPv6 address + var addr: std.net.Address = undefined; + addr.in6 = std.net.Ip6Address.init([16]u8{ + 0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, + }, 12345); + addr.any.family = std.posix.AF.INET6; + + // Should handle IPv6 without crashing + const response = handler.handle(&query, addr, allocator); + + if (response) |r| { + defer allocator.free(r); + try testing.expect(r.len >= types.DNS_HEADER_SIZE); + } +} + +test "Cache - TTL clamping" { + const allocator = testing.allocator; + + // Configure with min_ttl=60, max_ttl=3600 + var dns_cache = cache.DnsCache.initWithConfig(allocator, 100, 60, 3600); + defer dns_cache.deinit(); + + // Put with TTL below min - should be clamped to 60 + dns_cache.put("test1.com", types.QType.A, "response", 10); + const cached1 = try dns_cache.getCopy("test1.com", types.QType.A); + try testing.expect(cached1 != null); + if (cached1) |c| allocator.free(c); + + // Put with TTL above max - should be clamped to 3600 + dns_cache.put("test2.com", types.QType.A, "response", 100000); + const cached2 = try dns_cache.getCopy("test2.com", types.QType.A); + try testing.expect(cached2 != null); + if (cached2) |c| allocator.free(c); +} +