diff --git a/src/Config.zig b/src/Config.zig index 3562346c6..9dd641380 100644 --- a/src/Config.zig +++ b/src/Config.zig @@ -254,6 +254,7 @@ pub const ExperimentalFeatures = packed struct(u2) { const CommonOptions = .{ .{ .name = "obey_robots", .type = bool }, .{ .name = "robot_store_entry_limit", .type = ?u32, .default = 1000 }, + .{ .name = "cors_store_entry_limit", .type = ?u32, .default = 1000 }, .{ .name = "proxy_bearer_token", .type = ?[:0]const u8 }, .{ .name = "http_proxy", .type = ?[:0]const u8 }, .{ .name = "http_max_concurrent", .type = ?u8 }, @@ -570,6 +571,13 @@ pub fn robotStoreEntryLimit(self: *const Config) u32 { }; } +pub fn corsStoreEntryLimit(self: *const Config) u32 { + return switch (self.mode) { + inline .serve, .fetch, .mcp, .agent => |opts| opts.cors_store_entry_limit.?, + else => 1000, + }; +} + pub fn httpVersion(self: *const Config) HttpVersion { return switch (self.mode) { inline .serve, .fetch, .mcp, .agent => |opts| opts.http_version, 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/help.zon b/src/help.zon index 4fcf2426d..f48bc9091 100644 --- a/src/help.zon +++ b/src/help.zon @@ -396,6 +396,9 @@ \\ --cookie-jar \\ Path to a JSON file to save cookies to on exit (write-only). \\ Defaults to no cookie saving. + \\ --cors-store-entry-limit + \\ Maximum number of entries kept in the CorsStore. 0 means no limit. + \\ Defaults to 1000. \\ --experimental-features \\ Enable an experimental, unstable feature. Can be passed multiple times. \\ Behavior may change or be removed without notice. @@ -502,12 +505,12 @@ \\ --obey-robots \\ Fetches and obeys robots.txt of the target page. \\ Defaults to false. - \\ --robot-store-entry-limit - \\ Maximum number of entries kept in the RobotStore. 0 means no limit. - \\ Defaults to 1000. \\ --proxy-bearer-token \\ Token sent for bearer authentication with the proxy: \\ Proxy-Authorization: Bearer . + \\ --robot-store-entry-limit + \\ Maximum number of entries kept in the RobotStore. 0 means no limit. + \\ Defaults to 1000. \\ --timezone \\ Time zone used by Date and Intl, e.g. Europe/Paris or UTC. \\ Defaults to the host time zone. diff --git a/src/network/ClockCache.zig b/src/network/ClockCache.zig index 4eaa21372..4a6b2ad53 100644 --- a/src/network/ClockCache.zig +++ b/src/network/ClockCache.zig @@ -64,6 +64,15 @@ pub fn ClockCache(comptime V: type) type { return &entry.value; } + pub fn remove(self: *Self, key: []const u8) ?V { + const index = self.map.getIndex(key) orelse return null; + const owned_key = self.map.keys()[index]; + const value = self.map.values()[index].value; + self.map.swapRemoveAt(index); + self.allocator.free(owned_key); + return value; + } + pub fn insert(self: *Self, key: []const u8, value: V) !InsertResult { const gop = try self.map.getOrPut(self.allocator, key); if (gop.found_existing) return .exists; diff --git a/src/network/CorsGate.zig b/src/network/CorsGate.zig index 8efbe0f45..109264584 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 @@ -42,6 +46,7 @@ const ACCESS_CONTROL_ALLOW_ORIGIN = "access-control-allow-origin"; const ACCESS_CONTROL_ALLOW_METHODS = "access-control-allow-methods"; const ACCESS_CONTROL_ALLOW_HEADERS = "access-control-allow-headers"; const ACCESS_CONTROL_ALLOW_CREDENTIALS = "access-control-allow-credentials"; +const ACCESS_CONTROL_MAX_AGE = "access-control-max-age"; pub fn deinit(self: *CorsGate) void { self.single_flight.deinit(); @@ -72,7 +77,7 @@ fn flushPending(self: *CorsGate, key: []const u8, allowed: bool) void { } } -fn isSafelistedMethod(value: http.Method) bool { +pub fn isSafelistedMethod(value: http.Method) bool { return switch (value) { .GET, .HEAD, .POST => true, else => false, @@ -212,6 +217,28 @@ pub fn check(self: *CorsGate, transfer: *Transfer) !Result { return .allowed; } + const wants_credentials = req.credentials_mode == .include; + + const authored = try collectAuthoredHeaders(transfer, transfer.arena.allocator()); + + const covered = try self.network.cors_store.coversRequest( + transfer.arena.allocator(), + .{ .origin = origin, .target = req.url, .credentials = wants_credentials }, + req.method, + authored.items, + ); + + if (covered) { + 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, @@ -219,10 +246,25 @@ pub fn check(self: *CorsGate, transfer: *Transfer) !Result { }); lp.metrics.cors_check.incr(.preflight); - try self.fetchThenResume(transfer); + try self.fetchThenResume(transfer, authored.items); 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, @@ -264,6 +306,9 @@ const CorsPreflightContext = struct { wants_credentials: bool, allowed: bool = false, + acam: ?[]const u8 = null, + acah: ?[]const u8 = null, + acma: ?[]const u8 = null, fn validateHeaders( self: *CorsPreflightContext, @@ -355,6 +400,57 @@ const CorsPreflightContext = struct { return true; } + fn cacheGrant(self: *CorsPreflightContext, acam: ?[]const u8, acah: ?[]const u8, acma: ?[]const u8) !void { + if (self.url.len == 0) return; + + const max_age_s: u64 = blk: { + const v = acma orelse break :blk 5; + if (v.len == 0) break :blk 5; + break :blk std.fmt.parseUnsigned(u64, v, 10) catch return; + }; + if (max_age_s == 0) return; + + const capped_s: u64 = @min(max_age_s, 7200); + const capped_ms = capped_s * 1000; + + const methods_wildcard = acam != null and std.mem.eql(u8, acam.?, "*") and !self.wants_credentials; + var methods = std.EnumSet(http.Method).initEmpty(); + if (!methods_wildcard) { + if (acam) |list| { + var it = std.mem.splitScalar(u8, list, ','); + while (it.next()) |raw| { + const token = std.mem.trim(u8, raw, &std.ascii.whitespace); + if (std.meta.stringToEnum(http.Method, token)) |m| methods.insert(m); + } + } + } + + const headers_wildcard = acah != null and std.mem.eql(u8, acah.?, "*") and !self.wants_credentials; + + var allowed_headers: std.ArrayList([]const u8) = .empty; + if (!headers_wildcard) { + if (acah) |list| { + var it = std.mem.splitScalar(u8, list, ','); + while (it.next()) |raw| { + const token = std.mem.trim(u8, raw, &std.ascii.whitespace); + if (token.len == 0) continue; + try allowed_headers.append(self.arena.allocator(), token); + } + } + } + + try self.gate.network.cors_store.put( + .{ .origin = self.origin, .target = self.url, .credentials = self.wants_credentials }, + .{ + .methods_wildcard = methods_wildcard, + .methods = methods, + .headers_wildcard = headers_wildcard, + .headers = allowed_headers.items, + .expires_at = lp.datetime.milliTimestamp(.real) + capped_ms, + }, + ); + } + fn methodAllowed(list: []const u8, method: http.Method) bool { const method_name = @tagName(method); var it = std.mem.splitScalar(u8, list, ','); @@ -393,6 +489,7 @@ const CorsPreflightContext = struct { var acam: ?[]const u8 = null; var acah: ?[]const u8 = null; var acac: ?[]const u8 = null; + var acma: ?[]const u8 = null; var iter = transfer.responseHeaderIterator(); while (iter.next()) |hdr| { @@ -404,15 +501,27 @@ const CorsPreflightContext = struct { acah = hdr.value; } else if (std.mem.eql(u8, hdr.name, ACCESS_CONTROL_ALLOW_CREDENTIALS)) { acac = hdr.value; + } else if (std.ascii.eqlIgnoreCase(ACCESS_CONTROL_MAX_AGE, hdr.name)) { + acma = hdr.value; } } self.allowed = self.validateHeaders(acao, acam, acah, acac); + if (self.allowed) { + self.acam = acam; + self.acah = acah; + self.acma = acma; + } return .proceed; } fn doneCallback(ctx_ptr: *anyopaque) anyerror!void { const self: *CorsPreflightContext = @ptrCast(@alignCast(ctx_ptr)); + if (self.allowed) { + self.cacheGrant(self.acam, self.acah, self.acma) catch |err| { + log.warn(.cors, "preflight cache store failed", .{ .url = self.url, .err = err }); + }; + } self.resolve(self.allowed); } @@ -441,31 +550,16 @@ const CorsPreflightContext = struct { } }; -fn fetchThenResume(self: *CorsGate, transfer: *Transfer) !void { +fn fetchThenResume(self: *CorsGate, transfer: *Transfer, authored_headers: []const []const u8) !void { const url = transfer.req.url; - const origin = transfer.req.origin orelse "null"; - - 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( - transfer.arena.allocator(), - try std.ascii.allocLowerString(transfer.arena.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); + const origin = transfer.effectiveOrigin(); const cors_key = CorsKey{ .url = url, .origin = origin, .method = transfer.req.method, .wants_credentials = transfer.req.credentials_mode == .include, - .authored_headers = header_names.items, + .authored_headers = authored_headers, }; const key = try cors_key.build(transfer.arena.allocator()); @@ -488,8 +582,8 @@ fn fetchThenResume(self: *CorsGate, transfer: *Transfer) !void { const referer: ?[]const u8 = transfer.findRequestHeader("referer"); - const owned_header_names = try arena.alloc([]const u8, header_names.items.len); - for (header_names.items, 0..) |name, i| { + const owned_header_names = try arena.alloc([]const u8, authored_headers.len); + for (authored_headers, 0..) |name, i| { owned_header_names[i] = try arena.dupe(u8, name); } @@ -527,11 +621,7 @@ fn fetchThenResume(self: *CorsGate, transfer: *Transfer) !void { errdefer fetch_transfer.deinit(); // Origin - try fetch_transfer.setHeader( - ORIGIN, - transfer.req.origin orelse "null", - .{}, - ); + try fetch_transfer.setHeader(ORIGIN, transfer.effectiveOrigin(), .{}); if (referer) |r| { try fetch_transfer.setHeader("Referer", r, .{}); @@ -549,8 +639,8 @@ fn fetchThenResume(self: *CorsGate, transfer: *Transfer) !void { ); // Access-Control-Allow-Headers - if (header_names.items.len > 0) { - const request_headers_value = try std.mem.join(arena.allocator(), ",", header_names.items); + if (authored_headers.len > 0) { + const request_headers_value = try std.mem.join(arena.allocator(), ",", authored_headers); try fetch_transfer.setHeader( ACCESS_CONTROL_REQUEST_HEADERS, request_headers_value, diff --git a/src/network/CorsStore.zig b/src/network/CorsStore.zig new file mode 100644 index 000000000..6674542b1 --- /dev/null +++ b/src/network/CorsStore.zig @@ -0,0 +1,418 @@ +// 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 isSafelistedMethod = @import("CorsGate.zig").isSafelistedMethod; +const ClockCache = @import("ClockCache.zig").ClockCache; + +const CorsStore = @This(); + +pub const Key = struct { + origin: []const u8, + target: []const u8, + credentials: bool, + + /// Serializes into a single string suitable as a ClockCache key. + fn build(self: Key, allocator: std.mem.Allocator) ![]const u8 { + var buf: std.ArrayList(u8) = .empty; + errdefer buf.deinit(allocator); + + try buf.appendSlice(allocator, self.origin); + try buf.append(allocator, 0); + try buf.appendSlice(allocator, self.target); + try buf.append(allocator, 0); + try buf.append(allocator, @intFromBool(self.credentials)); + + return buf.toOwnedSlice(allocator); + } +}; + +pub const Entry = struct { + methods_wildcard: bool, + methods: std.EnumSet(http.Method), + + headers_wildcard: bool, + headers: []const []const u8, + + expires_at: u64, + + 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 .{ + .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 = @min(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, + }; + } + + pub fn deinit(self: Entry, allocator: std.mem.Allocator) void { + for (self.headers) |h| allocator.free(h); + allocator.free(self.headers); + } +}; + +allocator: std.mem.Allocator, +map: ClockCache(Entry), +mutex: std.Io.Mutex = .init, + +pub fn init(allocator: std.mem.Allocator, capacity: usize) CorsStore { + return .{ .allocator = allocator, .map = .init(allocator, capacity) }; +} + +pub fn deinit(self: *CorsStore) void { + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + for (self.map.entries()) |*entry| { + entry.value.deinit(self.allocator); + } + self.map.deinit(); +} + +// Caller is expected to be holding mutex. +fn getWithExpiration(self: *CorsStore, cache_key: []const u8) ?*Entry { + const entry = self.map.get(cache_key) orelse return null; + + if (entry.expires_at <= lp.datetime.milliTimestamp(.real)) { + if (self.map.remove(cache_key)) |e| { + e.deinit(self.allocator); + } + + return null; + } + + return entry; +} + +fn matches(entry: Entry, method: http.Method, authored_headers: []const []const u8) bool { + if (!isSafelistedMethod(method) and !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; +} + +pub fn coversRequest( + self: *CorsStore, + allocator: std.mem.Allocator, + key: Key, + method: http.Method, + authored_headers: []const []const u8, +) !bool { + const primary_key = try key.build(allocator); + defer allocator.free(primary_key); + + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + if (self.getWithExpiration(primary_key)) |entry| { + if (matches(entry.*, method, authored_headers)) return true; + } + + if (key.credentials) return false; + + const cred_key = try (Key{ .origin = key.origin, .target = key.target, .credentials = true }).build(allocator); + defer allocator.free(cred_key); + + const entry = self.getWithExpiration(cred_key) orelse return false; + return matches(entry.*, method, authored_headers); +} + +/// 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 { + const cache_key = try key.build(self.allocator); + defer self.allocator.free(cache_key); + + self.mutex.lockUncancelable(lp.io); + defer self.mutex.unlock(lp.io); + + if (self.getWithExpiration(cache_key)) |existing| { + const merged = try existing.merge(self.allocator, entry); + existing.deinit(self.allocator); + existing.* = merged; + return; + } + + const owned_entry = try entry.dupe(self.allocator); + errdefer owned_entry.deinit(self.allocator); + + switch (try self.map.insert(cache_key, owned_entry)) { + .exists => unreachable, + .inserted => |evicted| { + if (evicted) |v| { + var e = v; + e.deinit(self.allocator); + } + }, + } +} + +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 covers, miss on different origin/target/credentials" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator, 10); + 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.milliTimestamp(.real) + 60_000, + }); + + try testing.expect(try store.coversRequest( + allocator, + .{ .origin = "https://a.example", .target = "https://api.example", .credentials = false }, + .POST, + &.{}, + )); + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = "https://a.example", .target = "https://api.example", .credentials = false }, + .PUT, + &.{}, + )); + + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = "https://b.example", .target = "https://api.example", .credentials = false }, + .POST, + &.{}, + )); + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = "https://a.example", .target = "https://other.example", .credentials = false }, + .POST, + &.{}, + )); + + // Same origin/target but different credentials mode: separate entry, must miss. + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = "https://a.example", .target = "https://api.example", .credentials = true }, + .POST, + &.{}, + )); +} + +test "CorsStore: expired entries are treated as a miss on covers" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator, 10); + 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.milliTimestamp(.real) - 1, + }); + + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = "https://a.example", .target = "https://api.example", .credentials = false }, + .GET, + &.{}, + )); +} + +test "CorsStore: put merges into existing entry rather than clobbering" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator, 10); + defer store.deinit(); + + const key = Key{ .origin = "https://a.example", .target = "https://api.example", .credentials = false }; + + const h1 = try allocator.alloc([]const u8, 1); + h1[0] = try allocator.dupe(u8, "x-one"); + try store.put(key, .{ + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.POST), + .headers_wildcard = false, + .headers = h1, + .expires_at = lp.datetime.milliTimestamp(.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, .{ + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.PUT), + .headers_wildcard = false, + .headers = h2, + .expires_at = lp.datetime.milliTimestamp(.real) + 60_000, + }); + freeHeaders(allocator, h2); + + try testing.expect(try store.coversRequest(allocator, key, .POST, &.{"x-one"})); + try testing.expect(try store.coversRequest(allocator, key, .PUT, &.{"x-two"})); + try testing.expect(!try store.coversRequest(allocator, key, .DELETE, &.{})); +} + +test "CorsStore: credentialed and non-credentialed grants for same origin/target stay separate" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator, 10); + defer store.deinit(); + + const origin = "https://a.example"; + const target = "https://api.example"; + + // Non-credentialed grant: wildcard headers allowed (valid per spec for non-cred requests). + try store.put(.{ .origin = origin, .target = target, .credentials = false }, .{ + .methods_wildcard = true, + .methods = .initEmpty(), + .headers_wildcard = true, + .headers = &.{}, + .expires_at = lp.datetime.milliTimestamp(.real) + 60_000, + }); + + // Credentialed grant: explicit methods/headers only, no wildcard. + const h = try allocator.alloc([]const u8, 1); + h[0] = try allocator.dupe(u8, "x-custom"); + try store.put(.{ .origin = origin, .target = target, .credentials = true }, .{ + .methods_wildcard = false, + .methods = std.EnumSet(http.Method).initOne(.GET), + .headers_wildcard = false, + .headers = h, + .expires_at = lp.datetime.milliTimestamp(.real) + 60_000, + }); + freeHeaders(allocator, h); + + // A credentialed request asking for an arbitrary header must be rejected + // against the credentialed entry, even though the non-cred entry has a wildcard. + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = origin, .target = target, .credentials = true }, + .GET, + &.{"x-anything"}, + )); + try testing.expect(try store.coversRequest( + allocator, + .{ .origin = origin, .target = target, .credentials = true }, + .GET, + &.{"x-custom"}, + )); + try testing.expect(!try store.coversRequest( + allocator, + .{ .origin = origin, .target = target, .credentials = true }, + .PUT, + &.{}, + )); + + // The non-credentialed entry's wildcard still works for non-cred requests. + try testing.expect(try store.coversRequest(allocator, .{ .origin = origin, .target = target, .credentials = false }, .GET, &.{"x-anything"})); +} + +test "CorsStore: covers never lets a wildcard cover Authorization" { + const allocator = testing.allocator; + var store = CorsStore.init(allocator, 10); + defer store.deinit(); + + const key = Key{ .origin = "https://a.example", .target = "https://api.example", .credentials = false }; + try store.put(key, .{ + .methods_wildcard = true, + .methods = .initEmpty(), + .headers_wildcard = true, + .headers = &.{}, + .expires_at = std.math.maxInt(u64), + }); + + try testing.expect(!try store.coversRequest(allocator, key, .GET, &.{"authorization"})); + try testing.expect(try store.coversRequest(allocator, key, .GET, &.{"x-anything"})); +} diff --git a/src/network/HttpClient.zig b/src/network/HttpClient.zig index 695209ff8..245d4255e 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..b27dc4329 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, config.corsStoreEntryLimit()), .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(); }