lib/pluck/src/toplevel/lifecycle.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_arena = @import("alloc_arena");
3 const pluck = @import("../root.zig");
4 const log = pluck.logger;
5 const Allocator = std.mem.Allocator;
6
7 const pexpr = pluck.pexpr;
8 const Definitions = pexpr.Definitions;
9 const TypeRegistry = pexpr.TypeRegistry;
10
11 const lpsmc_module = pluck.lpsmc;
12
13 const bdd = pluck.bdd;
14 const Manager = bdd.Manager;
15
16 const context_owner = @import("context.zig");
17 const ToplevelContext = context_owner.ToplevelContext;
18 const top_types = @import("types.zig");
19 const ToplevelError = top_types.ToplevelError;
20 const config_owner = @import("config.zig");
21 const ToplevelConfig = config_owner.ToplevelConfig;
22 const ToplevelConfigPatch = config_owner.ToplevelConfigPatch;
23 const ConfigError = config_owner.ConfigError;
24 const shared_owner = @import("shared.zig");
25 const source_owner = @import("source.zig");
26
27 const stdlib_source = @embedFile("../stdlib.pluck");
28
29 pub fn init(allocator: Allocator) !ToplevelContext {
30 const arena_ptr = try allocator.create(alloc_arena.Arena);
31 errdefer allocator.destroy(arena_ptr);
32 arena_ptr.* = alloc_arena.Arena.init(allocator);
33 errdefer arena_ptr.deinit();
34
35 const query_arena_ptr = try allocator.create(alloc_arena.Arena);
36 errdefer allocator.destroy(query_arena_ptr);
37 query_arena_ptr.* = alloc_arena.Arena.init(allocator);
38 errdefer query_arena_ptr.deinit();
39
40 const arena_alloc = arena_ptr.allocator();
41
42 const types = try TypeRegistry.initWithDefaults(arena_alloc);
43
44 const definitions = Definitions.init(arena_alloc);
45
46 const manager = try allocator.create(Manager);
47 errdefer allocator.destroy(manager);
48 manager.* = try Manager.init(arena_alloc);
49 errdefer manager.deinit();
50
51 var self = ToplevelContext{
52 .allocator = allocator,
53 .owns_shared = true,
54 .owns_arena = true,
55 .arena = arena_ptr,
56 .query_arena = query_arena_ptr,
57 .types = types,
58 .definitions = definitions,
59 .manager = manager,
60 .config = .{},
61 };
62
63 try loadStdlib(&self);
64
65 return self;
66 }
67
68 pub fn initWorker(self: *const ToplevelContext, allocator: Allocator) !ToplevelContext {
69 const arena_ptr = try allocator.create(alloc_arena.Arena);
70 errdefer allocator.destroy(arena_ptr);
71 arena_ptr.* = alloc_arena.Arena.init(allocator);
72 errdefer arena_ptr.deinit();
73
74 const query_arena_ptr = try allocator.create(alloc_arena.Arena);
75 errdefer allocator.destroy(query_arena_ptr);
76 query_arena_ptr.* = alloc_arena.Arena.init(allocator);
77 errdefer query_arena_ptr.deinit();
78
79 const manager = try allocator.create(Manager);
80 errdefer allocator.destroy(manager);
81 manager.* = try Manager.init(arena_ptr.allocator());
82 errdefer manager.deinit();
83
84 const arena_alloc = arena_ptr.allocator();
85
86 var cloned_defs = try shared_owner.cloneDefinitionsShallow(arena_alloc, &self.definitions);
87 errdefer cloned_defs.defs.deinit();
88
89 var cloned_types = try shared_owner.cloneTypesShallow(arena_alloc, &self.types);
90 errdefer {
91 cloned_types.type_of_constructor.deinit();
92 cloned_types.constructors_of_type.deinit();
93 cloned_types.args_of_constructor.deinit();
94 }
95
96 var worker = self.*;
97 worker.allocator = allocator;
98 worker.owns_shared = false;
99 worker.owns_arena = true;
100 worker.arena = arena_ptr;
101 worker.query_arena = query_arena_ptr;
102 worker.manager = manager;
103 worker.definitions = cloned_defs;
104 worker.types = cloned_types;
105 worker.incremental_lpsmc = null;
106 worker.last_query_constructor_typo = null;
107 worker.current_source = null;
108 worker.loading_stdlib = false;
109 return worker;
110 }
111
112 pub fn loadStdlib(self: *ToplevelContext) error{ OutOfMemory, StdlibLoadFailed }!void {
113 self.loading_stdlib = true;
114 defer self.loading_stdlib = false;
115
116 source_owner.processSourceSexprNoWriter(self, stdlib_source) catch |err| {
117 if (@import("builtin").mode == .debug) {
118 log.err("stdlib load failed: {}", .{err});
119 }
120 return switch (err) {
121 ToplevelError.OutOfMemory => error.OutOfMemory,
122 else => error.StdlibLoadFailed,
123 };
124 };
125 }
126
127 pub fn deinit(self: *ToplevelContext) void {
128 if (!self.owns_shared) {
129 self.definitions.defs.deinit();
130 self.types.type_of_constructor.deinit();
131 self.types.constructors_of_type.deinit();
132 self.types.args_of_constructor.deinit();
133 }
134
135 if (self.incremental_lpsmc) |lpsmc| {
136 lpsmc_module.deinit(lpsmc);
137 self.allocator.destroy(lpsmc);
138 }
139
140 self.manager.deinit();
141 self.allocator.destroy(self.manager);
142 if (self.owns_arena) {
143 self.arena.deinit();
144 self.allocator.destroy(self.arena);
145 }
146 self.query_arena.deinit();
147 self.allocator.destroy(self.query_arena);
148 }
149
150 pub fn setConfig(self: *ToplevelContext, config: ToplevelConfig) void {
151 self.config = config;
152 }
153
154 pub fn patchConfig(self: *ToplevelContext, patch: ToplevelConfigPatch) ConfigError!void {
155 try config_owner.applyPatch(&self.config, patch);
156 }
157
158 const query_context = pluck.query_context;
159 pub const SharedContext = query_context.SharedContext;
160 pub const QueryContext = query_context.QueryContext;
161 pub const RunContext = query_context.RunContext;
162 pub const QueryConfig = query_context.QueryConfig;
163 pub const SessionConfig = query_context.SessionConfig;
164
165 pub fn asSharedContext(self: *const ToplevelContext) SharedContext {
166 return SharedContext{
167 .allocator = self.allocator,
168 .arena = self.arena,
169 .types = &self.types,
170 .definitions = &self.definitions,
171 .config = SessionConfig{
172 .max_depth = self.config.max_depth,
173 .ite_limit = self.config.ite_limit,
174 .time_limit = self.config.time_limit,
175 .sample_after_max_depth = self.config.sample_after_max_depth,
176 .use_strict_order = self.config.use_strict_order,
177 .use_reverse_order = self.config.use_reverse_order,
178 .definition_order_mode = self.config.definition_order_mode,
179 .parallel_wmc = self.config.parallel_wmc,
180 .verbose = self.config.verbose,
181 .rng_seed = self.config.rng_seed,
182 .lpsmc_rng_seed = self.config.lpsmc_rng_seed,
183 },
184 };
185 }
186
187 pub fn createQueryContext(self: *const ToplevelContext, query_config: QueryConfig) !QueryContext {
188 return try query_context.initQueryContext(asSharedContext(self), query_config);
189 }
190
191 pub fn reset(self: *ToplevelContext) error{ OutOfMemory, StdlibLoadFailed, MacroAlreadyExists }!void {
192 if (self.incremental_lpsmc) |lpsmc| {
193 lpsmc_module.clearCaches(lpsmc);
194 }
195
196 _ = self.arena.reset(.retain_capacity);
197 const alloc = self.arena.allocator();
198
199 self.types = try TypeRegistry.initWithDefaults(alloc);
200 self.definitions = Definitions.init(alloc);
201
202 self.manager.* = try Manager.init(alloc);
203
204 self.current_source = null;
205 self.last_query_constructor_typo = null;
206
207 try loadStdlib(self);
208 }
209
210 pub fn resetManagerForQuery(self: *ToplevelContext) !void {
211 const manager_alloc = self.manager.allocator;
212 self.manager.deinit();
213 self.manager.* = try Manager.init(manager_alloc);
214 }
215
216 test "query type registry matches Pluck.jl public constructors" {
217 const allocator = std.testing.allocator;
218 var ctx = try init(allocator);
219 defer ctx.deinit();
220
221 try std.testing.expect(ctx.types.hasConstructor("Marginal"));
222 try std.testing.expect(ctx.types.hasConstructor("Posterior"));
223 try std.testing.expect(ctx.types.hasConstructor("PosteriorSamples"));
224 try std.testing.expect(ctx.types.hasConstructor("AdaptiveRejection"));
225 try std.testing.expect(!ctx.types.hasConstructor("SubproblemMonteCarlo"));
226
227 try ctx.reset();
228 try std.testing.expect(ctx.types.hasConstructor("Marginal"));
229 try std.testing.expect(ctx.types.hasConstructor("Posterior"));
230 try std.testing.expect(ctx.types.hasConstructor("PosteriorSamples"));
231 try std.testing.expect(ctx.types.hasConstructor("AdaptiveRejection"));
232 try std.testing.expect(!ctx.types.hasConstructor("SubproblemMonteCarlo"));
233 }
234
235 test "patchConfig updates context config" {
236 var ctx = try init(std.testing.allocator);
237 defer ctx.deinit();
238
239 try ctx.patchConfig(.{
240 .time_limit = .{ .set = 0.5 },
241 .var_order = .reverse,
242 .fallback_lpsmc_k = 3,
243 });
244
245 try std.testing.expectEqual(@as(?f64, 0.5), ctx.config.time_limit);
246 try std.testing.expect(ctx.config.use_strict_order);
247 try std.testing.expect(ctx.config.use_reverse_order);
248 try std.testing.expectEqual(@as(usize, 3), ctx.config.fallback_lpsmc_k);
249 }
250
251 test "worker context validates definitions without mutating parent" {
252 const allocator = std.testing.allocator;
253 var ctx = try init(allocator);
254 defer ctx.deinit();
255
256 try ctx.processSource("(define p 0.25)");
257
258 var worker = try ctx.initWorker(allocator);
259 defer worker.deinit();
260 try worker.processSource(
261 \\(define-type result (Yes) (No))
262 \\(define p 0.75)
263 );
264
265 try std.testing.expect(worker.types.hasConstructor("Yes"));
266 try std.testing.expect(!ctx.types.hasConstructor("Yes"));
267
268 var parent_result = try ctx.processForm("(Marginal (flip p))") orelse return error.TestExpectedEqual;
269 defer parent_result.deinit();
270 var parent_true: f64 = 0;
271 for (parent_result.outcomes) |outcome| {
272 if (std.mem.eql(u8, outcome.value_str, "True")) parent_true = outcome.probability;
273 }
274 try std.testing.expectApproxEqAbs(@as(f64, 0.25), parent_true, 1e-9);
275
276 var worker_result = try worker.processForm("(Marginal (flip p))") orelse return error.TestExpectedEqual;
277 defer worker_result.deinit();
278 var worker_true: f64 = 0;
279 for (worker_result.outcomes) |outcome| {
280 if (std.mem.eql(u8, outcome.value_str, "True")) worker_true = outcome.probability;
281 }
282 try std.testing.expectApproxEqAbs(@as(f64, 0.75), worker_true, 1e-9);
283 }
284
285 test "stdlib list and nat helpers support comparison programs" {
286 const allocator = std.testing.allocator;
287 var ctx = try init(allocator);
288 defer ctx.deinit();
289
290 const maybe_result = try ctx.processForm(
291 "(Marginal (and (eq_nat (length (Cons True (Cons False Nil))) 2) (list_eq eq_nat (take 2 (Cons 0 (Cons 1 (Cons 2 Nil)))) (Cons 0 (Cons 1 Nil)))))",
292 );
293 try std.testing.expect(maybe_result != null);
294 var result = maybe_result.?;
295 defer result.deinit();
296
297 try std.testing.expect(!result.program_error);
298 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
299 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
300 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
301 }
302
303 test "stdlib operator names support Pluck.jl programs" {
304 const allocator = std.testing.allocator;
305 var ctx = try init(allocator);
306 defer ctx.deinit();
307
308 const maybe_result = try ctx.processForm(
309 "(Marginal (and (nat=? (+ 1 2) 3) (and (nat=? (- 3 1) 2) (and (nat=? (* 2 3) 6) (and (> 3 2) (and (< 2 3) (and (>= 3 3) (and (<= 3 3) (and (list=? nat=? [1 2] (Cons 1 (Cons 2 Nil))) (=? (Cons True Nil) (Cons True Nil)))))))))))",
310 );
311 try std.testing.expect(maybe_result != null);
312 var result = maybe_result.?;
313 defer result.deinit();
314
315 try std.testing.expect(!result.program_error);
316 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
317 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
318 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
319 }
320
321 test "stdlib reference helpers support Pluck.jl programs" {
322 const allocator = std.testing.allocator;
323 var ctx = try init(allocator);
324 defer ctx.deinit();
325
326 const maybe_result = try ctx.processForm(
327 "(Marginal (and (nat=? (dec 2) 1) (and (iseven 2) (list=? nat=? (zip_with + [1 2] [3 4]) [4 6]))))",
328 );
329 try std.testing.expect(maybe_result != null);
330 var result = maybe_result.?;
331 defer result.deinit();
332
333 try std.testing.expect(!result.program_error);
334 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
335 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
336 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
337 }
338
339 test "stdlib reference list helpers match Pluck.jl names" {
340 const allocator = std.testing.allocator;
341 var ctx = try init(allocator);
342 defer ctx.deinit();
343
344 const maybe_result = try ctx.processForm(
345 "(Marginal (and (nat=? (mod 7 3) 1) (and (isempty (cdr_safe Nil)) (and (nat=? (fold + 0 [1 2 3]) 6) (and (list=? nat=? (mapi + [2 3]) [2 4]) (and (list=? nat=? (filter (fn x -> (< x 3)) [1 3 2]) [1 2]) (and (list=? nat=? (filteri (fn x i -> (< i 2)) [5 6 7]) [5 6]) (and (list=? nat=? (range 3) [0 1 2]) (and (list=? nat=? (append_one [1 2] 3) [1 2 3]) (and (nat=? (head-or Nil 4) 4) (and (nat=? (head-or [5] 4) 5) (and (int2=? (Int2 True False) (Int2 True False)) (constructor=? (suspendible-list=? nat=? [1] [1]) (Suspend FinallyTrue))))))))))))))",
346 );
347 try std.testing.expect(maybe_result != null);
348 var result = maybe_result.?;
349 defer result.deinit();
350
351 try std.testing.expect(!result.program_error);
352 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
353 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
354 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
355 }
356
357 test "stdlib reference distribution helpers parse and evaluate" {
358 const allocator = std.testing.allocator;
359 var ctx = try init(allocator);
360 defer ctx.deinit();
361
362 const uniform_result = try ctx.processForm("(Marginal (make-uniform [True False]))");
363 try std.testing.expect(uniform_result != null);
364 var uniform = uniform_result.?;
365 defer uniform.deinit();
366
367 try std.testing.expect(!uniform.program_error);
368 try std.testing.expectEqual(@as(usize, 1), uniform.outcomes.len);
369 try std.testing.expectEqualStrings("(Cons (Pair True 0.5) (Cons (Pair False 0.5) (Nil)))", uniform.outcomes[0].value_str);
370 try std.testing.expectApproxEqAbs(@as(f64, 1.0), uniform.outcomes[0].probability, 1e-10);
371
372 const normalized_result = try ctx.processForm("(Marginal (normalize-seq [(Pair True 1.0) (Pair False 1.0)]))");
373 try std.testing.expect(normalized_result != null);
374 var normalized = normalized_result.?;
375 defer normalized.deinit();
376
377 try std.testing.expect(!normalized.program_error);
378 try std.testing.expectEqual(@as(usize, 1), normalized.outcomes.len);
379 try std.testing.expectEqualStrings("(Cons (Pair True 0.5) (Cons (Pair False 1) (Nil)))", normalized.outcomes[0].value_str);
380 try std.testing.expectApproxEqAbs(@as(f64, 1.0), normalized.outcomes[0].probability, 1e-10);
381
382 const sample_result = try ctx.processForm("(Marginal (sample-seq (normalize-seq [(Pair True 1.0) (Pair False 1.0)])))");
383 try std.testing.expect(sample_result != null);
384 var sample = sample_result.?;
385 defer sample.deinit();
386
387 try std.testing.expect(!sample.program_error);
388 try std.testing.expectEqual(@as(usize, 2), sample.outcomes.len);
389 try std.testing.expectEqualStrings("False", sample.outcomes[0].value_str);
390 try std.testing.expectApproxEqAbs(@as(f64, 0.5), sample.outcomes[0].probability, 1e-10);
391
392 const randnat_result = try ctx.processForm("(Marginal (< (randnat) 2))");
393 try std.testing.expect(randnat_result != null);
394 var randnat = randnat_result.?;
395 defer randnat.deinit();
396
397 try std.testing.expect(!randnat.program_error);
398 try std.testing.expectEqual(@as(usize, 2), randnat.outcomes.len);
399 try std.testing.expectEqualStrings("False", randnat.outcomes[0].value_str);
400 try std.testing.expectApproxEqAbs(@as(f64, 0.25), randnat.outcomes[0].probability, 1e-10);
401 try std.testing.expectEqualStrings("True", randnat.outcomes[1].value_str);
402 try std.testing.expectApproxEqAbs(@as(f64, 0.75), randnat.outcomes[1].probability, 1e-10);
403 }
404
405 test "lowercase constructors from Pluck.jl programs parse and evaluate" {
406 const allocator = std.testing.allocator;
407 var ctx = try init(allocator);
408 defer ctx.deinit();
409
410 try std.testing.expect(try ctx.processForm("(define-type char (a_) (b_))") == null);
411
412 const maybe_result = try ctx.processForm("(Marginal (constructor=? (a_) (a_)))");
413 try std.testing.expect(maybe_result != null);
414 var result = maybe_result.?;
415 defer result.deinit();
416
417 try std.testing.expect(!result.program_error);
418 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
419 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
420 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
421 }
422
423 test "full distribution resolves finite lazy values beyond one thousand thunks" {
424 const allocator = std.testing.allocator;
425 var ctx = try init(allocator);
426 defer ctx.deinit();
427
428 try std.testing.expect(try ctx.processForm("(define (long-list n) (case n of O => Nil | S k => (Cons O (long-list k))))") == null);
429
430 const maybe_result = try ctx.processForm("(Marginal (long-list 1005))");
431 try std.testing.expect(maybe_result != null);
432 var result = maybe_result.?;
433 defer result.deinit();
434
435 try std.testing.expect(!result.program_error);
436 try std.testing.expect(result.limit_reason == null);
437 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
438 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
439 }
440
441 test "stdlib ADT equality helpers compare constructed values" {
442 const allocator = std.testing.allocator;
443 var ctx = try init(allocator);
444 defer ctx.deinit();
445
446 const maybe_result = try ctx.processForm(
447 "(Marginal (adt_eq (Cons True (Cons False Nil)) (Cons True (Cons False Nil))))",
448 );
449 try std.testing.expect(maybe_result != null);
450 var result = maybe_result.?;
451 defer result.deinit();
452
453 try std.testing.expect(!result.program_error);
454 try std.testing.expectEqual(@as(usize, 1), result.outcomes.len);
455 try std.testing.expectEqualStrings("True", result.outcomes[0].value_str);
456 try std.testing.expectApproxEqAbs(@as(f64, 1.0), result.outcomes[0].probability, 1e-10);
457 }