From a246a67d9312c0fce842d9ce5a734c651cdc6b4b Mon Sep 17 00:00:00 2001 From: Muki Kiboigo Date: Thu, 10 Sep 2026 11:42:06 -0700 Subject: [PATCH] add CorsStore --- src/Metrics.zig | 2 +- src/network/CorsGate.zig | 35 ++++ src/network/CorsStore.zig | 348 +++++++++++++++++++++++++++++++++++++ src/network/HttpClient.zig | 5 +- src/network/Network.zig | 5 + 5 files changed, 393 insertions(+), 2 deletions(-) create mode 100644 src/network/CorsStore.zig diff --git a/src/Metrics.zig b/src/Metrics.zig index c7dabee96..c92278d6f 100644 --- a/src/Metrics.zig +++ b/src/Metrics.zig @@ -93,7 +93,7 @@ http_navigation_delay_ms: Histogram(&.{ robots_status: CounterEnum("category", @import("network/http.zig").StatusCategory) = .{}, robots_access: CounterEnum("result", enum { allow, deny }) = .{}, robots_evictions: Counter = .{}, -cors_check: CounterEnum("result", enum { same_origin, no_cors, simple, preflight }) = .{}, +cors_check: CounterEnum("result", enum { same_origin, no_cors, simple, preflight, cached }) = .{}, cors_preflight: CounterEnum("result", enum { allowed, blocked }) = .{}, cors_response: CounterEnum("result", enum { allowed, blocked }) = .{}, adblock_verdicts: CounterEnum("verdict", @import("network/adblock/AdBlocker.zig").Verdict) = .{}, diff --git a/src/network/CorsGate.zig b/src/network/CorsGate.zig index 7302ab3b9..ba9310eb8 100644 --- a/src/network/CorsGate.zig +++ b/src/network/CorsGate.zig @@ -25,11 +25,15 @@ const http = @import("http.zig"); const Transfer = @import("HttpClient.zig").Transfer; const SingleFlight = @import("SingleFlight.zig"); const HttpClient = @import("HttpClient.zig"); +const Network = @import("Network.zig"); + +const CorsStore = @import("CorsStore.zig"); const log = lp.log; const CorsGate = @This(); +network: *Network, single_flight: SingleFlight, // CORS Request Headers @@ -212,6 +216,22 @@ pub fn check(self: *CorsGate, transfer: *Transfer) !Result { return .allowed; } + if (try URL.getOrigin(transfer.arena.allocator(), req.url)) |target| { + if (self.network.cors_store.get(.{ .origin = origin, .target = target })) |cached| { + const authored = try collectAuthoredHeaders(transfer, transfer.arena.allocator()); + if (CorsStore.covers(cached, req.method, req.credentials_mode == .include, authored.items)) { + log.debug(.cors, "cross origin", .{ + .url = req.url, + .origin = origin, + .preflight = false, + .cached = true, + }); + lp.metrics.cors_check.incr(.cached); + return .allowed; + } + } + } + log.debug(.cors, "cross origin", .{ .url = req.url, .origin = origin, @@ -223,6 +243,21 @@ pub fn check(self: *CorsGate, transfer: *Transfer) !Result { return .pending; } +fn collectAuthoredHeaders(transfer: *Transfer, allocator: std.mem.Allocator) !std.ArrayList([]const u8) { + var header_names: std.ArrayList([]const u8) = .empty; + for (transfer.req_headers.items) |hdr| { + if (hdr.source != .author) continue; + if (isSafelistedHeader(hdr.name, hdr.value)) continue; + try header_names.append(allocator, try std.ascii.allocLowerString(allocator, hdr.name)); + } + std.mem.sort([]const u8, header_names.items, {}, struct { + fn lessThan(_: void, a: []const u8, b: []const u8) bool { + return std.mem.lessThan(u8, a, b); + } + }.lessThan); + return header_names; +} + const CorsKey = struct { url: []const u8, origin: []const u8, diff --git a/src/network/CorsStore.zig b/src/network/CorsStore.zig new file mode 100644 index 000000000..696c5676a --- /dev/null +++ b/src/network/CorsStore.zig @@ -0,0 +1,348 @@ +// Copyright (C) 2023-2026 Lightpanda (Selecy SAS) +// +// Francis Bouvier +// Pierre Tachoire +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as +// published by the Free Software Foundation, either version 3 of the +// License, or (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +const std = @import("std"); +const lp = @import("lightpanda"); + +const http = @import("http.zig"); + +const CorsStore = @This(); + +const Key = struct { + origin: []const u8, + target: []const u8, + + fn dupe(self: Key, allocator: std.mem.Allocator) !Key { + return .{ + .origin = try allocator.dupe(u8, self.origin), + .target = try allocator.dupe(u8, self.target), + }; + } + + fn deinit(self: Key, allocator: std.mem.Allocator) void { + allocator.free(self.origin); + allocator.free(self.target); + } +}; + +const KeyContext = struct { + pub fn hash(_: KeyContext, key: Key) u64 { + var hasher = std.hash.Wyhash.init(0); + hasher.update(key.origin); + hasher.update(&.{0}); + hasher.update(key.target); + return hasher.final(); + } + + pub fn eql(_: KeyContext, a: Key, b: Key) bool { + return std.ascii.eqlIgnoreCase(a.origin, b.origin) and std.ascii.eqlIgnoreCase(a.target, b.target); + } +}; + +const Entry = struct { + methods_wildcard: bool, + methods: std.EnumSet(http.Method), + + headers_wildcard: bool, + headers: []const []const u8, + + expires_at: u64, + + credentials: bool, + + fn unionHeaders( + allocator: std.mem.Allocator, + a: []const []const u8, + b: []const []const u8, + ) ![]const []const u8 { + var list: std.ArrayList([]const u8) = .empty; + errdefer { + for (list.items) |s| allocator.free(s); + list.deinit(allocator); + } + + outerA: for (a) |s| { + for (list.items) |existing| { + if (std.ascii.eqlIgnoreCase(existing, s)) continue :outerA; + } + try list.append(allocator, try allocator.dupe(u8, s)); + } + outerB: for (b) |s| { + for (list.items) |existing| { + if (std.ascii.eqlIgnoreCase(existing, s)) continue :outerB; + } + try list.append(allocator, try allocator.dupe(u8, s)); + } + + return list.toOwnedSlice(allocator); + } + + fn merge(self: Entry, allocator: std.mem.Allocator, new: Entry) !Entry { + return .{ + .credentials = self.credentials or new.credentials, + .methods_wildcard = self.methods_wildcard or new.methods_wildcard, + .methods = self.methods.unionWith(new.methods), + .headers_wildcard = self.headers_wildcard or new.headers_wildcard, + .headers = try unionHeaders(allocator, self.headers, new.headers), + .expires_at = @max(self.expires_at, new.expires_at), + }; + } + + fn dupe(self: Entry, allocator: std.mem.Allocator) !Entry { + var new_headers: std.ArrayList([]const u8) = try .initCapacity(allocator, self.headers.len); + errdefer { + for (new_headers.items) |hdr| allocator.free(hdr); + new_headers.deinit(allocator); + } + + for (self.headers) |hdr| { + new_headers.appendAssumeCapacity(try allocator.dupe(u8, hdr)); + } + + return .{ + .methods_wildcard = self.methods_wildcard, + .methods = self.methods, + .headers_wildcard = self.headers_wildcard, + .headers = new_headers.items, + .expires_at = self.expires_at, + .credentials = self.credentials, + }; + } + + fn deinit(self: Entry, allocator: std.mem.Allocator) void { + for (self.headers) |h| allocator.free(h); + allocator.free(self.headers); + } +}; + +const Map = std.HashMapUnmanaged(Key, Entry, KeyContext, std.hash_map.default_max_load_percentage); + +allocator: std.mem.Allocator, +map: Map = .empty, +mutex: std.Io.Mutex = .init, + +pub fn init(allocator: std.mem.Allocator) CorsStore { + return .{ .allocator = allocator }; +} + +pub fn deinit(self: *CorsStore) void { + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + var iter = self.map.iterator(); + while (iter.next()) |entry| { + entry.key_ptr.deinit(self.allocator); + entry.value_ptr.deinit(self.allocator); + } + + self.map.deinit(self.allocator); +} + +pub fn get(self: *CorsStore, key: Key) ?Entry { + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + const entry = self.map.get(key) orelse return null; + + if (entry.expires_at <= lp.datetime.timestamp(.real)) { + const kv = self.map.fetchRemove(key).?; + kv.key.deinit(self.allocator); + kv.value.deinit(self.allocator); + return null; + } + + return entry; +} + +/// Insert or merge a CORS grant for (origin, target). `entry` is not +/// consumed: `put` copies whatever it needs (via `dupe`/`merge`, which +/// always allocate their own copies) and never takes ownership of +/// `entry.headers` or its contents. +/// +/// Callers remain responsible for +/// freeing `entry.headers` after this call, on both the insert and +/// the merge path. +pub fn put(self: *CorsStore, key: Key, entry: Entry) !void { + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + const gop = try self.map.getOrPut(self.allocator, key); + if (!gop.found_existing) { + errdefer _ = self.map.remove(key); + gop.key_ptr.* = try key.dupe(self.allocator); + gop.value_ptr.* = try entry.dupe(self.allocator); + return; + } + + const old = gop.value_ptr.*; + const merged = old.merge(self.allocator, entry) catch |err| { + return err; + }; + old.deinit(self.allocator); + gop.value_ptr.* = merged; +} + +pub fn covers( + entry: Entry, + method: http.Method, + wants_credentials: bool, + authored_headers: []const []const u8, +) bool { + if (wants_credentials and !entry.credentials) { + return false; + } + + if (!entry.methods_wildcard and !entry.methods.contains(method)) { + return false; + } + + for (authored_headers) |name| { + const is_authorization = std.ascii.eqlIgnoreCase(name, "authorization"); + if (entry.headers_wildcard and !is_authorization) continue; + + var found = false; + for (entry.headers) |allowed| { + if (std.ascii.eqlIgnoreCase(allowed, name)) { + found = true; + break; + } + } + if (!found) return false; + } + + return true; +} + +const testing = @import("../testing.zig"); + +fn freeHeaders(allocator: std.mem.Allocator, headers: []const []const u8) void { + for (headers) |h| allocator.free(h); + allocator.free(headers); +} + +test "CorsStore: put then get, miss on different origin/target" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator); + defer store.deinit(); + + const headers = try allocator.alloc([]const u8, 1); + headers[0] = try allocator.dupe(u8, "x-custom"); + defer freeHeaders(allocator, headers); + + try store.put(.{ .origin = "https://a.example", .target = "https://api.example" }, .{ + .credentials = false, + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.POST), + .headers_wildcard = false, + .headers = headers, + .expires_at = lp.datetime.timestamp(.real) + 60_000, + }); + + // store.get returns the store's own copy — not caller-owned, don't free it. + const hit = store.get(.{ .origin = "https://a.example", .target = "https://api.example" }).?; + try testing.expect(hit.methods.contains(.POST)); + try testing.expect(!hit.methods.contains(.GET)); + + try testing.expectEqual(null, store.get(.{ .origin = "https://b.example", .target = "https://api.example" })); + try testing.expectEqual(null, store.get(.{ .origin = "https://a.example", .target = "https://other.example" })); +} + +test "CorsStore: expired entries are evicted on get" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator); + defer store.deinit(); + + try store.put(.{ .origin = "https://a.example", .target = "https://api.example" }, .{ + .credentials = false, + .methods_wildcard = true, + .methods = .initEmpty(), + .headers_wildcard = true, + .headers = &.{}, // empty slice, nothing to free + .expires_at = lp.datetime.timestamp(.real) - 1, + }); + + try testing.expectEqual(null, store.get(.{ .origin = "https://a.example", .target = "https://api.example" })); + try testing.expectEqual(0, store.map.count()); +} + +test "CorsStore: put merges into existing entry rather than clobbering" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator); + defer store.deinit(); + + const key = Key{ .origin = "https://a.example", .target = "https://api.example" }; + + const h1 = try allocator.alloc([]const u8, 1); + h1[0] = try allocator.dupe(u8, "x-one"); + try store.put(key, .{ + .credentials = false, + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.POST), + .headers_wildcard = false, + .headers = h1, + .expires_at = lp.datetime.timestamp(.real) + 60_000, + }); + freeHeaders(allocator, h1); + + const h2 = try allocator.alloc([]const u8, 1); + h2[0] = try allocator.dupe(u8, "x-two"); + try store.put(key, .{ + .credentials = false, + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.PUT), + .headers_wildcard = false, + .headers = h2, + .expires_at = lp.datetime.timestamp(.real) + 60_000, + }); + freeHeaders(allocator, h2); + + const merged = store.get(key).?; + try testing.expect(merged.methods.contains(.POST)); + try testing.expect(merged.methods.contains(.PUT)); + try testing.expectEqual(2, merged.headers.len); + + try testing.expect(CorsStore.covers(merged, .POST, false, &.{"x-one"})); + try testing.expect(CorsStore.covers(merged, .PUT, false, &.{"x-two"})); + try testing.expect(!CorsStore.covers(merged, .DELETE, false, &.{})); +} + +test "CorsStore: covers rejects credentialed request against uncredentialed wildcard" { + const entry = CorsStore.Entry{ + .credentials = false, + .methods_wildcard = true, + .methods = .initEmpty(), + .headers_wildcard = true, + .headers = &.{}, + .expires_at = std.math.maxInt(i64), + }; + try testing.expect(!CorsStore.covers(entry, .GET, true, &.{})); + try testing.expect(CorsStore.covers(entry, .GET, false, &.{})); +} + +test "CorsStore: covers never lets a wildcard cover Authorization" { + const entry = CorsStore.Entry{ + .credentials = false, + .methods_wildcard = true, + .methods = .initEmpty(), + .headers_wildcard = true, + .headers = &.{}, + .expires_at = std.math.maxInt(i64), + }; + try testing.expect(!CorsStore.covers(entry, .GET, false, &.{"authorization"})); + try testing.expect(CorsStore.covers(entry, .GET, false, &.{"x-anything"})); +} diff --git a/src/network/HttpClient.zig b/src/network/HttpClient.zig index 45e155065..a1ee7bcdc 100644 --- a/src/network/HttpClient.zig +++ b/src/network/HttpClient.zig @@ -243,7 +243,10 @@ pub fn init(self: *Client, app: *lp.App) !void { .network = network, .single_flight = .init(allocator), }, - .cors = .{ .single_flight = .init(allocator) }, + .cors = .{ + .network = network, + .single_flight = .init(allocator), + }, .url_blocklist = url_blocklist, .arena_pool = &app.arena_pool, }; diff --git a/src/network/Network.zig b/src/network/Network.zig index f6593c008..7a78d5a24 100644 --- a/src/network/Network.zig +++ b/src/network/Network.zig @@ -27,6 +27,7 @@ const libcurl = @import("../sys/libcurl.zig"); const http = @import("http.zig"); const IpFilter = @import("IpFilter.zig"); const RobotStore = @import("Robots.zig").RobotStore; +const CorsStore = @import("CorsStore.zig"); const WebBotAuth = @import("WebBotAuth.zig"); const RateLimiter = @import("RateLimiter.zig"); const Certificates = @import("Certificates.zig"); @@ -43,6 +44,7 @@ cache: Cache, allocator: Allocator, config: *const Config, robot_store: RobotStore, +cors_store: CorsStore, web_bot_auth: ?WebBotAuth, rate_limiter: ?RateLimiter, certificates: Certificates, @@ -122,6 +124,7 @@ pub fn init(app: *App) !Network { .cache = cache, .robot_store = RobotStore.init(allocator, config.robotStoreEntryLimit()), + .cors_store = CorsStore.init(allocator), .web_bot_auth = web_bot_auth, .rate_limiter = if (config.httpNavDelay()) |ms| RateLimiter.init(allocator, ms, config.httpNavBurst()) else null, .adblocker = adblocker, @@ -144,6 +147,8 @@ pub fn deinit(self: *Network) void { self.ws_pool.deinit(self.allocator); self.robot_store.deinit(); + self.cors_store.deinit(); + if (self.rate_limiter) |*rl| { rl.deinit(); }