lib/pluck/src/main.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const pretty = @import("pretty");
  3 const pretty_usage = @import("pretty_usage");
  4 const sys = @import("sys");
  5 
  6 pub const bdd = @import("bdd.zig");
  7 pub const Manager = bdd.Manager;
  8 pub const Bdd = bdd.Bdd;
  9 pub const VarLabel = bdd.VarLabel;
 10 
 11 pub const weight_dd = @import("weight.zig");
 12 pub const WeightDD = weight_dd.WeightDD;
 13 pub const WeightNode = weight_dd.WeightNode;
 14 pub const Weight = weight_dd.Weight;
 15 pub const GuardedWeight = weight_dd.GuardedWeight;
 16 
 17 pub const pexpr = @import("pexpr.zig");
 18 pub const PExpr = pexpr.PExpr;
 19 pub const Head = pexpr.Head;
 20 pub const ConstructorDef = pexpr.ConstructorDef;
 21 pub const TypeRegistry = pexpr.TypeRegistry;
 22 pub const Definitions = pexpr.Definitions;
 23 pub const parseExpr = pexpr.parseExpr;
 24 
 25 pub const runtime = @import("runtime.zig");
 26 pub const Env = runtime.Env;
 27 pub const RuntimeValue = runtime.RuntimeValue;
 28 pub const Closure = runtime.Closure;
 29 pub const LazyKCThunk = runtime.LazyKCThunk;
 30 pub const LazyKCThunkUnion = runtime.LazyKCThunkUnion;
 31 pub const LazyEnumeratorThunk = runtime.LazyEnumeratorThunk;
 32 pub const StateVars = runtime.StateVars;
 33 
 34 pub const evaluator = @import("evaluator.zig");
 35 pub const LazyKCState = evaluator.LazyKCState;
 36 pub const LazyKCConfig = evaluator.LazyKCConfig;
 37 pub const LazyKCStats = evaluator.LazyKCStats;
 38 pub const compile = evaluator.compile;
 39 pub const CompileResult = evaluator.CompileResult;
 40 pub const WeightedResult = evaluator.WeightedResult;
 41 
 42 pub const toplevel = @import("root.zig").toplevel;
 43 pub const limits = @import("root.zig").limits;
 44 pub const ToplevelContext = toplevel.ToplevelContext;
 45 pub const ToplevelConfig = toplevel.ToplevelConfig;
 46 pub const QueryResult = toplevel.QueryResult;
 47 
 48 pub const cli = @import("cli/root.zig");
 49 
 50 pub const query_context = @import("query.zig");
 51 pub const QueryContext = query_context.QueryContext;
 52 pub const RunContext = query_context.RunContext;
 53 pub const SharedContext = query_context.SharedContext;
 54 
 55 pub const time = @import("time.zig");
 56 
 57 test "pluck package namespace" {
 58     std.testing.refAllDecls(@This());
 59     _ = bdd;
 60     _ = weight_dd;
 61     _ = pexpr;
 62     _ = runtime;
 63     _ = evaluator;
 64     _ = toplevel;
 65     _ = cli;
 66     _ = query_context;
 67 }
 68 
 69 const CLIOptions = cli.CLIOptions;
 70 const ExitCode = cli.ExitCode;
 71 
 72 pub fn main(init: sys.process.Init) u8 {
 73     return mainWithExitCode(init.minimal.args) catch |err| {
 74         writeStderrErrorFmt("fatal error: {s}", .{@errorName(err)});
 75         return @backingInt(ExitCode.@"error");
 76     };
 77 }
 78 
 79 fn mainWithExitCode(process_args: sys.process.Args) !u8 {
 80     var gpa: std.heap.DebugAllocator(.{}) = .init;
 81     defer _ = gpa.deinit();
 82     const allocator = gpa.allocator();
 83 
 84     const stdout_file = sys.stdio.stdout();
 85     const stdout_options = cli.prettyOptions(stdout_file);
 86     var stdout_buf: [8192]u8 = undefined;
 87     var stdout_writer = stdout_file.writer(sys.stdio.debugIo(), &stdout_buf);
 88     const stdout = &stdout_writer.interface;
 89     defer stdout.flush() catch {};
 90 
 91     var opts = try cli.parseArgs(allocator, process_args);
 92     defer cli.freeOptions(allocator, &opts);
 93 
 94     if (opts.version) {
 95         try writePrettyFmt(allocator, stdout, "pluck {s}\n", .{cli.version});
 96         return @backingInt(ExitCode.success);
 97     }
 98 
 99     if (opts.help) {
100         if (opts.parse_error) {
101             try cli.printParseError(allocator, stdout, &opts, stdout_options);
102             return @backingInt(ExitCode.@"error");
103         }
104         if (opts.help_topic) |topic| {
105             if (try cli.printHelpTopic(allocator, stdout, topic, stdout_options)) {
106                 return @backingInt(ExitCode.success);
107             }
108             try cli.printUnknownHelpTopic(allocator, stdout, topic, stdout_options);
109             return @backingInt(ExitCode.@"error");
110         }
111         try cli.printHelp(allocator, stdout, stdout_options);
112         return @backingInt(ExitCode.success);
113     }
114 
115     var ctx_obj = ToplevelContext.init(allocator) catch {
116         writeStderrErrorFmt("failed to initialize toplevel context", .{});
117         return @backingInt(ExitCode.@"error");
118     };
119     defer ctx_obj.deinit();
120     const ctx = &ctx_obj;
121 
122     ctx.setConfig(.{
123         .time_limit = opts.time_limit,
124         .max_depth = opts.max_depth,
125         .ite_limit = opts.ite_limit,
126         .sample_after_max_depth = opts.sample_after_max_depth,
127         .verbose = opts.verbose and !opts.silent,
128         .parallel_wmc = opts.parallel,
129         .rng_seed = opts.rng_seed,
130         .fallback_mode = opts.fallback_mode,
131         .fallback_lpsmc_k = opts.fallback_lpsmc_k,
132         .lpsmc_adaptive_k = opts.lpsmc_adaptive_k,
133         .lpsmc_adaptive_k_max = opts.lpsmc_adaptive_k_max,
134         .lpsmc_workers = opts.lpsmc_workers,
135         .factor_max_branches = opts.factor_max_branches,
136         .weight_dd_max_nodes = opts.weight_dd_max_nodes,
137         .use_strict_order = opts.use_strict_order,
138         .use_reverse_order = opts.use_reverse_order,
139         .var_order_fallback = opts.var_order_fallback,
140         .definition_order_mode = opts.definition_order_mode,
141     });
142 
143     var exit_code = ExitCode.success;
144 
145     if (opts.eval_expr) |expr| {
146         const result = ctx.processForm(expr) catch |err| {
147             if (opts.json) {
148                 const json_err = toplevel.JsonError{
149                     .type = "eval_error",
150                     .message = @errorName(err),
151                     .span = ctx.lastErrorSpan(),
152                 };
153                 try json_err.write(stdout);
154             } else {
155                 if (err == error.ParseError) {
156                     try ctx.formatError(stdout);
157                 } else {
158                     try writePrettyFmt(allocator, stdout, "Error: {any}\n", .{err});
159                 }
160             }
161             return @backingInt(ExitCode.@"error");
162         };
163         if (result) |*query_result| {
164             defer @constCast(query_result).deinit();
165             exit_code = ExitCode.combine(exit_code, ExitCode.fromQueryResult(query_result));
166             if (opts.json) {
167                 try toplevel.writeQueryResult(query_result, stdout);
168             } else {
169                 try query_result.print(stdout);
170             }
171         }
172         return @backingInt(exit_code);
173     }
174 
175     const has_stdin_file = for (opts.files.items) |f| {
176         if (std.mem.eql(u8, f, "-")) break true;
177     } else false;
178     if (has_stdin_file or (!cli.isStdinTty() and opts.files.items.len == 0)) {
179         var stdin_buf: [8192]u8 = undefined;
180         var stdin_reader = sys.stdio.stdin().readerStreaming(sys.stdio.debugIo(), &stdin_buf);
181         const input = stdin_reader.interface.allocRemaining(allocator, .limited(10 * 1024 * 1024)) catch |err| {
182             if (opts.json) {
183                 const json_err = toplevel.JsonError{
184                     .type = "io_error",
185                     .message = @errorName(err),
186                 };
187                 try json_err.write(stdout);
188             } else {
189                 try writePrettyFmt(allocator, stdout, "Error reading stdin: {any}\n", .{err});
190             }
191             return @backingInt(ExitCode.@"error");
192         };
193         defer allocator.free(input);
194 
195         if (opts.json) {
196             const CallbackState = struct {
197                 exit_code: ExitCode = .success,
198                 writer_ptr: *const anyopaque,
199                 write_fn: *const fn (*const anyopaque, *toplevel.QueryResult) anyerror!void,
200 
201                 fn callback(result: *toplevel.QueryResult, ctx_ptr: *anyopaque) void {
202                     const self: *@This() = @ptrCast(@alignCast(ctx_ptr));
203                     defer result.deinit();
204                     self.exit_code = ExitCode.combine(self.exit_code, ExitCode.fromQueryResult(result));
205                     const writer_typed: *@TypeOf(stdout) = @ptrCast(@alignCast(@constCast(self.writer_ptr)));
206                     toplevel.writeQueryResult(result, writer_typed.*) catch {};
207                 }
208             };
209 
210             var state = CallbackState{ .writer_ptr = @ptrCast(&stdout), .write_fn = undefined };
211             ctx.processSourceWithCallback(input, CallbackState.callback, @ptrCast(&state)) catch |err| {
212                 const json_err = toplevel.JsonError{
213                     .type = "parse_error",
214                     .message = @errorName(err),
215                     .span = ctx.lastErrorSpan(),
216                 };
217                 try json_err.write(stdout);
218                 return @backingInt(ExitCode.@"error");
219             };
220             exit_code = state.exit_code;
221         } else {
222             const TextCallbackState = struct {
223                 exit_code: ExitCode = .success,
224                 writer_ptr: *const anyopaque,
225 
226                 fn callback(result: *toplevel.QueryResult, ctx_ptr: *anyopaque) void {
227                     const self: *@This() = @ptrCast(@alignCast(ctx_ptr));
228                     defer result.deinit();
229                     self.exit_code = ExitCode.combine(self.exit_code, ExitCode.fromQueryResult(result));
230                     const writer_typed: *@TypeOf(stdout) = @ptrCast(@alignCast(@constCast(self.writer_ptr)));
231                     result.print(writer_typed.*) catch {};
232                 }
233             };
234 
235             var state = TextCallbackState{ .writer_ptr = @ptrCast(&stdout) };
236             ctx.processSourceWithCallback(input, TextCallbackState.callback, @ptrCast(&state)) catch |err| {
237                 if (err == error.ParseError) {
238                     try ctx.formatError(stdout);
239                 } else {
240                     try writePrettyFmt(allocator, stdout, "Error: {any}\n", .{err});
241                 }
242                 return @backingInt(ExitCode.@"error");
243             };
244             exit_code = state.exit_code;
245         }
246         return @backingInt(exit_code);
247     }
248 
249     if (opts.files.items.len == 0) {
250         try cli.printHelp(allocator, stdout, stdout_options);
251         return @backingInt(ExitCode.success);
252     }
253 
254     for (opts.files.items) |file| {
255         if (std.mem.eql(u8, file, "-")) continue;
256         if (opts.verbose and !opts.silent) {
257             try writePrettyFmt(allocator, stdout, "Loading {s}...\n", .{file});
258             try stdout.flush();
259         }
260         if (opts.json) {
261             exit_code = ExitCode.combine(exit_code, try loadFileJson(ctx, file, stdout, allocator));
262         } else {
263             ctx.loadFileWithWriter(file, stdout) catch |err| {
264                 try writePrettyFmt(allocator, stdout, "Error loading {s}:\n", .{file});
265                 if (err == toplevel.ToplevelError.ParseError) {
266                     try ctx.formatError(stdout);
267                 } else {
268                     try writePrettyFmt(allocator, stdout, "  {any}\n", .{err});
269                 }
270                 return @backingInt(ExitCode.@"error");
271             };
272         }
273     }
274 
275     return @backingInt(exit_code);
276 }
277 
278 fn loadFileJson(ctx: *ToplevelContext, file: []const u8, writer: anytype, allocator: std.mem.Allocator) !ExitCode {
279     const content = sys.fs.readFileAlloc(allocator, file, 10 * 1024 * 1024) catch |err| {
280         const json_err = toplevel.JsonError{
281             .type = "io_error",
282             .message = @errorName(err),
283             .file = file,
284         };
285         try json_err.write(writer);
286         return ExitCode.@"error";
287     };
288     defer allocator.free(content);
289 
290     const CallbackState = struct {
291         exit_code: ExitCode = .success,
292         writer_ptr: @TypeOf(&writer),
293 
294         fn callback(result: *toplevel.QueryResult, ctx_ptr: *anyopaque) void {
295             const self: *@This() = @ptrCast(@alignCast(ctx_ptr));
296             defer result.deinit();
297             self.exit_code = ExitCode.combine(self.exit_code, ExitCode.fromQueryResult(result));
298             toplevel.writeQueryResult(result, self.writer_ptr.*) catch {};
299         }
300     };
301 
302     var state = CallbackState{ .writer_ptr = &writer };
303 
304     ctx.processSourceWithCallback(content, CallbackState.callback, @ptrCast(&state)) catch |err| {
305         const json_err = toplevel.JsonError{
306             .type = "parse_error",
307             .message = @errorName(err),
308             .file = file,
309             .span = ctx.lastErrorSpan(),
310         };
311         try json_err.write(writer);
312         return ExitCode.@"error";
313     };
314 
315     return state.exit_code;
316 }
317 
318 fn writePrettyFmt(allocator: std.mem.Allocator, writer: *std.Io.Writer, comptime fmt: []const u8, args: anytype) !void {
319     const text = try std.fmt.allocPrint(allocator, fmt, args);
320     defer allocator.free(text);
321     try writePrettyText(writer, text);
322 }
323 
324 fn writePrettyText(writer: *std.Io.Writer, text: []const u8) !void {
325     try pretty.write(writer, .{ .text = text }, .{});
326 }
327 
328 fn writeStderrErrorFmt(comptime fmt: []const u8, args: anytype) void {
329     var buffer: [4096]u8 = undefined;
330     var fixed_allocator = std.heap.FixedBufferAllocator.init(&buffer);
331     pretty_usage.writeFileErrorTextFmt(
332         fixed_allocator.allocator(),
333         sys.stdio.stderr(),
334         "pluck",
335         fmt,
336         args,
337         .{},
338     ) catch {};
339 }