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;