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 }