From fff0957ba3a3355b6a81e6fbee58dc73c43f527e Mon Sep 17 00:00:00 2001 From: zzstoatzz Date: Sun, 1 Mar 2026 20:43:52 -0600 Subject: [PATCH] feat: add per-host rate limiting and Server header - per-host sliding window rate limits (50/s, 2500/hr, 20000/day baseline; 5000/s, 50M/hr, 500M/day for trusted *.host.bsky.network hosts) matches Go relay's slidingwindow approach - Server: zlay (atproto-relay) header on all HTTP responses, enabling relay loop detection by other relays - new relay_rate_limited_total prometheus counter Co-Authored-By: Claude Opus 4.6 --- src/broadcaster.zig | 5 ++ src/main.zig | 20 +++++-- src/subscriber.zig | 128 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 149 insertions(+), 4 deletions(-) diff --git a/src/broadcaster.zig b/src/broadcaster.zig index d0dbda3..d1ac1b3 100644 --- a/src/broadcaster.zig +++ b/src/broadcaster.zig @@ -27,6 +27,7 @@ pub const Stats = struct { validated: std.atomic.Value(u64) = .{ .raw = 0 }, failed: std.atomic.Value(u64) = .{ .raw = 0 }, skipped: std.atomic.Value(u64) = .{ .raw = 0 }, + rate_limited: std.atomic.Value(u64) = .{ .raw = 0 }, decode_errors: std.atomic.Value(u64) = .{ .raw = 0 }, cache_hits: std.atomic.Value(u64) = .{ .raw = 0 }, cache_misses: std.atomic.Value(u64) = .{ .raw = 0 }, @@ -531,6 +532,9 @@ pub fn formatPrometheusMetrics(stats: *const Stats, buf: []u8) []const u8 { \\relay_validation_total{{result="failed"}} {d} \\relay_validation_total{{result="skipped"}} {d} \\ + \\# TYPE relay_rate_limited_total counter + \\relay_rate_limited_total {d} + \\ \\# TYPE relay_decode_errors_total counter \\relay_decode_errors_total {d} \\ @@ -560,6 +564,7 @@ pub fn formatPrometheusMetrics(stats: *const Stats, buf: []u8) []const u8 { stats.validated.load(.acquire), stats.failed.load(.acquire), stats.skipped.load(.acquire), + stats.rate_limited.load(.acquire), stats.decode_errors.load(.acquire), stats.cache_hits.load(.acquire), stats.cache_misses.load(.acquire), diff --git a/src/main.zig b/src/main.zig index e59baff..d29dde4 100644 --- a/src/main.zig +++ b/src/main.zig @@ -247,7 +247,10 @@ fn handleGet(request: *http.Server.Request, path: []const u8, query: []const u8, } else if (std.mem.eql(u8, path, "/metrics")) { var metrics_buf: [4096]u8 = undefined; const body = broadcaster.formatPrometheusMetrics(stats, &metrics_buf); - request.respond(body, .{ .status = .ok, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "text/plain; version=0.0.4; charset=utf-8" }} }) catch {}; + request.respond(body, .{ .status = .ok, .keep_alive = false, .extra_headers = &.{ + .{ .name = "content-type", .value = "text/plain; version=0.0.4; charset=utf-8" }, + .{ .name = "server", .value = "zlay (atproto-relay)" }, + } }) catch {}; } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.listRepos")) { handleListRepos(request, query, persist); } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.getRepoStatus")) { @@ -279,7 +282,10 @@ fn handleGet(request: *http.Server.Request, path: []const u8, query: []const u8, \\ \\Z \\ - , .{ .status = .ok, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "image/svg+xml" }} }) catch {}; + , .{ .status = .ok, .keep_alive = false, .extra_headers = &.{ + .{ .name = "content-type", .value = "image/svg+xml" }, + .{ .name = "server", .value = "zlay (atproto-relay)" }, + } }) catch {}; } else { respondText(request, .not_found, "not found"); } @@ -869,11 +875,17 @@ fn hexVal(c: u8) ?u4 { // --- response helpers --- fn respondJson(request: *http.Server.Request, status: http.Status, body: []const u8) void { - request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "application/json" }} }) catch {}; + request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{ + .{ .name = "content-type", .value = "application/json" }, + .{ .name = "server", .value = "zlay (atproto-relay)" }, + } }) catch {}; } fn respondText(request: *http.Server.Request, status: http.Status, body: []const u8) void { - request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "text/plain" }} }) catch {}; + request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{ + .{ .name = "content-type", .value = "text/plain" }, + .{ .name = "server", .value = "zlay (atproto-relay)" }, + } }) catch {}; } fn parseEnvInt(comptime T: type, key: []const u8, default: T) T { diff --git a/src/subscriber.zig b/src/subscriber.zig index 677bfd7..6b4a916 100644 --- a/src/subscriber.zig +++ b/src/subscriber.zig @@ -20,12 +20,83 @@ const log = std.log.scoped(.relay); const max_consecutive_failures = 15; const cursor_flush_interval_sec = 4; // flush cursor to DB every N seconds (Go relay: 4s) +// per-host rate limits (Go relay: 50/s baseline, 2500/hr, 20000/day for public hosts) +const default_per_second_limit: u64 = 50; +const default_per_hour_limit: u64 = 2500; +const default_per_day_limit: u64 = 20_000; + +// trusted hosts get much higher limits (Go relay: 5000/s, 50M/hr, 500M/day) +const trusted_per_second_limit: u64 = 5_000; +const trusted_per_hour_limit: u64 = 50_000_000; +const trusted_per_day_limit: u64 = 500_000_000; + +// Go relay: TrustedDomains config — hosts matching these suffixes get trusted limits +const trusted_suffixes: []const []const u8 = &.{".host.bsky.network"}; + +fn isTrustedHost(hostname: []const u8) bool { + for (trusted_suffixes) |suffix| { + if (std.mem.endsWith(u8, hostname, suffix)) return true; + } + return false; +} + pub const Options = struct { hostname: []const u8 = "bsky.network", max_message_size: usize = 5 * 1024 * 1024, host_id: u64 = 0, }; +/// simple sliding window rate limiter — tracks event counts per second/hour/day. +/// Go relay uses github.com/RussellLuo/slidingwindow; this is a simpler fixed-window +/// approximation that resets counters at window boundaries. +const RateLimiter = struct { + // per-second + sec_count: u64 = 0, + sec_epoch: i64 = 0, + sec_limit: u64 = default_per_second_limit, + + // per-hour + hour_count: u64 = 0, + hour_epoch: i64 = 0, + hour_limit: u64 = default_per_hour_limit, + + // per-day + day_count: u64 = 0, + day_epoch: i64 = 0, + day_limit: u64 = default_per_day_limit, + + /// returns true if the event is allowed, false if rate-limited. + fn allow(self: *RateLimiter, now: i64) bool { + // per-second window + if (now != self.sec_epoch) { + self.sec_epoch = now; + self.sec_count = 0; + } + if (self.sec_count >= self.sec_limit) return false; + + // per-hour window + const hour = @divTrunc(now, 3600); + if (hour != self.hour_epoch) { + self.hour_epoch = hour; + self.hour_count = 0; + } + if (self.hour_count >= self.hour_limit) return false; + + // per-day window + const day = @divTrunc(now, 86400); + if (day != self.day_epoch) { + self.day_epoch = day; + self.day_count = 0; + } + if (self.day_count >= self.day_limit) return false; + + self.sec_count += 1; + self.hour_count += 1; + self.day_count += 1; + return true; + } +}; + pub const Subscriber = struct { allocator: Allocator, options: Options, @@ -36,6 +107,7 @@ pub const Subscriber = struct { shutdown: *std.atomic.Value(bool), last_upstream_seq: ?u64 = null, last_cursor_flush: i64 = 0, + rate_limiter: RateLimiter = .{}, // per-host shutdown (e.g. FutureCursor — stops only this subscriber) host_shutdown: std.atomic.Value(bool) = .{ .raw = false }, @@ -48,6 +120,7 @@ pub const Subscriber = struct { shutdown: *std.atomic.Value(bool), options: Options, ) Subscriber { + const trusted = isTrustedHost(options.hostname); return .{ .allocator = allocator, .options = options, @@ -55,6 +128,11 @@ pub const Subscriber = struct { .validator = val, .persist = persist, .shutdown = shutdown, + .rate_limiter = .{ + .sec_limit = if (trusted) trusted_per_second_limit else default_per_second_limit, + .hour_limit = if (trusted) trusted_per_hour_limit else default_per_hour_limit, + .day_limit = if (trusted) trusted_per_day_limit else default_per_day_limit, + }, }; } @@ -220,6 +298,7 @@ const FrameHandler = struct { _ = sub.bc.stats.frames_in.fetchAdd(1, .monotonic); // extract seq for cursor tracking (all event types have seq) + // must happen before rate limiting so we don't re-process dropped events on reconnect const upstream_seq = payload.getUint("seq"); if (upstream_seq) |s| { sub.last_upstream_seq = s; @@ -235,6 +314,13 @@ const FrameHandler = struct { } } + // per-host rate limiting (Go relay: slidingwindow per-second/hour/day) + // applied after cursor tracking so dropped events aren't re-processed on reconnect + if (!sub.rate_limiter.allow(std.time.timestamp())) { + _ = sub.bc.stats.rate_limited.fetchAdd(1, .monotonic); + return; + } + // route by frame type const is_commit = std.mem.eql(u8, frame_type, "#commit"); const is_account = std.mem.eql(u8, frame_type, "#account"); @@ -436,6 +522,48 @@ test "decode identity frame via SDK" { try std.testing.expectEqual(@as(i64, 99), p.getInt("seq").?); } +test "rate limiter enforces per-second limit" { + var rl: RateLimiter = .{}; + rl.sec_limit = 3; + rl.hour_limit = 1000; + rl.day_limit = 10000; + + const now: i64 = 1000000; + try std.testing.expect(rl.allow(now)); + try std.testing.expect(rl.allow(now)); + try std.testing.expect(rl.allow(now)); + // 4th should be rejected + try std.testing.expect(!rl.allow(now)); + + // next second resets + try std.testing.expect(rl.allow(now + 1)); +} + +test "rate limiter enforces per-hour limit" { + var rl: RateLimiter = .{}; + rl.sec_limit = 1000; // high per-second so it doesn't interfere + rl.hour_limit = 5; + rl.day_limit = 10000; + + const now: i64 = 3600 * 100; // some hour boundary + for (0..5) |_| { + try std.testing.expect(rl.allow(now)); + } + // 6th should be rejected (same hour) + try std.testing.expect(!rl.allow(now + 1)); + + // next hour resets + try std.testing.expect(rl.allow(now + 3600)); +} + +test "trusted host detection" { + try std.testing.expect(isTrustedHost("pds-123.host.bsky.network")); + try std.testing.expect(isTrustedHost("abc.host.bsky.network")); + try std.testing.expect(!isTrustedHost("bsky.network")); + try std.testing.expect(!isTrustedHost("evil.bsky.network")); + try std.testing.expect(!isTrustedHost("pds.example.com")); +} + test "error frame (op=-1) is detected" { const cbor = zat.cbor; -- 2.51.2