547 lines
19 KiB
Python
547 lines
19 KiB
Python
#!/usr/bin/env python
|
||
# -*- coding: utf-8 -*-
|
||
"""
|
||
单据计算引擎
|
||
|
||
支持:
|
||
- 基本运算:+, -, *, /, %, **
|
||
- 聚合函数:sum, avg, max, min, count
|
||
- 内置函数:round, abs, ceil, floor, numberToChinese
|
||
- 条件表达式:if-else (三元运算符)
|
||
- 安全执行:使用 simpleeval 库防止代码注入
|
||
"""
|
||
import logging
|
||
import math
|
||
import re
|
||
from decimal import Decimal, ROUND_HALF_UP
|
||
from typing import Any, Dict, List, Optional, Union
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
# 中文数字映射
|
||
CHINESE_DIGITS = ['零', '壹', '贰', '叁', '肆', '伍', '陆', '柒', '捌', '玖']
|
||
CHINESE_UNITS = ['', '拾', '佰', '仟']
|
||
CHINESE_GROUP_UNITS = ['', '万', '亿', '兆']
|
||
CHINESE_DECIMAL_UNITS = ['角', '分', '厘', '毫']
|
||
|
||
|
||
def number_to_chinese(num: Union[int, float, Decimal, str]) -> str:
|
||
"""
|
||
将数字转换为中文大写金额
|
||
|
||
Args:
|
||
num: 数字(支持整数、浮点数、Decimal、字符串)
|
||
|
||
Returns:
|
||
中文大写金额字符串
|
||
|
||
Examples:
|
||
>>> number_to_chinese(1234.56)
|
||
'壹仟贰佰叁拾肆元伍角陆分'
|
||
>>> number_to_chinese(0)
|
||
'零元整'
|
||
"""
|
||
if num is None:
|
||
return ''
|
||
|
||
try:
|
||
# 转换为 Decimal 以保证精度
|
||
if isinstance(num, str):
|
||
num = Decimal(num.replace(',', ''))
|
||
elif isinstance(num, float):
|
||
num = Decimal(str(num))
|
||
elif isinstance(num, int):
|
||
num = Decimal(num)
|
||
elif not isinstance(num, Decimal):
|
||
num = Decimal(str(num))
|
||
except Exception:
|
||
return str(num)
|
||
|
||
# 处理负数
|
||
if num < 0:
|
||
return '负' + number_to_chinese(-num)
|
||
|
||
# 处理零
|
||
if num == 0:
|
||
return '零元整'
|
||
|
||
# 分离整数和小数部分
|
||
num = num.quantize(Decimal('0.0001'), rounding=ROUND_HALF_UP)
|
||
str_num = str(num)
|
||
|
||
if '.' in str_num:
|
||
int_part, dec_part = str_num.split('.')
|
||
else:
|
||
int_part, dec_part = str_num, ''
|
||
|
||
result = ''
|
||
|
||
# 处理整数部分
|
||
if int_part and int(int_part) > 0:
|
||
int_part = int_part.lstrip('0') or '0'
|
||
length = len(int_part)
|
||
|
||
# 按4位分组处理
|
||
groups = []
|
||
while int_part:
|
||
groups.insert(0, int_part[-4:])
|
||
int_part = int_part[:-4]
|
||
|
||
for i, group in enumerate(groups):
|
||
group_result = ''
|
||
group = group.zfill(4)
|
||
|
||
for j, digit in enumerate(group):
|
||
d = int(digit)
|
||
unit_index = 3 - j
|
||
|
||
if d != 0:
|
||
group_result += CHINESE_DIGITS[d] + CHINESE_UNITS[unit_index]
|
||
else:
|
||
# 处理连续零
|
||
if group_result and not group_result.endswith('零'):
|
||
group_result += '零'
|
||
|
||
# 移除末尾的零
|
||
group_result = group_result.rstrip('零')
|
||
|
||
if group_result:
|
||
group_unit_index = len(groups) - 1 - i
|
||
group_result += CHINESE_GROUP_UNITS[group_unit_index] if group_unit_index < len(CHINESE_GROUP_UNITS) else ''
|
||
result += group_result
|
||
|
||
result += '元'
|
||
else:
|
||
result = '零元'
|
||
|
||
# 处理小数部分
|
||
if dec_part:
|
||
dec_part = dec_part[:4] # 最多4位小数
|
||
has_decimal = False
|
||
|
||
for i, digit in enumerate(dec_part):
|
||
d = int(digit)
|
||
if d != 0:
|
||
result += CHINESE_DIGITS[d] + CHINESE_DECIMAL_UNITS[i]
|
||
has_decimal = True
|
||
|
||
if not has_decimal:
|
||
result += '整'
|
||
else:
|
||
result += '整'
|
||
|
||
return result
|
||
|
||
|
||
def safe_get_value(data: Dict[str, Any], key: str, default: Any = 0) -> Any:
|
||
"""
|
||
安全地从字典中获取值,支持点号路径
|
||
|
||
Args:
|
||
data: 数据字典
|
||
key: 键名(支持点号路径,如 'order.total')
|
||
default: 默认值
|
||
|
||
Returns:
|
||
获取的值或默认值
|
||
"""
|
||
if not key:
|
||
return default
|
||
|
||
parts = key.split('.')
|
||
value = data
|
||
|
||
for part in parts:
|
||
if isinstance(value, dict):
|
||
value = value.get(part)
|
||
elif isinstance(value, list):
|
||
# 支持数组索引
|
||
if part.lstrip('-').isdigit():
|
||
index = int(part)
|
||
if -len(value) <= index < len(value):
|
||
value = value[index]
|
||
else:
|
||
return default
|
||
else:
|
||
return default
|
||
else:
|
||
return default
|
||
|
||
if value is None:
|
||
return default
|
||
|
||
return value if value is not None else default
|
||
|
||
|
||
class CalculationEngine:
|
||
"""单据计算引擎"""
|
||
|
||
# 支持的聚合函数
|
||
AGGREGATE_FUNCTIONS = {
|
||
'sum': lambda values: sum(v for v in values if v is not None),
|
||
'avg': lambda values: sum(v for v in values if v is not None) / len([v for v in values if v is not None]) if values else 0,
|
||
'max': lambda values: max((v for v in values if v is not None), default=0),
|
||
'min': lambda values: min((v for v in values if v is not None), default=0),
|
||
'count': lambda values: len([v for v in values if v is not None]),
|
||
}
|
||
|
||
# 安全的内置函数
|
||
SAFE_FUNCTIONS = {
|
||
'abs': abs,
|
||
'round': round,
|
||
'ceil': math.ceil,
|
||
'floor': math.floor,
|
||
'max': max,
|
||
'min': min,
|
||
'sum': sum,
|
||
'len': len,
|
||
'float': float,
|
||
'int': int,
|
||
'str': str,
|
||
'numberToChinese': number_to_chinese,
|
||
'toChineseAmount': number_to_chinese,
|
||
}
|
||
|
||
# 安全的运算符
|
||
SAFE_OPERATORS = {
|
||
'+', '-', '*', '/', '//', '%', '**',
|
||
'==', '!=', '<', '>', '<=', '>=',
|
||
'and', 'or', 'not',
|
||
'(', ')', ',', '.',
|
||
}
|
||
|
||
@classmethod
|
||
def evaluate_formula(cls, formula: str, context: Dict[str, Any]) -> Any:
|
||
"""
|
||
安全地执行计算公式
|
||
|
||
Args:
|
||
formula: 计算公式,如 "quantity * unit_price * (1 - discount_rate)"
|
||
context: 上下文数据
|
||
|
||
Returns:
|
||
计算结果
|
||
"""
|
||
if not formula:
|
||
return None
|
||
|
||
try:
|
||
# 替换公式中的变量
|
||
evaluated_formula = cls._replace_variables(formula, context)
|
||
logger.info(f"公式: {formula} -> 替换后: {evaluated_formula}")
|
||
|
||
# 使用 eval 执行(在受限环境中)
|
||
# 注意:这里使用了安全的方式,只允许特定的函数和运算
|
||
result = cls._safe_eval(evaluated_formula, context)
|
||
logger.info(f"公式执行结果: {result}")
|
||
|
||
return result
|
||
except Exception as e:
|
||
logger.error(f"公式计算失败: {formula}, 错误: {e}", exc_info=True)
|
||
return None
|
||
|
||
@classmethod
|
||
def _replace_variables(cls, formula: str, context: Dict[str, Any]) -> str:
|
||
"""
|
||
替换公式中的变量为实际值
|
||
|
||
支持的变量格式:
|
||
- 简单变量:quantity, unit_price
|
||
- 点号路径:order.total, items[0].price
|
||
"""
|
||
# 匹配变量名(字母开头,可包含字母、数字、下划线、点号、方括号)
|
||
pattern = r'\b([a-zA-Z_][a-zA-Z0-9_]*(?:\.[a-zA-Z_][a-zA-Z0-9_]*|\[\d+\])*)\b'
|
||
|
||
def replace_var(match):
|
||
var_name = match.group(1)
|
||
|
||
# 跳过函数名
|
||
if var_name in cls.SAFE_FUNCTIONS:
|
||
return var_name
|
||
|
||
# 跳过 Python 关键字
|
||
if var_name in ('and', 'or', 'not', 'if', 'else', 'True', 'False', 'None'):
|
||
return var_name
|
||
|
||
# 获取变量值
|
||
value = safe_get_value(context, var_name, 0)
|
||
|
||
# 转换为字符串表示
|
||
if value is None:
|
||
return '0'
|
||
elif isinstance(value, str):
|
||
# 尝试转换为数字
|
||
try:
|
||
return str(float(value))
|
||
except ValueError:
|
||
return f'"{value}"'
|
||
elif isinstance(value, bool):
|
||
return str(value)
|
||
elif isinstance(value, (int, float, Decimal)):
|
||
return str(float(value))
|
||
else:
|
||
return '0'
|
||
|
||
return re.sub(pattern, replace_var, formula)
|
||
|
||
@classmethod
|
||
def _safe_eval(cls, expression: str, context: Dict[str, Any]) -> Any:
|
||
"""
|
||
安全地执行表达式
|
||
|
||
使用受限的 eval 环境,只允许特定的函数和运算
|
||
"""
|
||
# 构建安全的执行环境
|
||
safe_globals = {
|
||
'__builtins__': {},
|
||
**cls.SAFE_FUNCTIONS,
|
||
}
|
||
|
||
# 添加上下文数据
|
||
safe_locals = dict(context)
|
||
|
||
try:
|
||
result = eval(expression, safe_globals, safe_locals)
|
||
return result
|
||
except Exception as e:
|
||
logger.warning(f"表达式执行失败: {expression}, 错误: {e}")
|
||
raise
|
||
|
||
@classmethod
|
||
def calculate_aggregation(
|
||
cls,
|
||
data: Dict[str, Any],
|
||
source: str,
|
||
field: str,
|
||
function: str
|
||
) -> Any:
|
||
"""
|
||
计算聚合值
|
||
|
||
Args:
|
||
data: 数据字典
|
||
source: 数据源(子表名)
|
||
field: 聚合字段
|
||
function: 聚合函数名
|
||
|
||
Returns:
|
||
聚合结果
|
||
"""
|
||
logger.info(f"聚合计算开始: source={source}, field={field}, function={function}")
|
||
logger.info(f"数据中的顶层键: {list(data.keys())}")
|
||
|
||
# 获取子表数据
|
||
sub_table_data = None
|
||
|
||
# 1. 优先从 sub_tables 中查找(表单数据的标准结构)
|
||
if 'sub_tables' in data and isinstance(data['sub_tables'], dict):
|
||
sub_tables = data['sub_tables']
|
||
logger.info(f"sub_tables 中的键: {list(sub_tables.keys())}")
|
||
|
||
# 直接匹配
|
||
if source in sub_tables:
|
||
sub_table_data = sub_tables[source]
|
||
logger.info(f"从 sub_tables 中直接匹配到 ({source}): 找到 {len(sub_table_data) if isinstance(sub_table_data, list) else 0} 条")
|
||
else:
|
||
# 尝试模糊匹配(source 可能是字段名,sub_tables 的键可能是表名)
|
||
# 例如:source='product_details',sub_tables 键可能是 'fd_product_details' 或 'contract_product_details'
|
||
for key in sub_tables.keys():
|
||
if source in key or key in source or key.endswith(f'_{source}') or key.endswith(source):
|
||
sub_table_data = sub_tables[key]
|
||
logger.info(f"从 sub_tables 中模糊匹配到 ({source} -> {key}): 找到 {len(sub_table_data) if isinstance(sub_table_data, list) else 0} 条")
|
||
break
|
||
|
||
# 2. 如果 sub_tables 中没有,尝试直接从顶层获取
|
||
if not sub_table_data:
|
||
sub_table_data = safe_get_value(data, source, [])
|
||
logger.info(f"从顶层获取子表数据 ({source}): {type(sub_table_data)}")
|
||
|
||
if not isinstance(sub_table_data, list):
|
||
logger.warning(f"聚合数据源不是数组: {source}, 实际类型: {type(sub_table_data)}")
|
||
return 0
|
||
|
||
if len(sub_table_data) == 0:
|
||
logger.warning(f"聚合数据源为空数组: {source}")
|
||
return 0
|
||
|
||
# 提取字段值
|
||
values = []
|
||
for i, item in enumerate(sub_table_data):
|
||
if isinstance(item, dict):
|
||
logger.info(f"子表第{i}行的键: {list(item.keys())}")
|
||
value = safe_get_value(item, field, None)
|
||
logger.info(f"子表第{i}行的 {field} 值: {value}")
|
||
if value is not None:
|
||
try:
|
||
values.append(float(value))
|
||
except (ValueError, TypeError) as e:
|
||
logger.warning(f"无法转换为数字: {value}, 错误: {e}")
|
||
|
||
logger.info(f"提取到的数值列表: {values}")
|
||
|
||
# 执行聚合函数
|
||
agg_func = cls.AGGREGATE_FUNCTIONS.get(function.lower())
|
||
if not agg_func:
|
||
logger.warning(f"不支持的聚合函数: {function}")
|
||
return 0
|
||
|
||
try:
|
||
result = agg_func(values)
|
||
logger.info(f"聚合计算结果: {result}")
|
||
return result
|
||
except Exception as e:
|
||
logger.warning(f"聚合计算失败: {e}")
|
||
return 0
|
||
|
||
@classmethod
|
||
def format_value(
|
||
cls,
|
||
value: Any,
|
||
format_type: str = 'number',
|
||
decimal_places: int = 2
|
||
) -> Any:
|
||
"""
|
||
格式化计算结果
|
||
|
||
Args:
|
||
value: 原始值
|
||
format_type: 格式化类型 (number/money/percent/chinese)
|
||
decimal_places: 小数位数
|
||
|
||
Returns:
|
||
格式化后的值(字符串,保留指定小数位数)
|
||
"""
|
||
if value is None:
|
||
return None
|
||
|
||
try:
|
||
num_value = float(value)
|
||
except (ValueError, TypeError):
|
||
return value
|
||
|
||
if format_type == 'chinese':
|
||
return number_to_chinese(num_value)
|
||
elif format_type == 'percent':
|
||
# 百分比:乘以100后保留指定小数位,返回格式化字符串
|
||
percent_value = num_value * 100
|
||
if decimal_places == 0:
|
||
return str(int(round(percent_value)))
|
||
return f"{percent_value:.{decimal_places}f}"
|
||
else:
|
||
# number 和 money:保留指定小数位,返回格式化字符串
|
||
if decimal_places == 0:
|
||
return str(int(round(num_value)))
|
||
return f"{num_value:.{decimal_places}f}"
|
||
|
||
@classmethod
|
||
async def calculate_all(
|
||
cls,
|
||
calculation_rules: Optional[Dict[str, Any]],
|
||
form_data: Dict[str, Any]
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
执行所有计算规则
|
||
|
||
Args:
|
||
calculation_rules: 计算规则配置
|
||
form_data: 表单数据
|
||
|
||
Returns:
|
||
计算结果字典
|
||
"""
|
||
if not calculation_rules:
|
||
return {}
|
||
|
||
results = {}
|
||
|
||
# 创建计算上下文(包含原始数据和已计算的结果)
|
||
context = dict(form_data)
|
||
|
||
# 1. 先执行聚合计算(因为计算字段可能依赖聚合结果)
|
||
aggregations = calculation_rules.get('aggregations', [])
|
||
logger.info(f"开始执行聚合计算,共 {len(aggregations)} 个")
|
||
for agg in aggregations:
|
||
try:
|
||
logger.info(f"聚合配置原始数据: {agg}")
|
||
name = agg.get('name')
|
||
source = agg.get('source')
|
||
field = agg.get('field')
|
||
function = agg.get('function') or 'sum'
|
||
format_type = agg.get('format') or 'number'
|
||
# 确保 decimal_places 是整数,处理 None 和非数字情况
|
||
decimal_places_raw = agg.get('decimal_places')
|
||
decimal_places = int(decimal_places_raw) if decimal_places_raw is not None else 2
|
||
|
||
logger.info(f"聚合字段解析: name={name}, source={source}, field={field}, function={function}, format={format_type}, decimal_places={decimal_places}")
|
||
|
||
if not all([name, source, field]):
|
||
logger.warning(f"聚合字段配置不完整,跳过: name={name}, source={source}, field={field}")
|
||
continue
|
||
|
||
# 计算聚合值
|
||
raw_value = cls.calculate_aggregation(context, source, field, function)
|
||
logger.info(f"聚合原始值: {raw_value}, 类型: {type(raw_value)}")
|
||
|
||
# 格式化
|
||
formatted_value = cls.format_value(raw_value, format_type, decimal_places)
|
||
logger.info(f"格式化后: {formatted_value}, decimal_places={decimal_places}")
|
||
|
||
results[name] = formatted_value
|
||
context[name] = raw_value # 使用原始值用于后续计算
|
||
|
||
# 如果是中文格式,同时保存原始数值
|
||
if format_type == 'chinese':
|
||
results[f'{name}_raw'] = raw_value
|
||
|
||
logger.info(f"聚合计算完成: {name} = {formatted_value}")
|
||
|
||
except Exception as e:
|
||
logger.warning(f"聚合计算失败: {agg}, 错误: {e}")
|
||
|
||
# 2. 执行计算字段(按顺序,支持依赖)
|
||
fields = calculation_rules.get('fields', [])
|
||
logger.info(f"开始执行计算字段,共 {len(fields)} 个")
|
||
for field_config in fields:
|
||
try:
|
||
name = field_config.get('name')
|
||
formula = field_config.get('formula')
|
||
format_type = field_config.get('format') or 'number'
|
||
# 确保 decimal_places 是整数,处理 None 和非数字情况
|
||
decimal_places_raw = field_config.get('decimal_places')
|
||
decimal_places = int(decimal_places_raw) if decimal_places_raw is not None else 2
|
||
|
||
logger.info(f"处理计算字段: name={name}, formula={formula}, format={format_type}, decimal_places={decimal_places}")
|
||
|
||
if not all([name, formula]):
|
||
logger.warning(f"计算字段配置不完整,跳过: {field_config}")
|
||
continue
|
||
|
||
# 计算公式
|
||
raw_value = cls.evaluate_formula(formula, context)
|
||
logger.info(f"计算字段 {name} 原始值: {raw_value}")
|
||
|
||
if raw_value is not None:
|
||
# 格式化
|
||
formatted_value = cls.format_value(raw_value, format_type, decimal_places)
|
||
logger.info(f"计算字段 {name} 格式化后: {formatted_value}")
|
||
|
||
results[name] = formatted_value
|
||
context[name] = raw_value # 使用原始值用于后续计算
|
||
|
||
# 如果是中文格式,同时保存原始数值
|
||
if format_type == 'chinese':
|
||
results[f'{name}_raw'] = raw_value
|
||
|
||
logger.info(f"公式计算完成: {name} = {formatted_value}")
|
||
else:
|
||
logger.warning(f"计算字段 {name} 返回 None")
|
||
|
||
except Exception as e:
|
||
logger.error(f"公式计算失败: {field_config}, 错误: {e}", exc_info=True)
|
||
|
||
return results
|
||
|
||
|
||
# 导出
|
||
__all__ = ['CalculationEngine', 'number_to_chinese', 'safe_get_value']
|