use crate::ast::{BinaryOpKind as Op, Expr, Statement};
use crate::std as vexity_std;
use serde::{
de::{self, Deserialize, Deserializer, MapAccess, SeqAccess, Visitor},
ser::{Serialize, SerializeMap, SerializeSeq, Serializer},
};
use std::{collections::HashMap, fmt, sync::Arc};
#[derive(Debug, Clone)]
pub enum VarValue {
Bool(bool),
String(String),
Int(i32),
Float(f64),
Array(Vec<VarValue>),
HashMap(HashMap<String, VarValue>),
Lambda { param: String, body: Vec<Statement> },
}
impl fmt::Display for VarValue {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
VarValue::Bool(b) => write!(f, "{b}"),
VarValue::String(s) => write!(f, "{s}"),
VarValue::Int(v) => write!(f, "{v}"),
VarValue::Float(v) => write!(f, "{v}"),
VarValue::Array(a) => write!(f, "{a:?}"),
VarValue::HashMap(m) => {
write!(f, "{{")?;
for (i, (k, v)) in m.iter().enumerate() {
write!(f, "{k}: {v}")?;
if i + 1 < m.len() {
write!(f, ", ")?;
}
}
write!(f, "}}")
}
VarValue::Lambda { param, body } => write!(f, "|{param}| {body:?}"),
}
}
}
impl<'de> Deserialize<'de> for VarValue {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct VarValueVisitor;
impl<'de> Visitor<'de> for VarValueVisitor {
type Value = VarValue;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a valid VarValue")
}
fn visit_bool<E>(self, v: bool) -> Result<Self::Value, E> {
Ok(VarValue::Bool(v))
}
fn visit_i64<E>(self, v: i64) -> Result<Self::Value, E> {
Ok(VarValue::Int(v as i32))
}
fn visit_u64<E>(self, v: u64) -> Result<Self::Value, E> {
Ok(VarValue::Int(v as i32))
}
fn visit_f64<E>(self, v: f64) -> Result<Self::Value, E> {
Ok(VarValue::Float(v))
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
where
E: de::Error,
{
Ok(VarValue::String(v.to_string()))
}
fn visit_string<E>(self, v: String) -> Result<Self::Value, E> {
Ok(VarValue::String(v))
}
fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let mut vec = Vec::new();
while let Some(elem) = seq.next_element()? {
vec.push(elem);
}
Ok(VarValue::Array(vec))
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: MapAccess<'de>,
{
let mut obj = HashMap::new();
while let Some((key, value)) = map.next_entry()? {
obj.insert(key, value);
}
Ok(VarValue::HashMap(obj))
}
}
deserializer.deserialize_any(VarValueVisitor)
}
}
impl Serialize for VarValue {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
VarValue::Bool(b) => serializer.serialize_bool(*b),
VarValue::String(s) => serializer.serialize_str(s),
VarValue::Int(i) => serializer.serialize_i32(*i),
VarValue::Float(f) => serializer.serialize_f64(*f),
VarValue::Array(arr) => {
let mut seq = serializer.serialize_seq(Some(arr.len()))?;
for v in arr {
seq.serialize_element(v)?;
}
seq.end()
}
VarValue::HashMap(map) => {
let mut ser_map = serializer.serialize_map(Some(map.len()))?;
for (k, v) in map {
ser_map.serialize_entry(k, v)?;
}
ser_map.end()
}
VarValue::Lambda { .. } => Err(serde::ser::Error::custom(
"Cannot serialize VarValue::Lambda",
)),
}
}
}
pub enum Func {
Builtin(Arc<dyn Fn(Vec<VarValue>) -> VarValue + Send + Sync>),
Script {
params: Vec<String>,
body: Vec<Statement>,
},
}
impl Clone for Func {
fn clone(&self) -> Self {
match self {
Func::Builtin(f) => Func::Builtin(f.clone()),
Func::Script { params, body } => Func::Script {
params: params.clone(),
body: body.clone(),
},
}
}
}
pub struct Interpreter {
pub vars: HashMap<String, VarValue>,
pub funcs: HashMap<String, Func>,
}
impl Default for Interpreter {
fn default() -> Self {
Self::new()
}
}
impl Interpreter {
pub fn new() -> Self {
let funcs = HashMap::from([
("sma".into(), vexity_std::sma()),
("sma_batch".into(), vexity_std::sma_batch()),
("ema".into(), vexity_std::ema()),
("ema_batch".into(), vexity_std::ema_batch()),
("vwma".into(), vexity_std::vwma()),
("vwma_batch".into(), vexity_std::vwma_batch()),
("lsma".into(), vexity_std::lsma()),
("lsma_batch".into(), vexity_std::lsma_batch()),
]);
Interpreter {
vars: HashMap::new(),
funcs,
}
}
pub fn run(&mut self, stmts: Vec<Statement>) {
for stmt in stmts {
match stmt {
Statement::LetExpression { name, value } => {
let val = self.eval_expr(value);
self.vars.insert(name, val);
}
Statement::Expression { value } => {
self.eval_expr(value);
}
Statement::Print { name } => {
let v = self
.vars
.get(&name)
.unwrap_or_else(|| panic!("Variable `{name}` not found"));
println!("{v}");
}
Statement::Call { function, args } => {
let evaluated = args.into_iter().map(|e| self.eval_expr(e)).collect();
self.call_func(&function, evaluated);
}
Statement::FunctionDef { name, params, body } => {
self.funcs.insert(name, Func::Script { params, body });
}
}
}
}
fn eval_expr(&mut self, expr: Expr) -> VarValue {
match expr {
Expr::Boolean(b) => VarValue::Bool(b),
Expr::String(s) => VarValue::String(s),
Expr::Integer(i) => VarValue::Int(i),
Expr::Float(f) => VarValue::Float(f),
Expr::Array(items) => {
VarValue::Array(items.into_iter().map(|e| self.eval_expr(e)).collect())
}
Expr::HashMap(map) => VarValue::HashMap(
map.into_iter()
.map(|(k, v)| (k, self.eval_expr(v)))
.collect(),
),
Expr::Identifier(name) => self
.vars
.get(&name)
.cloned()
.unwrap_or_else(|| panic!("Variable `{name}` not found")),
Expr::Field { target, name } => {
let value = self.eval_expr(*target);
match (value, name.as_str()) {
(VarValue::Array(arr), "length") => VarValue::Int(arr.len() as i32),
(VarValue::HashMap(mut map), key) => map
.remove(key)
.unwrap_or_else(|| panic!("key `{key}` not found")),
(other, key) => {
panic!("Cannot access field `{key}` on non-object: {other:?}")
}
}
}
Expr::BinaryOp { lhs, op, rhs } => {
let l = self.eval_expr(*lhs);
let r = self.eval_expr(*rhs);
self.eval_binary(l, op, r)
}
Expr::Lambda { param, body } => VarValue::Lambda { param, body },
Expr::Block(stmts) => self.eval_statements_in_scope(&stmts, &mut self.vars.clone()),
Expr::Call { function, args } => {
let evaluated = args.into_iter().map(|e| self.eval_expr(e)).collect();
self.call_func(&function, evaluated)
}
Expr::MethodCall {
target,
method,
arg,
} => {
let tgt = self.eval_expr(*target);
let scope = self.vars.clone();
self.eval_method_call(tgt, &method, *arg, &scope)
}
Expr::Index { target, index } => {
let arr = self.eval_expr(*target);
let idx = self.eval_expr(*index);
match (arr, idx) {
(VarValue::Array(a), VarValue::Int(i)) => a
.get(i as usize)
.cloned()
.unwrap_or_else(|| panic!("Index {i} is out of bounds")),
_ => panic!("Indexing is only supported for arrays using integer indices"),
}
}
Expr::Range { start, end } => {
let start = self.eval_expr(*start);
let end = self.eval_expr(*end);
let (start, end) = match (start, end) {
(VarValue::Int(s), VarValue::Int(e)) => (s, e),
_ => panic!("Range bounds must be integers"),
};
let values = (start..end).map(VarValue::Int).collect();
VarValue::Array(values)
}
}
}
fn eval_binary(&self, l: VarValue, op: Op, r: VarValue) -> VarValue {
let (a, b) = match (l, r) {
(VarValue::Int(a), VarValue::Int(b)) => (a as f64, b as f64),
(VarValue::Int(a), VarValue::Float(b)) => (a as f64, b),
(VarValue::Float(a), VarValue::Int(b)) => (a, b as f64),
(VarValue::Float(a), VarValue::Float(b)) => (a, b),
(x, y) => panic!("Cannot apply {op:?} to {x:?} and {y:?}"),
};
let res = match op {
Op::Add => a + b,
Op::Sub => a - b,
Op::Mul => a * b,
Op::Div => a / b,
};
if res.fract() == 0.0 {
VarValue::Int(res as i32)
} else {
VarValue::Float(res)
}
}
pub fn call_func(&mut self, name: &str, args: Vec<VarValue>) -> VarValue {
let func = self
.funcs
.get(name)
.unwrap_or_else(|| panic!("unknown function `{name}`"))
.clone();
match func {
Func::Builtin(f) => f(args),
Func::Script { params, body } => {
if params.len() != args.len() {
panic!("{name} expects {} args, got {}", params.len(), args.len());
}
let mut scope = self.vars.clone();
for (p, v) in params.into_iter().zip(args) {
scope.insert(p, v);
}
self.eval_statements_in_scope(&body, &mut scope)
}
}
}
fn eval_statements_in_scope(
&mut self,
stmts: &[Statement],
scope: &mut HashMap<String, VarValue>,
) -> VarValue {
let mut last = VarValue::Float(0.0);
for stmt in stmts {
match stmt {
Statement::LetExpression { name, value } => {
let val = self.eval_expr_in_scope(value, scope);
scope.insert(name.clone(), val.clone());
self.vars.insert(name.clone(), val.clone());
last = val;
}
Statement::Expression { value } => {
last = self.eval_expr_in_scope(value, scope);
}
Statement::Call { function, args } => {
let evaluated = args
.iter()
.map(|e| self.eval_expr_in_scope(e, scope))
.collect();
last = self.call_func(function, evaluated);
}
Statement::Print { name } => {
let v = scope
.get(name)
.unwrap_or_else(|| panic!("Variable `{name}` not in scope"));
println!("{v}");
}
Statement::FunctionDef { name, params, body } => {
self.funcs.insert(
name.clone(),
Func::Script {
params: params.clone(),
body: body.clone(),
},
);
last = VarValue::Float(0.0);
}
}
}
last
}
fn eval_expr_in_scope(
&mut self,
expr: &Expr,
scope: &mut HashMap<String, VarValue>,
) -> VarValue {
match expr {
Expr::Boolean(b) => VarValue::Bool(*b),
Expr::String(s) => VarValue::String(s.clone()),
Expr::Integer(i) => VarValue::Int(*i),
Expr::Float(f) => VarValue::Float(*f),
Expr::Array(items) => VarValue::Array(
items
.iter()
.map(|e| self.eval_expr_in_scope(e, scope))
.collect(),
),
Expr::HashMap(m) => VarValue::HashMap(
m.iter()
.map(|(k, v)| (k.clone(), self.eval_expr_in_scope(v, scope)))
.collect(),
),
Expr::Identifier(name) => scope
.get(name)
.cloned()
.or_else(|| self.vars.get(name).cloned())
.unwrap_or_else(|| panic!("unknown identifier `{name}`")),
Expr::Field { target, name } => {
let obj = self.eval_expr_in_scope(target, scope);
match (obj, name.as_str()) {
(VarValue::Array(arr), "length") => VarValue::Int(arr.len() as i32),
(VarValue::HashMap(map), field) => map
.get(field)
.cloned()
.expect("invalid key: `{field}` for hashmap."),
(other, _) => panic!("invalid field target: {other:?}"),
}
}
Expr::BinaryOp { lhs, op, rhs } => {
let l = self.eval_expr_in_scope(lhs, scope);
let r = self.eval_expr_in_scope(rhs, scope);
self.eval_binary(l, op.clone(), r)
}
Expr::Lambda { param, body } => VarValue::Lambda {
param: param.clone(),
body: body.clone(),
},
Expr::Block(stmts) => self.eval_statements_in_scope(stmts, scope),
Expr::Call { function, args } => {
let evaluated = args
.iter()
.map(|e| self.eval_expr_in_scope(e, scope))
.collect();
self.call_func(function, evaluated)
}
Expr::MethodCall {
target,
method,
arg,
} => {
let tgt = self.eval_expr_in_scope(target, scope);
self.eval_method_call(tgt, method, *arg.clone(), scope)
}
Expr::Index { target, index } => {
let arr = self.eval_expr_in_scope(target, scope);
let idx = self.eval_expr_in_scope(index, scope);
match (arr, idx) {
(VarValue::Array(a), VarValue::Int(i)) => a
.get(i as usize)
.cloned()
.unwrap_or_else(|| panic!("Index {i} is out of bounds")),
_ => panic!("Indexing is only supported for arrays using integer indices"),
}
}
Expr::Range { start, end } => {
let start_val = self.eval_expr(*start.clone());
let end_val = self.eval_expr(*end.clone());
let start_i = match start_val {
VarValue::Int(i) => i,
_ => panic!("range start must evaluate to an integer"),
};
let end_i = match end_val {
VarValue::Int(i) => i,
_ => panic!("range end must evaluate to an integer"),
};
let values = (start_i..end_i).map(VarValue::Int).collect();
VarValue::Array(values)
}
}
}
fn eval_method_call(
&mut self,
target: VarValue,
method: &str,
arg: Expr,
scope: &HashMap<String, VarValue>,
) -> VarValue {
match (target, method) {
(VarValue::Array(arr), "map") => {
let mapped = arr
.into_iter()
.map(|v| self.eval_lambda_callable(&arg, v, scope))
.collect();
VarValue::Array(mapped)
}
(VarValue::Array(arr), "for_each") => {
for v in &arr {
self.eval_lambda_callable(&arg, v.clone(), scope);
}
VarValue::Array(arr)
}
(other, m) => panic!("Unsupported method '{m}' on {other:?}"),
}
}
fn eval_lambda_callable(
&mut self,
callable: &Expr,
input: VarValue,
scope: &HashMap<String, VarValue>,
) -> VarValue {
match callable {
Expr::Identifier(name) => self.call_func(name, vec![input]),
Expr::Lambda { param, body } => {
let mut temp = scope.clone();
temp.insert(param.clone(), input);
self.eval_statements_in_scope(body, &mut temp)
}
other => panic!("map expects a function name or lambda, got {other:?}"),
}
}
}