use super::*;
use crate::error::{ OfficeError, Result, XlsxError };
use crate::xlsx::cell::{ CellReference, CellValue };
use std::collections::HashMap;
impl FormulaCalculator {
pub fn new(cell_provider: Box<dyn CellProvider>) -> Self {
let mut calculator = Self {
cell_provider,
function_library: FunctionLibrary::new(),
cache: HashMap::new(),
};
calculator.register_builtin_functions();
calculator
}
pub fn evaluate(&mut self, expr: &FormulaExpression) -> Result<FormulaValue> {
match expr {
FormulaExpression::Constant(value) => Ok(value.clone()),
FormulaExpression::CellRef(cell_ref) => {
let cell_value = self.cell_provider.get_cell_value(cell_ref)?;
Ok(self.cell_value_to_formula_value(cell_value))
}
FormulaExpression::RangeRef(start, end) => {
let values = self.cell_provider.get_range_values(start, end)?;
let formula_values = values
.into_iter()
.map(|row| {
row.into_iter()
.map(|cell| self.cell_value_to_formula_value(cell))
.collect()
})
.collect();
Ok(FormulaValue::Array(formula_values))
}
FormulaExpression::Function { name, args } => self.evaluate_function(name, args),
FormulaExpression::BinaryOp { op, left, right } => {
let left_val = self.evaluate(left)?;
let right_val = self.evaluate(right)?;
self.evaluate_binary_op(op, &left_val, &right_val)
}
FormulaExpression::UnaryOp { op, operand } => {
let operand_val = self.evaluate(operand)?;
self.evaluate_unary_op(op, &operand_val)
}
}
}
fn cell_value_to_formula_value(&self, cell_value: CellValue) -> FormulaValue {
match cell_value {
CellValue::Empty => FormulaValue::Number(0.0),
CellValue::Number(n) => FormulaValue::Number(n),
CellValue::Text(s) => FormulaValue::Text(s),
CellValue::Boolean(b) => FormulaValue::Boolean(b),
CellValue::DateTime(d) => FormulaValue::Number(d),
CellValue::Formula(_) => FormulaValue::Error(FormulaError::ReferenceError),
CellValue::Error(e) => FormulaValue::Error(FormulaError::ValueError),
}
}
fn evaluate_binary_op(
&self,
op: &BinaryOperator,
left: &FormulaValue,
right: &FormulaValue
) -> Result<FormulaValue> {
if left.is_error() {
return Ok(left.clone());
}
if right.is_error() {
return Ok(right.clone());
}
match op {
BinaryOperator::Add => {
let left_num = left.as_number()?;
let right_num = right.as_number()?;
Ok(FormulaValue::Number(left_num + right_num))
}
BinaryOperator::Subtract => {
let left_num = left.as_number()?;
let right_num = right.as_number()?;
Ok(FormulaValue::Number(left_num - right_num))
}
BinaryOperator::Multiply => {
let left_num = left.as_number()?;
let right_num = right.as_number()?;
Ok(FormulaValue::Number(left_num * right_num))
}
BinaryOperator::Divide => {
let left_num = left.as_number()?;
let right_num = right.as_number()?;
if right_num == 0.0 {
Ok(FormulaValue::Error(FormulaError::DivisionByZero))
} else {
Ok(FormulaValue::Number(left_num / right_num))
}
}
BinaryOperator::Power => {
let left_num = left.as_number()?;
let right_num = right.as_number()?;
Ok(FormulaValue::Number(left_num.powf(right_num)))
}
BinaryOperator::Equal => Ok(FormulaValue::Boolean(self.values_equal(left, right))),
BinaryOperator::NotEqual => Ok(FormulaValue::Boolean(!self.values_equal(left, right))),
BinaryOperator::LessThan => {
let result = self.compare_values(left, right)?;
Ok(FormulaValue::Boolean(result < 0))
}
BinaryOperator::LessThanOrEqual => {
let result = self.compare_values(left, right)?;
Ok(FormulaValue::Boolean(result <= 0))
}
BinaryOperator::GreaterThan => {
let result = self.compare_values(left, right)?;
Ok(FormulaValue::Boolean(result > 0))
}
BinaryOperator::GreaterThanOrEqual => {
let result = self.compare_values(left, right)?;
Ok(FormulaValue::Boolean(result >= 0))
}
BinaryOperator::Concatenate => {
let left_text = left.as_text();
let right_text = right.as_text();
Ok(FormulaValue::Text(format!("{}{}", left_text, right_text)))
}
BinaryOperator::LogicalOr => {
let left_bool = left.as_boolean()?;
let right_bool = right.as_boolean()?;
Ok(FormulaValue::Boolean(left_bool || right_bool))
}
BinaryOperator::LogicalAnd => {
let left_bool = left.as_boolean()?;
let right_bool = right.as_boolean()?;
Ok(FormulaValue::Boolean(left_bool && right_bool))
}
}
}
fn evaluate_unary_op(
&self,
op: &UnaryOperator,
operand: &FormulaValue
) -> Result<FormulaValue> {
if operand.is_error() {
return Ok(operand.clone());
}
match op {
UnaryOperator::Plus => {
let num = operand.as_number()?;
Ok(FormulaValue::Number(num))
}
UnaryOperator::Minus => {
let num = operand.as_number()?;
Ok(FormulaValue::Number(-num))
}
UnaryOperator::Percent => {
let num = operand.as_number()?;
Ok(FormulaValue::Number(num / 100.0))
}
UnaryOperator::Factorial => {
let num = operand.as_number()?;
if num < 0.0 || num.fract() != 0.0 {
return Ok(FormulaValue::Error(FormulaError::ValueError));
}
let mut result = 1.0;
for i in 1..=num as u64 {
result *= i as f64;
}
Ok(FormulaValue::Number(result))
}
}
}
fn values_equal(&self, left: &FormulaValue, right: &FormulaValue) -> bool {
match (left, right) {
(FormulaValue::Number(a), FormulaValue::Number(b)) => (a - b).abs() < f64::EPSILON,
(FormulaValue::Text(a), FormulaValue::Text(b)) => a == b,
(FormulaValue::Boolean(a), FormulaValue::Boolean(b)) => a == b,
(FormulaValue::Error(a), FormulaValue::Error(b)) => a == b,
_ => false,
}
}
fn compare_values(&self, left: &FormulaValue, right: &FormulaValue) -> Result<i32> {
match (left, right) {
(FormulaValue::Number(a), FormulaValue::Number(b)) => {
Ok(a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal) as i32)
}
(FormulaValue::Text(a), FormulaValue::Text(b)) => Ok(a.cmp(b) as i32),
(FormulaValue::Boolean(a), FormulaValue::Boolean(b)) => Ok(a.cmp(b) as i32),
_ =>
Err(
OfficeError::Xlsx(XlsxError::InvalidFormula {
formula: "Cannot compare different types".to_string(),
})
),
}
}
fn evaluate_function(
&mut self,
name: &str,
args: &[FormulaExpression]
) -> Result<FormulaValue> {
let mut arg_values = Vec::new();
for arg in args {
arg_values.push(self.evaluate(arg)?);
}
self.function_library.call_function(name, &arg_values)
}
fn register_builtin_functions(&mut self) {
self.function_library.register(Box::new(SumFunction));
self.function_library.register(Box::new(AverageFunction));
self.function_library.register(Box::new(MaxFunction));
self.function_library.register(Box::new(MinFunction));
self.function_library.register(Box::new(CountFunction));
self.function_library.register(Box::new(RoundFunction));
self.function_library.register(Box::new(AbsFunction));
self.function_library.register(Box::new(SqrtFunction));
self.function_library.register(Box::new(IfFunction));
self.function_library.register(Box::new(AndFunction));
self.function_library.register(Box::new(OrFunction));
self.function_library.register(Box::new(NotFunction));
self.function_library.register(Box::new(ConcatenateFunction));
self.function_library.register(Box::new(LeftFunction));
self.function_library.register(Box::new(RightFunction));
self.function_library.register(Box::new(MidFunction));
self.function_library.register(Box::new(LenFunction));
self.function_library.register(Box::new(UpperFunction));
self.function_library.register(Box::new(LowerFunction));
}
}
impl FunctionLibrary {
pub fn new() -> Self {
Self {
functions: HashMap::new(),
}
}
pub fn register(&mut self, function: Box<dyn FormulaFunction>) {
self.functions.insert(function.name().to_uppercase(), function);
}
pub fn call_function(&self, name: &str, args: &[FormulaValue]) -> Result<FormulaValue> {
let function = self.functions.get(&name.to_uppercase()).ok_or_else(|| {
OfficeError::Xlsx(XlsxError::InvalidFormula {
formula: format!("Unknown function: {}", name),
})
})?;
if args.len() < function.min_args() {
return Err(
OfficeError::Xlsx(XlsxError::InvalidFormula {
formula: format!(
"Function {} requires at least {} arguments",
name,
function.min_args()
),
})
);
}
if let Some(max_args) = function.max_args() {
if args.len() > max_args {
return Err(
OfficeError::Xlsx(XlsxError::InvalidFormula {
formula: format!(
"Function {} accepts at most {} arguments",
name,
max_args
),
})
);
}
}
function.execute(args)
}
}
struct SumFunction;
impl FormulaFunction for SumFunction {
fn name(&self) -> &str {
"SUM"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut sum = 0.0;
for arg in args {
match arg {
FormulaValue::Number(n) => {
sum += n;
}
FormulaValue::Boolean(b) => {
sum += if *b { 1.0 } else { 0.0 };
}
FormulaValue::Array(arr) => {
for row in arr {
for cell in row {
if let FormulaValue::Number(n) = cell {
sum += n;
}
}
}
}
FormulaValue::Error(e) => {
return Ok(FormulaValue::Error(e.clone()));
}
_ => {} }
}
Ok(FormulaValue::Number(sum))
}
}
struct AverageFunction;
impl FormulaFunction for AverageFunction {
fn name(&self) -> &str {
"AVERAGE"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut sum = 0.0;
let mut count = 0;
for arg in args {
match arg {
FormulaValue::Number(n) => {
sum += n;
count += 1;
}
FormulaValue::Boolean(b) => {
sum += if *b { 1.0 } else { 0.0 };
count += 1;
}
FormulaValue::Array(arr) => {
for row in arr {
for cell in row {
if let FormulaValue::Number(n) = cell {
sum += n;
count += 1;
}
}
}
}
FormulaValue::Error(e) => {
return Ok(FormulaValue::Error(e.clone()));
}
_ => {} }
}
if count == 0 {
Ok(FormulaValue::Error(FormulaError::DivisionByZero))
} else {
Ok(FormulaValue::Number(sum / (count as f64)))
}
}
}
struct MaxFunction;
impl FormulaFunction for MaxFunction {
fn name(&self) -> &str {
"MAX"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut max_val = f64::NEG_INFINITY;
let mut has_number = false;
for arg in args {
match arg {
FormulaValue::Number(n) => {
max_val = max_val.max(*n);
has_number = true;
}
FormulaValue::Array(arr) => {
for row in arr {
for cell in row {
if let FormulaValue::Number(n) = cell {
max_val = max_val.max(*n);
has_number = true;
}
}
}
}
FormulaValue::Error(e) => {
return Ok(FormulaValue::Error(e.clone()));
}
_ => {} }
}
if has_number {
Ok(FormulaValue::Number(max_val))
} else {
Ok(FormulaValue::Number(0.0))
}
}
}
struct MinFunction;
impl FormulaFunction for MinFunction {
fn name(&self) -> &str {
"MIN"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut min_val = f64::INFINITY;
let mut has_number = false;
for arg in args {
match arg {
FormulaValue::Number(n) => {
min_val = min_val.min(*n);
has_number = true;
}
FormulaValue::Array(arr) => {
for row in arr {
for cell in row {
if let FormulaValue::Number(n) = cell {
min_val = min_val.min(*n);
has_number = true;
}
}
}
}
FormulaValue::Error(e) => {
return Ok(FormulaValue::Error(e.clone()));
}
_ => {} }
}
if has_number {
Ok(FormulaValue::Number(min_val))
} else {
Ok(FormulaValue::Number(0.0))
}
}
}
struct CountFunction;
impl FormulaFunction for CountFunction {
fn name(&self) -> &str {
"COUNT"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut count = 0;
for arg in args {
match arg {
FormulaValue::Number(_) => {
count += 1;
}
FormulaValue::Array(arr) => {
for row in arr {
for cell in row {
if let FormulaValue::Number(_) = cell {
count += 1;
}
}
}
}
FormulaValue::Error(e) => {
return Ok(FormulaValue::Error(e.clone()));
}
_ => {} }
}
Ok(FormulaValue::Number(count as f64))
}
}
struct RoundFunction;
impl FormulaFunction for RoundFunction {
fn name(&self) -> &str {
"ROUND"
}
fn min_args(&self) -> usize {
2
}
fn max_args(&self) -> Option<usize> {
Some(2)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let number = args[0].as_number()?;
let digits = args[1].as_number()? as i32;
let multiplier = (10.0_f64).powi(digits);
let rounded = (number * multiplier).round() / multiplier;
Ok(FormulaValue::Number(rounded))
}
}
struct AbsFunction;
impl FormulaFunction for AbsFunction {
fn name(&self) -> &str {
"ABS"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let number = args[0].as_number()?;
Ok(FormulaValue::Number(number.abs()))
}
}
struct SqrtFunction;
impl FormulaFunction for SqrtFunction {
fn name(&self) -> &str {
"SQRT"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let number = args[0].as_number()?;
if number < 0.0 {
Ok(FormulaValue::Error(FormulaError::NumError))
} else {
Ok(FormulaValue::Number(number.sqrt()))
}
}
}
struct IfFunction;
impl FormulaFunction for IfFunction {
fn name(&self) -> &str {
"IF"
}
fn min_args(&self) -> usize {
2
}
fn max_args(&self) -> Option<usize> {
Some(3)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let condition = args[0].as_boolean()?;
if condition {
Ok(args[1].clone())
} else if args.len() > 2 {
Ok(args[2].clone())
} else {
Ok(FormulaValue::Boolean(false))
}
}
}
struct AndFunction;
impl FormulaFunction for AndFunction {
fn name(&self) -> &str {
"AND"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
for arg in args {
if !arg.as_boolean()? {
return Ok(FormulaValue::Boolean(false));
}
}
Ok(FormulaValue::Boolean(true))
}
}
struct OrFunction;
impl FormulaFunction for OrFunction {
fn name(&self) -> &str {
"OR"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
for arg in args {
if arg.as_boolean()? {
return Ok(FormulaValue::Boolean(true));
}
}
Ok(FormulaValue::Boolean(false))
}
}
struct NotFunction;
impl FormulaFunction for NotFunction {
fn name(&self) -> &str {
"NOT"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let value = args[0].as_boolean()?;
Ok(FormulaValue::Boolean(!value))
}
}
struct ConcatenateFunction;
impl FormulaFunction for ConcatenateFunction {
fn name(&self) -> &str {
"CONCATENATE"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
None
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let mut result = String::new();
for arg in args {
result.push_str(&arg.as_text());
}
Ok(FormulaValue::Text(result))
}
}
struct LeftFunction;
impl FormulaFunction for LeftFunction {
fn name(&self) -> &str {
"LEFT"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(2)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
let num_chars = if args.len() > 1 { args[1].as_number()? as usize } else { 1 };
let result = text.chars().take(num_chars).collect::<String>();
Ok(FormulaValue::Text(result))
}
}
struct RightFunction;
impl FormulaFunction for RightFunction {
fn name(&self) -> &str {
"RIGHT"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(2)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
let num_chars = if args.len() > 1 { args[1].as_number()? as usize } else { 1 };
let chars: Vec<char> = text.chars().collect();
let start = chars.len().saturating_sub(num_chars);
let result = chars[start..].iter().collect::<String>();
Ok(FormulaValue::Text(result))
}
}
struct MidFunction;
impl FormulaFunction for MidFunction {
fn name(&self) -> &str {
"MID"
}
fn min_args(&self) -> usize {
3
}
fn max_args(&self) -> Option<usize> {
Some(3)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
let start = (args[1].as_number()? as usize).saturating_sub(1); let length = args[2].as_number()? as usize;
let chars: Vec<char> = text.chars().collect();
let end = (start + length).min(chars.len());
if start >= chars.len() {
Ok(FormulaValue::Text(String::new()))
} else {
let result = chars[start..end].iter().collect::<String>();
Ok(FormulaValue::Text(result))
}
}
}
struct LenFunction;
impl FormulaFunction for LenFunction {
fn name(&self) -> &str {
"LEN"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
Ok(FormulaValue::Number(text.chars().count() as f64))
}
}
struct UpperFunction;
impl FormulaFunction for UpperFunction {
fn name(&self) -> &str {
"UPPER"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
Ok(FormulaValue::Text(text.to_uppercase()))
}
}
struct LowerFunction;
impl FormulaFunction for LowerFunction {
fn name(&self) -> &str {
"LOWER"
}
fn min_args(&self) -> usize {
1
}
fn max_args(&self) -> Option<usize> {
Some(1)
}
fn execute(&self, args: &[FormulaValue]) -> Result<FormulaValue> {
let text = args[0].as_text();
Ok(FormulaValue::Text(text.to_lowercase()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::xlsx::cell::CellValue;
struct MockCellProvider {
cells: HashMap<String, CellValue>,
}
impl MockCellProvider {
fn new() -> Self {
let mut cells = HashMap::new();
cells.insert("A1".to_string(), CellValue::Number(10.0));
cells.insert("A2".to_string(), CellValue::Number(20.0));
cells.insert("A3".to_string(), CellValue::Number(30.0));
Self { cells }
}
}
impl CellProvider for MockCellProvider {
fn get_cell_value(&self, reference: &CellReference) -> Result<CellValue> {
let key = reference.to_a1();
Ok(self.cells.get(&key).cloned().unwrap_or(CellValue::Empty))
}
fn get_range_values(
&self,
start: &CellReference,
end: &CellReference
) -> Result<Vec<Vec<CellValue>>> {
let mut result = Vec::new();
for row in start.row..=end.row {
let mut row_values = Vec::new();
for col in start.column..=end.column {
let ref_key = CellReference::new(col, row).to_a1();
row_values.push(self.cells.get(&ref_key).cloned().unwrap_or(CellValue::Empty));
}
result.push(row_values);
}
Ok(result)
}
}
#[test]
fn test_sum_function() {
let provider = Box::new(MockCellProvider::new());
let mut calculator = FormulaCalculator::new(provider);
let args = vec![
FormulaValue::Number(1.0),
FormulaValue::Number(2.0),
FormulaValue::Number(3.0)
];
let result = calculator.function_library.call_function("SUM", &args).unwrap();
assert_eq!(result, FormulaValue::Number(6.0));
}
#[test]
fn test_if_function() {
let provider = Box::new(MockCellProvider::new());
let mut calculator = FormulaCalculator::new(provider);
let args = vec![
FormulaValue::Boolean(true),
FormulaValue::Text("Yes".to_string()),
FormulaValue::Text("No".to_string())
];
let result = calculator.function_library.call_function("IF", &args).unwrap();
assert_eq!(result, FormulaValue::Text("Yes".to_string()));
}
#[test]
fn test_binary_operations() {
let provider = Box::new(MockCellProvider::new());
let mut calculator = FormulaCalculator::new(provider);
let left = FormulaValue::Number(10.0);
let right = FormulaValue::Number(3.0);
let result = calculator.evaluate_binary_op(&BinaryOperator::Add, &left, &right).unwrap();
assert_eq!(result, FormulaValue::Number(13.0));
let result = calculator.evaluate_binary_op(&BinaryOperator::Divide, &left, &right).unwrap();
assert_eq!(result, FormulaValue::Number(10.0 / 3.0));
}
#[test]
fn test_factorial() {
let provider = Box::new(MockCellProvider::new());
let calculator = FormulaCalculator::new(provider);
let result = calculator
.evaluate_unary_op(&UnaryOperator::Factorial, &FormulaValue::Number(5.0))
.unwrap();
assert_eq!(result, FormulaValue::Number(120.0));
let result = calculator
.evaluate_unary_op(&UnaryOperator::Factorial, &FormulaValue::Number(0.0))
.unwrap();
assert_eq!(result, FormulaValue::Number(1.0));
let result = calculator
.evaluate_unary_op(&UnaryOperator::Factorial, &FormulaValue::Number(-3.0))
.unwrap();
assert_eq!(result, FormulaValue::Error(FormulaError::ValueError));
let result = calculator
.evaluate_unary_op(&UnaryOperator::Factorial, &FormulaValue::Number(3.5))
.unwrap();
assert_eq!(result, FormulaValue::Error(FormulaError::ValueError));
let result = calculator
.evaluate_unary_op(&UnaryOperator::Factorial, &FormulaValue::Number(10.0))
.unwrap();
assert_eq!(result, FormulaValue::Number(3628800.0)); }
}