#!/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']