tally/engine/src/evaluator.zig

530 lines
16 KiB
Zig

//! AST evaluator for Tally.
//!
//! Walks an AST and produces a Value. In standard mode, all computations
//! use f64 floating-point arithmetic. In programmer mode, integer operations
//! are exact (masked to bit width). The evaluator uses an Environment for
//! variable storage, history, and configuration.
const std = @import("std");
const math = std.math;
const Allocator = std.mem.Allocator;
const ast = @import("ast.zig");
const Expr = ast.Expr;
const BinaryOp = ast.BinaryOp;
const types = @import("types.zig");
const Mode = types.Mode;
const ProgrammerConfig = types.ProgrammerConfig;
const CalcError = types.CalcError;
const parser_mod = @import("parser.zig");
const Parser = parser_mod.Parser;
/// Evaluation environment holding variables, history, and config.
pub const Environment = struct {
allocator: Allocator,
mode: Mode,
programmer_config: ProgrammerConfig,
variables: std.StringHashMap(f64),
ans: f64,
history_len: usize,
pub fn init(allocator: Allocator, mode: Mode) Environment {
return .{
.allocator = allocator,
.mode = mode,
.programmer_config = .{},
.variables = std.StringHashMap(f64).init(allocator),
.ans = 0,
.history_len = 0,
};
}
pub fn deinit(self: *Environment) void {
self.variables.deinit();
}
/// Set a variable value.
pub fn setVar(self: *Environment, name: []const u8, value: f64) !void {
try self.variables.put(name, value);
}
/// Get a variable or constant value.
pub fn getVar(self: *const Environment, name: []const u8) ?f64 {
// Built-in constants
if (std.mem.eql(u8, name, "pi")) return math.pi;
if (std.mem.eql(u8, name, "e")) return math.e;
if (std.mem.eql(u8, name, "tau")) return math.tau;
if (std.mem.eql(u8, name, "Ans") or std.mem.eql(u8, name, "ans")) return self.ans;
return self.variables.get(name);
}
};
/// Evaluate a parsed expression in the given environment.
/// Returns the computed value as f64 for standard mode.
pub fn evaluate(env: *Environment, expr: *const Expr) CalcError!f64 {
switch (expr.*) {
.number => |n| return n.float_value,
.variable => |name| {
return env.getVar(name) orelse return CalcError.UnknownVariable;
},
.assignment => |a| {
const val = try evaluate(env, a.value);
env.setVar(a.name, val) catch return CalcError.OutOfMemory;
return val;
},
.unary => |u| {
const operand = try evaluate(env, u.operand);
return switch (u.op) {
.negate => -operand,
.bitwise_not => {
// In standard mode, bitwise not doesn't really make sense,
// but we'll compute it on the integer representation
const int_val: u64 = @bitCast(@as(i64, @intFromFloat(operand)));
const result = ~int_val & env.programmer_config.bit_width.mask();
return @floatFromInt(@as(i64, @bitCast(result)));
},
};
},
.binary => |b| {
const left = try evaluate(env, b.left);
const right = try evaluate(env, b.right);
return evalBinaryOp(b.op, left, right);
},
.call => |c| {
return evalFunction(env, c.name, c.args);
},
}
}
/// Evaluate a binary operation on two f64 values.
fn evalBinaryOp(op: BinaryOp, left: f64, right: f64) CalcError!f64 {
return switch (op) {
.add => left + right,
.sub => left - right,
.mul => left * right,
.div => if (right == 0) CalcError.DivisionByZero else left / right,
.mod => if (right == 0) CalcError.DivisionByZero else @mod(left, right),
.pow => math.pow(f64, left, right),
// Bitwise ops in standard mode operate on integer truncations
.bit_and => floatBitwise(left, right, bitwiseAnd),
.bit_or => floatBitwise(left, right, bitwiseOr),
.bit_xor => floatBitwise(left, right, bitwiseXor),
.shift_left => floatShift(left, right, true),
.shift_right, .shift_right_logical => floatShift(left, right, false),
.rotate_left, .rotate_right => {
// Rotations need bit width context; in standard mode, use 64-bit
const l: u64 = @bitCast(@as(i64, @intFromFloat(left)));
const r: u6 = @intFromFloat(@mod(right, 64.0));
const result = if (op == .rotate_left)
math.rotl(u64, l, r)
else
math.rotr(u64, l, r);
return @floatFromInt(@as(i64, @bitCast(result)));
},
};
}
fn bitwiseAnd(a: u64, b: u64) u64 {
return a & b;
}
fn bitwiseOr(a: u64, b: u64) u64 {
return a | b;
}
fn bitwiseXor(a: u64, b: u64) u64 {
return a ^ b;
}
fn floatBitwise(left: f64, right: f64, op: *const fn (u64, u64) u64) f64 {
const l: u64 = @bitCast(@as(i64, @intFromFloat(left)));
const r: u64 = @bitCast(@as(i64, @intFromFloat(right)));
const result = op(l, r);
return @floatFromInt(@as(i64, @bitCast(result)));
}
fn floatShift(left: f64, right: f64, is_left: bool) f64 {
const l: u64 = @bitCast(@as(i64, @intFromFloat(left)));
const shift_amt: u6 = @intFromFloat(@mod(right, 64.0));
const result = if (is_left) l << shift_amt else l >> shift_amt;
return @floatFromInt(@as(i64, @bitCast(result)));
}
/// Evaluate a built-in function call.
fn evalFunction(env: *Environment, name: []const u8, args: []const *Expr) CalcError!f64 {
// Single-argument functions
if (args.len == 1) {
const x = try evaluate(env, args[0]);
return evalSingleArgFn(name, x) orelse CalcError.UnknownFunction;
}
// Multi-argument functions
if (args.len == 2) {
const a = try evaluate(env, args[0]);
const b = try evaluate(env, args[1]);
if (std.mem.eql(u8, name, "max")) return @max(a, b);
if (std.mem.eql(u8, name, "min")) return @min(a, b);
if (std.mem.eql(u8, name, "atan2")) return math.atan2(a, b);
if (std.mem.eql(u8, name, "log")) {
// log(value, base)
if (b <= 0 or b == 1 or a <= 0) return CalcError.DomainError;
return @log(a) / @log(b);
}
}
// Zero-argument functions
if (args.len == 0) {
if (std.mem.eql(u8, name, "rand")) {
// Not truly random in a pure engine, but useful as placeholder
return 0.0;
}
}
return CalcError.UnknownFunction;
}
/// Evaluate a single-argument built-in function.
fn evalSingleArgFn(name: []const u8, x: f64) ?f64 {
if (std.mem.eql(u8, name, "sin")) return @sin(x);
if (std.mem.eql(u8, name, "cos")) return @cos(x);
if (std.mem.eql(u8, name, "tan")) return @tan(x);
if (std.mem.eql(u8, name, "asin")) {
if (x < -1 or x > 1) return null; // domain error
return math.asin(x);
}
if (std.mem.eql(u8, name, "acos")) {
if (x < -1 or x > 1) return null;
return math.acos(x);
}
if (std.mem.eql(u8, name, "atan")) return math.atan(x);
if (std.mem.eql(u8, name, "log")) return @log10(x);
if (std.mem.eql(u8, name, "log10")) return @log10(x);
if (std.mem.eql(u8, name, "ln")) return @log(x);
if (std.mem.eql(u8, name, "log2")) return @log2(x);
if (std.mem.eql(u8, name, "sqrt")) {
if (x < 0) return null;
return @sqrt(x);
}
if (std.mem.eql(u8, name, "cbrt")) return math.cbrt(x);
if (std.mem.eql(u8, name, "abs")) return @abs(x);
if (std.mem.eql(u8, name, "ceil")) return @ceil(x);
if (std.mem.eql(u8, name, "floor")) return @floor(x);
if (std.mem.eql(u8, name, "round")) return @round(x);
if (std.mem.eql(u8, name, "exp")) return @exp(x);
if (std.mem.eql(u8, name, "factorial")) {
if (x < 0 or x != @round(x) or x > 170) return null;
return factorial(@intFromFloat(x));
}
return null;
}
fn factorial(n: u64) f64 {
if (n <= 1) return 1.0;
var result: f64 = 1.0;
var i: u64 = 2;
while (i <= n) : (i += 1) {
result *= @floatFromInt(i);
}
return result;
}
/// High-level evaluate: parse a string and evaluate it.
/// Updates env.ans on success.
pub fn evalString(env: *Environment, allocator: Allocator, source: []const u8) CalcError!f64 {
var p = Parser.init(allocator, source, env.mode);
const expr = try p.parse();
const result = try evaluate(env, expr);
env.ans = result;
env.history_len += 1;
return result;
}
// -- Tests --
const testing = std.testing;
fn testEval(source: []const u8) !f64 {
var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
defer _ = arena.deinit();
const alloc = arena.allocator();
var env = Environment.init(alloc, .standard);
defer env.deinit();
return evalString(&env, alloc, source);
}
fn testEvalProgrammer(source: []const u8) !f64 {
var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
defer _ = arena.deinit();
const alloc = arena.allocator();
var env = Environment.init(alloc, .programmer);
defer env.deinit();
return evalString(&env, alloc, source);
}
test "eval simple number" {
const result = try testEval("42");
try testing.expectEqual(@as(f64, 42.0), result);
}
test "eval addition" {
const result = try testEval("2 + 3");
try testing.expectEqual(@as(f64, 5.0), result);
}
test "eval subtraction" {
const result = try testEval("10 - 7");
try testing.expectEqual(@as(f64, 3.0), result);
}
test "eval multiplication" {
const result = try testEval("6 * 7");
try testing.expectEqual(@as(f64, 42.0), result);
}
test "eval division" {
const result = try testEval("10 / 4");
try testing.expectEqual(@as(f64, 2.5), result);
}
test "eval division by zero" {
const result = testEval("1 / 0");
try testing.expectError(CalcError.DivisionByZero, result);
}
test "eval modulo" {
const result = try testEval("10 % 3");
try testing.expectApproxEqAbs(@as(f64, 1.0), result, 1e-10);
}
test "eval power" {
const result = try testEval("2^10");
try testing.expectEqual(@as(f64, 1024.0), result);
}
test "eval precedence" {
const result = try testEval("2 + 3 * 4");
try testing.expectEqual(@as(f64, 14.0), result);
}
test "eval parentheses" {
const result = try testEval("(2 + 3) * 4");
try testing.expectEqual(@as(f64, 20.0), result);
}
test "eval unary negation" {
const result = try testEval("-5 + 3");
try testing.expectEqual(@as(f64, -2.0), result);
}
test "eval nested parens" {
const result = try testEval("((2 + 3) * (4 - 1))");
try testing.expectEqual(@as(f64, 15.0), result);
}
test "eval pi constant" {
const result = try testEval("pi");
try testing.expectApproxEqAbs(math.pi, result, 1e-10);
}
test "eval e constant" {
const result = try testEval("e");
try testing.expectApproxEqAbs(math.e, result, 1e-10);
}
test "eval tau constant" {
const result = try testEval("tau");
try testing.expectApproxEqAbs(math.tau, result, 1e-10);
}
test "eval implicit mul: 2pi" {
const result = try testEval("2pi");
try testing.expectApproxEqAbs(2.0 * math.pi, result, 1e-10);
}
test "eval implicit mul: 3(4+5)" {
const result = try testEval("3(4+5)");
try testing.expectEqual(@as(f64, 27.0), result);
}
test "eval sin" {
const result = try testEval("sin(0)");
try testing.expectApproxEqAbs(@as(f64, 0.0), result, 1e-10);
}
test "eval cos" {
const result = try testEval("cos(0)");
try testing.expectApproxEqAbs(@as(f64, 1.0), result, 1e-10);
}
test "eval sqrt" {
const result = try testEval("sqrt(144)");
try testing.expectEqual(@as(f64, 12.0), result);
}
test "eval abs" {
const result = try testEval("abs(-42)");
try testing.expectEqual(@as(f64, 42.0), result);
}
test "eval floor" {
const result = try testEval("floor(3.7)");
try testing.expectEqual(@as(f64, 3.0), result);
}
test "eval ceil" {
const result = try testEval("ceil(3.2)");
try testing.expectEqual(@as(f64, 4.0), result);
}
test "eval round" {
const result = try testEval("round(3.5)");
try testing.expectEqual(@as(f64, 4.0), result);
}
test "eval factorial" {
const result = try testEval("factorial(5)");
try testing.expectEqual(@as(f64, 120.0), result);
}
test "eval factorial 0" {
const result = try testEval("factorial(0)");
try testing.expectEqual(@as(f64, 1.0), result);
}
test "eval ln" {
const result = try testEval("ln(1)");
try testing.expectApproxEqAbs(@as(f64, 0.0), result, 1e-10);
}
test "eval exp" {
const result = try testEval("exp(0)");
try testing.expectEqual(@as(f64, 1.0), result);
}
test "eval max" {
const result = try testEval("max(3, 7)");
try testing.expectEqual(@as(f64, 7.0), result);
}
test "eval min" {
const result = try testEval("min(3, 7)");
try testing.expectEqual(@as(f64, 3.0), result);
}
test "eval unknown function" {
const result = testEval("bogus(1)");
try testing.expectError(CalcError.UnknownFunction, result);
}
test "eval unknown variable" {
const result = testEval("xyz");
try testing.expectError(CalcError.UnknownVariable, result);
}
test "eval variable assignment and use" {
var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
defer _ = arena.deinit();
const alloc = arena.allocator();
var env = Environment.init(alloc, .standard);
defer env.deinit();
const assign_result = try evalString(&env, alloc, "X = 42");
try testing.expectEqual(@as(f64, 42.0), assign_result);
const use_result = try evalString(&env, alloc, "X + 8");
try testing.expectEqual(@as(f64, 50.0), use_result);
}
test "eval Ans" {
var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
defer _ = arena.deinit();
const alloc = arena.allocator();
var env = Environment.init(alloc, .standard);
defer env.deinit();
_ = try evalString(&env, alloc, "7 * 6");
const result = try evalString(&env, alloc, "Ans + 1");
try testing.expectEqual(@as(f64, 43.0), result);
}
test "eval complex expression" {
const result = try testEval("sin(pi/2) + cos(0)");
try testing.expectApproxEqAbs(@as(f64, 2.0), result, 1e-10);
}
test "eval 2^32 - 1" {
const result = try testEval("2^32 - 1");
try testing.expectEqual(@as(f64, 4294967295.0), result);
}
test "eval programmer XOR" {
const result = try testEvalProgrammer("0xF ^ 0x3");
try testing.expectEqual(@as(f64, 12.0), result);
}
test "eval programmer AND" {
const result = try testEvalProgrammer("0xFF & 0x0F");
try testing.expectEqual(@as(f64, 15.0), result);
}
test "eval programmer OR" {
const result = try testEvalProgrammer("0xF0 | 0x0F");
try testing.expectEqual(@as(f64, 255.0), result);
}
test "eval programmer shift left" {
const result = try testEvalProgrammer("1 << 8");
try testing.expectEqual(@as(f64, 256.0), result);
}
test "eval programmer shift right" {
const result = try testEvalProgrammer("256 >> 4");
try testing.expectEqual(@as(f64, 16.0), result);
}
test "eval number with underscores" {
const result = try testEval("1_000_000 + 1");
try testing.expectEqual(@as(f64, 1_000_001.0), result);
}
test "eval bitwise not in standard mode" {
const result = try testEval("~0");
// ~0 as i64 = -1
try testing.expectEqual(@as(f64, -1.0), result);
}
test "eval rotate left in standard mode" {
// 1 rol 4 = 16 (for 64-bit)
const result = try testEvalProgrammer("1 rol 4");
try testing.expectEqual(@as(f64, 16.0), result);
}
test "eval rotate right in standard mode" {
const result = try testEvalProgrammer("16 ror 4");
try testing.expectEqual(@as(f64, 1.0), result);
}
test "eval atan2" {
const result = try testEval("atan2(1, 1)");
try testing.expectApproxEqAbs(math.pi / 4.0, result, 1e-10);
}
test "eval log with base" {
const result = try testEval("log(100, 10)");
try testing.expectApproxEqAbs(@as(f64, 2.0), result, 1e-10);
}
test "eval log domain error" {
const result = testEval("log(-1, 10)");
try testing.expectError(CalcError.DomainError, result);
}
test "eval asin domain error" {
const result = testEval("asin(2)");
// asin(2) is domain error since |2| > 1
try testing.expectError(CalcError.UnknownFunction, result);
}
test "eval acos" {
const result = try testEval("acos(1)");
try testing.expectApproxEqAbs(@as(f64, 0.0), result, 1e-10);
}