Skip to content

Commit b5fe4b7

Browse files
Add trust firewall proxy hooks for TrustedModules auto-reshim.
Co-authored-by: Cursor
1 parent cc36a66 commit b5fe4b7

3 files changed

Lines changed: 308 additions & 0 deletions

File tree

‎build.zig‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,20 @@ pub fn build(b: *std.Build) void {
153153
if (std.mem.eql(u8, app, "proxy")) {
154154
exe.root_module.addImport("install_safety", install_safety_module);
155155
exe.root_module.addImport("wintrust", wintrust_module);
156+
const module_firewall_module = b.createModule(.{
157+
.root_source_file = b.path("shared/module_firewall.zig"),
158+
.target = target,
159+
.optimize = optimize,
160+
});
161+
module_firewall_module.addImport("config", config_module);
162+
module_firewall_module.addImport("registry", registry_module);
163+
exe.root_module.addImport("module_firewall", module_firewall_module);
164+
const module_firewall_tests = b.addTest(.{
165+
.root_module = module_firewall_module,
166+
});
167+
module_firewall_tests.root_module.linkSystemLibrary("advapi32", .{});
168+
const run_module_firewall_tests = b.addRunArtifact(module_firewall_tests);
169+
test_step.dependOn(&run_module_firewall_tests.step);
156170
}
157171

158172
exe.root_module.addImport("errors", b.createModule(.{

‎proxy/main.zig‎

Lines changed: 124 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,9 @@ const errors = @import("errors");
77
const shimintegrity = @import("shimintegrity");
88
const verifycache = @import("verifycache");
99
const install_safety = @import("install_safety");
10+
const module_firewall = @import("module_firewall");
11+
const config = @import("config");
12+
const registry = @import("registry");
1013

1114
const ParsedArgs = struct {
1215
override_version: ?[]const u8,
@@ -226,6 +229,8 @@ pub fn main() !void {
226229

227230
const needs_reshim = detectReshimNeeded(command_name, parsed_args.forwarded);
228231

232+
const digest_before = hashFileOptional(allocator, command_path);
233+
229234
const process_exit_code = runDelegatedCommand(
230235
allocator,
231236
node_install_dir_abs,
@@ -245,6 +250,8 @@ pub fn main() !void {
245250
if (needs_reshim) {
246251
eventlog.writeInfo(allocator, "proxy", "reshim scheduled");
247252
runReshim(allocator, cfg.root, node_install_dir_abs);
253+
} else {
254+
try maybeReshimAfterSelfUpdate(allocator, cfg.root, node_install_dir_abs, command_name, command_path, digest_before);
248255
}
249256

250257
std.process.exit(process_exit_code);
@@ -1020,6 +1027,123 @@ fn parseArgs(allocator: std.mem.Allocator, args: []const []const u8) !ParsedArgs
10201027
};
10211028
}
10221029

1030+
fn hashFileOptional(allocator: std.mem.Allocator, path: []const u8) ?[32]u8 {
1031+
_ = allocator;
1032+
var file = std.fs.openFileAbsolute(path, .{}) catch return null;
1033+
defer file.close();
1034+
var hasher = std.crypto.hash.sha2.Sha256.init(.{});
1035+
var buf: [8192]u8 = undefined;
1036+
while (true) {
1037+
const n = file.read(&buf) catch return null;
1038+
if (n == 0) break;
1039+
hasher.update(buf[0..n]);
1040+
}
1041+
var out: [32]u8 = undefined;
1042+
hasher.final(&out);
1043+
return out;
1044+
}
1045+
1046+
fn digestsEqual(a: ?[32]u8, b: ?[32]u8) bool {
1047+
if (a == null or b == null) return true; // skip trust path if unreadable
1048+
return std.mem.eql(u8, &a.?, &b.?);
1049+
}
1050+
1051+
fn loadTrustedModules(allocator: std.mem.Allocator) ![]const []const u8 {
1052+
const rules = module_firewall.loadMultiSzPolicy(allocator, module_firewall.reg_value_trusted_modules) catch &[_][]const u8{};
1053+
if (rules.len == 0) {
1054+
// Default NOT ALL
1055+
var out = try allocator.alloc([]const u8, 1);
1056+
out[0] = try allocator.dupe(u8, "NOT ALL");
1057+
return out;
1058+
}
1059+
return rules;
1060+
}
1061+
1062+
fn untrustedHandlerIsPrompt(allocator: std.mem.Allocator) bool {
1063+
const raw = registry.queryStringWithFallback(allocator, registry.preferenceHives(), config.preference_registry_root, module_firewall.reg_value_untrusted_handler) catch {
1064+
return false;
1065+
};
1066+
defer allocator.free(raw);
1067+
return std.ascii.eqlIgnoreCase(std.mem.trim(u8, raw, " \t"), "prompt");
1068+
}
1069+
1070+
fn promptTrustChange(allocator: std.mem.Allocator, command_name: []const u8) bool {
1071+
const nvm_path = nodeversion.resolveNvmExePath(allocator) catch {
1072+
return promptTrustChangeConsoleOnly(allocator, command_name);
1073+
};
1074+
defer allocator.free(nvm_path);
1075+
1076+
var child = std.process.Child.init(&.{ nvm_path, "firewall", "prompt-trust", command_name }, allocator);
1077+
child.stdin_behavior = .Inherit;
1078+
child.stdout_behavior = .Inherit;
1079+
child.stderr_behavior = .Inherit;
1080+
const term = child.spawnAndWait() catch {
1081+
return promptTrustChangeConsoleOnly(allocator, command_name);
1082+
};
1083+
return switch (term) {
1084+
.Exited => |code| code == 0,
1085+
else => false,
1086+
};
1087+
}
1088+
1089+
fn promptTrustChangeConsoleOnly(allocator: std.mem.Allocator, command_name: []const u8) bool {
1090+
const msg = std.fmt.allocPrint(allocator, "Untrusted module '{s}' changed after running. Approve and reshim? [y/N]: ", .{command_name}) catch {
1091+
return false;
1092+
};
1093+
defer allocator.free(msg);
1094+
std.debug.print("{s}", .{msg});
1095+
var stdin_buffer: [16]u8 = undefined;
1096+
const stdin = std.fs.File.stdin();
1097+
const n = stdin.read(stdin_buffer[0..]) catch return false;
1098+
if (n == 0) return false;
1099+
const answer = std.mem.trim(u8, stdin_buffer[0..n], " \t\r\n");
1100+
return answer.len > 0 and (answer[0] == 'y' or answer[0] == 'Y');
1101+
}
1102+
1103+
fn maybeReshimAfterSelfUpdate(
1104+
allocator: std.mem.Allocator,
1105+
install_root: []const u8,
1106+
node_install_dir: []const u8,
1107+
command_name: []const u8,
1108+
command_path: []const u8,
1109+
digest_before: ?[32]u8,
1110+
) !void {
1111+
const digest_after = hashFileOptional(allocator, command_path);
1112+
if (digestsEqual(digest_before, digest_after)) return;
1113+
1114+
// Package managers already handled via needs_reshim.
1115+
if (std.ascii.eqlIgnoreCase(command_name, "npm") or
1116+
std.ascii.eqlIgnoreCase(command_name, "npx") or
1117+
std.ascii.eqlIgnoreCase(command_name, "pnpm") or
1118+
std.ascii.eqlIgnoreCase(command_name, "yarn") or
1119+
std.ascii.eqlIgnoreCase(command_name, "corepack") or
1120+
std.ascii.eqlIgnoreCase(command_name, "vlt"))
1121+
{
1122+
return;
1123+
}
1124+
1125+
const rules = try loadTrustedModules(allocator);
1126+
defer module_firewall.freeMultiSz(allocator, rules);
1127+
1128+
const pkg = module_firewall.PackageSpec{ .name = command_name, .version = "", .raw = command_name };
1129+
const trusted = module_firewall.isPackageAllowed(pkg, rules) orelse false;
1130+
if (trusted) {
1131+
eventlog.writeInfo(allocator, "proxy", "firewall trusted module changed; scheduling reshim");
1132+
runReshim(allocator, install_root, node_install_dir);
1133+
return;
1134+
}
1135+
if (untrustedHandlerIsPrompt(allocator)) {
1136+
if (promptTrustChange(allocator, command_name)) {
1137+
eventlog.writeInfo(allocator, "proxy", "firewall trust prompt accepted; scheduling reshim");
1138+
runReshim(allocator, install_root, node_install_dir);
1139+
} else {
1140+
eventlog.writeInfo(allocator, "proxy", "firewall trust prompt declined; VerifyCache left stale");
1141+
}
1142+
return;
1143+
}
1144+
eventlog.writeInfo(allocator, "proxy", "firewall untrusted module changed; reshim not scheduled (deny)");
1145+
}
1146+
10231147
/// Returns true when the invoked package manager command is likely to install
10241148
/// or remove a globally-visible executable that reshim needs to reconcile.
10251149
fn detectReshimNeeded(command_name: []const u8, args: []const []const u8) bool {

‎shared/module_firewall.zig‎

Lines changed: 170 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,170 @@
1+
const std = @import("std");
2+
const config = @import("config");
3+
const registry = @import("registry");
4+
5+
pub const reg_value_trusted_modules = "TrustedModules";
6+
pub const reg_value_approved_modules = "ApprovedModules";
7+
pub const reg_value_approved_global_modules = "ApprovedGlobalModules";
8+
pub const reg_value_untrusted_handler = "UntrustedModuleHandlerAction";
9+
10+
pub const PackageSpec = struct {
11+
name: []const u8,
12+
version: []const u8,
13+
raw: []const u8,
14+
};
15+
16+
fn trimAscii(s: []const u8) []const u8 {
17+
return std.mem.trim(u8, s, " \t\r\n");
18+
}
19+
20+
fn eqlIgnoreCase(a: []const u8, b: []const u8) bool {
21+
return std.ascii.eqlIgnoreCase(a, b);
22+
}
23+
24+
fn startsWithIgnoreCase(hay: []const u8, needle: []const u8) bool {
25+
if (hay.len < needle.len) return false;
26+
return eqlIgnoreCase(hay[0..needle.len], needle);
27+
}
28+
29+
pub fn splitNameVersion(spec: []const u8) struct { name: []const u8, version: []const u8 } {
30+
const s = trimAscii(spec);
31+
if (s.len == 0) return .{ .name = "", .version = "" };
32+
if (s[0] == '@') {
33+
if (std.mem.indexOfScalar(u8, s[1..], '/')) |slash_rel| {
34+
const slash = slash_rel + 1;
35+
const after = s[slash + 1 ..];
36+
if (std.mem.indexOfScalar(u8, after, '@')) |at| {
37+
return .{ .name = s[0 .. slash + 1 + at], .version = after[at + 1 ..] };
38+
}
39+
return .{ .name = s, .version = "" };
40+
}
41+
return .{ .name = s, .version = "" };
42+
}
43+
if (std.mem.indexOfScalar(u8, s, '@')) |at| {
44+
return .{ .name = s[0..at], .version = s[at + 1 ..] };
45+
}
46+
return .{ .name = s, .version = "" };
47+
}
48+
49+
fn normalizeRule(raw: []const u8) struct { entry: []const u8, negated: bool } {
50+
var entry = trimAscii(raw);
51+
var negated = false;
52+
if (startsWithIgnoreCase(entry, "not ") or startsWithIgnoreCase(entry, "not\t")) {
53+
negated = true;
54+
entry = trimAscii(entry[3..]);
55+
} else if (entry.len > 0 and entry[0] == '!') {
56+
negated = true;
57+
entry = trimAscii(entry[1..]);
58+
}
59+
return .{ .entry = entry, .negated = negated };
60+
}
61+
62+
fn nameMatches(pattern: []const u8, pkg_name: []const u8) bool {
63+
if (eqlIgnoreCase(pattern, "all")) return true;
64+
if (std.mem.endsWith(u8, pattern, "/*")) {
65+
const org = pattern[0 .. pattern.len - 2];
66+
if (pkg_name.len <= org.len) return false;
67+
if (!eqlIgnoreCase(pkg_name[0..org.len], org)) return false;
68+
return pkg_name[org.len] == '/';
69+
}
70+
return eqlIgnoreCase(pattern, pkg_name);
71+
}
72+
73+
fn versionMatches(pattern_ver: []const u8, pkg_ver: []const u8) bool {
74+
if (pattern_ver.len == 0) return true;
75+
if (pkg_ver.len == 0) return true;
76+
if (std.mem.endsWith(u8, pattern_ver, ".*")) {
77+
const prefix = pattern_ver[0 .. pattern_ver.len - 2];
78+
if (eqlIgnoreCase(pkg_ver, prefix)) return true;
79+
if (pkg_ver.len > prefix.len and eqlIgnoreCase(pkg_ver[0..prefix.len], prefix) and pkg_ver[prefix.len] == '.') return true;
80+
return false;
81+
}
82+
// Exact or prefix equality for MVP ranges like >= are handled loosely: exact match only in Zig;
83+
// full semver ranges evaluated in Go path / HTTPS. For common 1.0.0 pins:
84+
return eqlIgnoreCase(pattern_ver, pkg_ver) or startsWithIgnoreCase(pattern_ver, ">=") or startsWithIgnoreCase(pattern_ver, "^") or startsWithIgnoreCase(pattern_ver, "~");
85+
}
86+
87+
fn ruleMatches(rule_entry: []const u8, pkg: PackageSpec) bool {
88+
if (eqlIgnoreCase(rule_entry, "all")) return true;
89+
const nv = splitNameVersion(rule_entry);
90+
if (!nameMatches(nv.name, pkg.name)) return false;
91+
// For range operators in Zig MVP, treat as name match (policy still useful); exact pins enforced.
92+
if (nv.version.len > 0 and (nv.version[0] == '>' or nv.version[0] == '^' or nv.version[0] == '~' or nv.version[0] == '<')) {
93+
return true;
94+
}
95+
return versionMatches(nv.version, pkg.version);
96+
}
97+
98+
/// True when policy list contains an https:// URL (remote evaluation mode).
99+
pub fn listHasHttps(rules: []const []const u8) bool {
100+
for (rules) |r| {
101+
const e = trimAscii(r);
102+
if (startsWithIgnoreCase(e, "https://")) return true;
103+
}
104+
return false;
105+
}
106+
107+
/// Local-list evaluation (VersionAllowList-compatible). HTTPS URL lists return null (caller does remote).
108+
pub fn isPackageAllowed(pkg: PackageSpec, rules: []const []const u8) ?bool {
109+
if (listHasHttps(rules)) return null;
110+
111+
var not_all = false;
112+
var has_exclusive = false;
113+
var i: usize = 0;
114+
while (i < rules.len) : (i += 1) {
115+
const norm = normalizeRule(rules[i]);
116+
if (norm.entry.len == 0) continue;
117+
if (norm.negated and eqlIgnoreCase(norm.entry, "all")) not_all = true;
118+
if (!norm.negated and !eqlIgnoreCase(norm.entry, "all")) has_exclusive = true;
119+
}
120+
121+
if (not_all) {
122+
i = 0;
123+
while (i < rules.len) : (i += 1) {
124+
const norm = normalizeRule(rules[i]);
125+
if (norm.negated) continue;
126+
if (ruleMatches(norm.entry, pkg)) return true;
127+
}
128+
return false;
129+
}
130+
131+
i = 0;
132+
while (i < rules.len) : (i += 1) {
133+
const norm = normalizeRule(rules[i]);
134+
if (!norm.negated) continue;
135+
if (ruleMatches(norm.entry, pkg)) return false;
136+
}
137+
i = 0;
138+
while (i < rules.len) : (i += 1) {
139+
const norm = normalizeRule(rules[i]);
140+
if (norm.negated) continue;
141+
if (ruleMatches(norm.entry, pkg)) return true;
142+
}
143+
if (has_exclusive) return false;
144+
return true;
145+
}
146+
147+
pub fn loadMultiSzPolicy(allocator: std.mem.Allocator, value_name: []const u8) ![]const []const u8 {
148+
const policy_hives = [_]registry.Hive{ .hkey_local_machine, .hkey_current_user };
149+
if (registry.queryMultiStringOptionalWithFallback(allocator, &policy_hives, config.policy_registry_root, value_name) catch null) |vals| {
150+
return vals;
151+
}
152+
if (registry.queryMultiStringOptionalWithFallback(allocator, registry.preferenceHives(), config.preference_registry_root, value_name) catch null) |vals| {
153+
return vals;
154+
}
155+
return &[_][]const u8{};
156+
}
157+
158+
pub fn freeMultiSz(allocator: std.mem.Allocator, vals: []const []const u8) void {
159+
if (vals.len == 0) return;
160+
for (vals) |v| allocator.free(v);
161+
allocator.free(vals);
162+
}
163+
164+
test "not all with exception" {
165+
const rules = [_][]const u8{ "NOT ALL", "porthog" };
166+
const ok = isPackageAllowed(.{ .name = "porthog", .version = "", .raw = "porthog" }, &rules);
167+
try std.testing.expect(ok != null and ok.?);
168+
const deny = isPackageAllowed(.{ .name = "eslint", .version = "", .raw = "eslint" }, &rules);
169+
try std.testing.expect(deny != null and !deny.?);
170+
}

0 commit comments

Comments
 (0)