lib/pluck/src/toplevel/script.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const pluck = @import("../root.zig");
3 const Allocator = std.mem.Allocator;
4
5 const pexpr = pluck.pexpr;
6 const context_owner = @import("context.zig");
7 const ToplevelContext = context_owner.ToplevelContext;
8 const top_types = @import("types.zig");
9 const ToplevelError = top_types.ToplevelError;
10 const SourceLocation = top_types.SourceLocation;
11 const SourceSpan = top_types.SourceSpan;
12 const result_owner = @import("result.zig");
13 const QueryResult = result_owner.QueryResult;
14 const forms_owner = @import("forms.zig");
15 const source_owner = @import("source.zig");
16
17 pub const FormKind = enum {
18 definition,
19 type_definition,
20 query,
21 expression,
22 };
23
24 pub const SourceForm = struct {
25 kind: FormKind,
26 source: []const u8,
27 name: ?[]const u8,
28 span: SourceSpan,
29 location: SourceLocation,
30 };
31
32 pub const SourceForms = struct {
33 allocator: Allocator,
34 items: []SourceForm,
35
36 pub fn deinit(self: *SourceForms) void {
37 self.allocator.free(self.items);
38 self.items = &[_]SourceForm{};
39 }
40 };
41
42 pub const QueryRecord = struct {
43 form_index: usize,
44 form: SourceForm,
45 result: QueryResult,
46 };
47
48 pub const Execution = struct {
49 allocator: Allocator,
50 forms: []SourceForm,
51 queries: []QueryRecord,
52
53 pub fn deinit(self: *Execution) void {
54 for (self.queries) |*query| {
55 query.result.deinit();
56 }
57 self.allocator.free(self.queries);
58 self.allocator.free(self.forms);
59 self.queries = &[_]QueryRecord{};
60 self.forms = &[_]SourceForm{};
61 }
62 };
63
64 pub fn parseSource(allocator: Allocator, source: []const u8) ToplevelError!SourceForms {
65 return scanSource(allocator, source, null);
66 }
67
68 pub fn parseSourceForms(self: *ToplevelContext, allocator: Allocator, source: []const u8) ToplevelError!SourceForms {
69 self.last_error_span = null;
70 return scanSource(allocator, source, &self.last_error_span);
71 }
72
73 pub fn executeSource(self: *ToplevelContext, allocator: Allocator, source: []const u8) ToplevelError!Execution {
74 var parsed = try parseSourceForms(self, allocator, source);
75 errdefer parsed.deinit();
76
77 const prebound_defs = try source_owner.prebindSexprDefs(self, source);
78 defer if (prebound_defs.len > 0) self.allocator.free(prebound_defs);
79 errdefer source_owner.rollbackPreboundDefs(self, prebound_defs);
80
81 var queries: std.ArrayListUnmanaged(QueryRecord) = .empty;
82 errdefer {
83 for (queries.items) |*query| {
84 query.result.deinit();
85 }
86 queries.deinit(allocator);
87 }
88
89 for (parsed.items, 0..) |form, form_index| {
90 const maybe_result = forms_owner.processFormSexpr(self, form.source) catch |err| {
91 self.last_error_span = form.span;
92 return err;
93 };
94
95 if (maybe_result) |result| {
96 var query_result = result;
97 query_result.source_location = form.location;
98 queries.append(allocator, .{
99 .form_index = form_index,
100 .form = form,
101 .result = query_result,
102 }) catch |err| {
103 query_result.deinit();
104 return switch (err) {
105 error.OutOfMemory => ToplevelError.OutOfMemory,
106 };
107 };
108 }
109 }
110
111 return .{
112 .allocator = allocator,
113 .forms = parsed.items,
114 .queries = queries.toOwnedSlice(allocator) catch return ToplevelError.OutOfMemory,
115 };
116 }
117
118 fn scanSource(allocator: Allocator, source: []const u8, error_span: ?*?SourceSpan) ToplevelError!SourceForms {
119 var forms: std.ArrayListUnmanaged(SourceForm) = .empty;
120 errdefer forms.deinit(allocator);
121
122 var pos: usize = 0;
123 while (pos < source.len) {
124 pos = source_owner.skipWhitespaceAndComments(source, pos);
125 if (pos >= source.len) break;
126
127 const form_start = pos;
128 const form_end = source_owner.findFormEnd(source, pos) orelse {
129 if (error_span) |out| out.* = source_owner.spanFromOffsets(source, form_start, source.len);
130 return ToplevelError.ParseError;
131 };
132 pos = form_end;
133
134 const form_source = source[form_start..form_end];
135 if (form_source.len == 0) continue;
136
137 const tokens = pexpr.tokenize(allocator, form_source) catch return ToplevelError.OutOfMemory;
138 defer pexpr.freeTokens(allocator, tokens);
139
140 if (tokens.len == 0) continue;
141
142 const span = source_owner.spanFromOffsets(source, form_start, form_end);
143 try forms.append(allocator, .{
144 .kind = classify(tokens),
145 .source = form_source,
146 .name = formName(tokens),
147 .span = span,
148 .location = source_owner.locationFromSpan(span),
149 });
150 }
151
152 return .{
153 .allocator = allocator,
154 .items = forms.toOwnedSlice(allocator) catch return ToplevelError.OutOfMemory,
155 };
156 }
157
158 fn classify(tokens: []const pexpr.Token) FormKind {
159 if (tokens.len >= 2 and std.mem.eql(u8, tokens[0], "(")) {
160 if (std.mem.eql(u8, tokens[1], "define")) return .definition;
161 if (std.mem.eql(u8, tokens[1], "define-type")) return .type_definition;
162 if (std.mem.eql(u8, tokens[1], "query")) return .query;
163 }
164 return .expression;
165 }
166
167 fn formName(tokens: []const pexpr.Token) ?[]const u8 {
168 return switch (classify(tokens)) {
169 .definition => definitionName(tokens),
170 .type_definition => if (tokens.len >= 3 and source_owner.isIdentifier(tokens[2])) tokens[2] else null,
171 .query => queryName(tokens),
172 .expression => null,
173 };
174 }
175
176 fn definitionName(tokens: []const pexpr.Token) ?[]const u8 {
177 if (tokens.len < 4) return null;
178 if (std.mem.eql(u8, tokens[2], "(")) {
179 if (tokens.len >= 4 and source_owner.isIdentifier(tokens[3])) return tokens[3];
180 return null;
181 }
182 if (source_owner.isIdentifier(tokens[2])) return tokens[2];
183 return null;
184 }
185
186 fn queryName(tokens: []const pexpr.Token) ?[]const u8 {
187 if (tokens.len <= 4) return null;
188 if (source_owner.isIdentifier(tokens[2]) and !std.mem.eql(u8, tokens[3], ")")) return tokens[2];
189 return null;
190 }
191
192 test "parseSource classifies top-level forms" {
193 const source =
194 \\(define p 0.5)
195 \\(define-type coin (Heads) (Tails))
196 \\(query coin-query (Marginal (flip p)))
197 \\(Marginal True)
198 ;
199
200 var parsed = try parseSource(std.testing.allocator, source);
201 defer parsed.deinit();
202
203 try std.testing.expectEqual(@as(usize, 4), parsed.items.len);
204 try std.testing.expectEqual(FormKind.definition, parsed.items[0].kind);
205 try std.testing.expectEqualStrings("p", parsed.items[0].name.?);
206 try std.testing.expectEqual(FormKind.type_definition, parsed.items[1].kind);
207 try std.testing.expectEqualStrings("coin", parsed.items[1].name.?);
208 try std.testing.expectEqual(FormKind.query, parsed.items[2].kind);
209 try std.testing.expectEqualStrings("coin-query", parsed.items[2].name.?);
210 try std.testing.expectEqual(FormKind.expression, parsed.items[3].kind);
211 try std.testing.expect(parsed.items[3].name == null);
212 }
213
214 test "parseSourceForms records unmatched form span" {
215 var ctx = try ToplevelContext.init(std.testing.allocator);
216 defer ctx.deinit();
217
218 try std.testing.expectError(ToplevelError.ParseError, ctx.parseSourceForms(std.testing.allocator, "(Marginal True"));
219 const span = ctx.lastErrorSpan().?;
220 try std.testing.expectEqual(@as(usize, 0), span.start.offset);
221 try std.testing.expectEqual(@as(usize, 14), span.end.offset);
222 }
223
224 test "executeSource collects query records" {
225 var ctx = try ToplevelContext.init(std.testing.allocator);
226 defer ctx.deinit();
227
228 const source =
229 \\(define p 0.5)
230 \\(Marginal (flip p))
231 ;
232
233 var execution = try ctx.executeSource(std.testing.allocator, source);
234 defer execution.deinit();
235
236 try std.testing.expectEqual(@as(usize, 2), execution.forms.len);
237 try std.testing.expectEqual(@as(usize, 1), execution.queries.len);
238 try std.testing.expectEqual(@as(usize, 1), execution.queries[0].form_index);
239 try std.testing.expectEqual(FormKind.expression, execution.queries[0].form.kind);
240 try std.testing.expectEqual(@as(u32, 1), execution.queries[0].result.source_location.?.start_line);
241
242 var found_true = false;
243 for (execution.queries[0].result.outcomes) |outcome| {
244 if (std.mem.eql(u8, outcome.value_str, "True")) {
245 found_true = true;
246 try std.testing.expectApproxEqAbs(@as(f64, 0.5), outcome.probability, 1e-9);
247 }
248 }
249 try std.testing.expect(found_true);
250 }