From a8df8cd187b21df4324c702fa704fdd460d70020 Mon Sep 17 00:00:00 2001 From: YuWuKunCheng Date: Tue, 9 Jun 2026 12:18:12 +0800 Subject: [PATCH] =?UTF-8?q?=E7=A7=BB=E9=99=A4=20=E6=97=A0=E6=84=8F?= =?UTF-8?q?=E4=B9=89=E9=85=8D=E7=BD=AE=E9=A1=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- chan.py | 37 +-- chanlun-py/Cargo.toml | 4 +- chanlun-py/chanlun.pyi | 3 +- chanlun-py/chanlun/__init__.pyi | 3 +- chanlun-py/chanlun/chan.py | 37 +-- chanlun-py/pyproject.toml | 2 +- chanlun-py/src/config_py.rs | 45 ++-- chanlun-py/tests/test_all.py | 265 ++++++++++++++++++++ chanlun/src/business/synthesizer.rs | 29 ++- chanlun/src/config.rs | 359 ++++++++++++++++++++++++---- 10 files changed, 631 insertions(+), 153 deletions(-) diff --git a/chan.py b/chan.py index 4f247dc..265f5d9 100644 --- a/chan.py +++ b/chan.py @@ -929,19 +929,6 @@ class 缠论配置: 买卖点_指标匹配_MACD: bool = True, # 买在负,卖在正! 买卖点_指标匹配_KDJ: bool = True, # 买在死叉之后,卖在金叉之后 买卖点_指标匹配_RSI: bool = True, # 买在均线之下,卖在均线之上 - # 以下字段酌情废弃 - 买卖点_背离率: float = float("inf"), - 买卖点_T2_回调阈值: float = 1.0, - 买卖点_T2S_最大层级: int = 3, - 买卖点_峰值条件: bool = False, - 买卖点_计算方式: str = "峰", - 买卖点_计算线段BSP1: bool = True, - 买卖点_处理BSP2: bool = True, - 买卖点_计算线段BSP3: bool = True, - 买卖点_依赖T1: bool = True, - 买卖点_中枢来源: str = "合", - 买卖点_调试输出: bool = False, - # 以上字段酌情废弃 线段内部背驰_MACD: bool = True, 线段内部背驰_斜率: bool = True, 线段内部背驰_测度: bool = True, @@ -1018,17 +1005,6 @@ class 缠论配置: self.买卖点_指标匹配_MACD = 买卖点_指标匹配_MACD self.买卖点_指标匹配_KDJ = 买卖点_指标匹配_KDJ self.买卖点_指标匹配_RSI = 买卖点_指标匹配_RSI - self.买卖点_背离率 = 买卖点_背离率 - self.买卖点_T2_回调阈值 = 买卖点_T2_回调阈值 - self.买卖点_T2S_最大层级 = 买卖点_T2S_最大层级 - self.买卖点_峰值条件 = 买卖点_峰值条件 - self.买卖点_计算方式 = 买卖点_计算方式 - self.买卖点_计算线段BSP1 = 买卖点_计算线段BSP1 - self.买卖点_处理BSP2 = 买卖点_处理BSP2 - self.买卖点_计算线段BSP3 = 买卖点_计算线段BSP3 - self.买卖点_依赖T1 = 买卖点_依赖T1 - self.买卖点_中枢来源 = 买卖点_中枢来源 - self.买卖点_调试输出 = 买卖点_调试输出 self.线段内部背驰_MACD = 线段内部背驰_MACD self.线段内部背驰_斜率 = 线段内部背驰_斜率 self.线段内部背驰_测度 = 线段内部背驰_测度 @@ -1115,17 +1091,6 @@ class 缠论配置: "买卖点_指标匹配_MACD": {"annotation": bool, "default": True}, "买卖点_指标匹配_KDJ": {"annotation": bool, "default": True}, "买卖点_指标匹配_RSI": {"annotation": bool, "default": True}, - "买卖点_背离率": {"annotation": float, "default": float("inf")}, - "买卖点_T2_回调阈值": {"annotation": float, "default": 1.0}, - "买卖点_T2S_最大层级": {"annotation": int, "default": 3}, - "买卖点_峰值条件": {"annotation": bool, "default": False}, - "买卖点_计算方式": {"annotation": str, "default": "峰"}, - "买卖点_计算线段BSP1": {"annotation": bool, "default": True}, - "买卖点_处理BSP2": {"annotation": bool, "default": True}, - "买卖点_计算线段BSP3": {"annotation": bool, "default": True}, - "买卖点_依赖T1": {"annotation": bool, "default": True}, - "买卖点_中枢来源": {"annotation": str, "default": "合"}, - "买卖点_调试输出": {"annotation": bool, "default": False}, "线段内部背驰_MACD": {"annotation": bool, "default": True}, "线段内部背驰_斜率": {"annotation": bool, "default": True}, "线段内部背驰_测度": {"annotation": bool, "default": True}, @@ -6598,7 +6563,7 @@ class 观察者: 传入差异 = 缠论配置().对比(配置) 传入差异.update(差异) 配置 = 缠论配置(**传入差异) - print("加载异常配置+传入差异", 传入差异) + logger.info(f"加载异常配置+传入差异: {传入差异}") name = Path(文件路径).name.split(".")[0] 符号, 周期, 起始时间戳, 结束时间戳 = name.split("-") diff --git a/chanlun-py/Cargo.toml b/chanlun-py/Cargo.toml index ae97e71..13c81e8 100644 --- a/chanlun-py/Cargo.toml +++ b/chanlun-py/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "chanlun-py" -version = "26.6.45" +version = "26.6.47" edition = "2024" description = "缠论技术分析库 — Rust 高性能 Python 绑定" authors = ["YuYuKunKun"] @@ -12,7 +12,7 @@ crate-type = ["cdylib"] name = "chanlun" [dependencies] -chanlun = { path = "../chanlun" } +chanlun = "26.6.3" # { path = "../chanlun" } lru = "0.18" pyo3 = { version = "0.28", features = ["experimental-inspect"] } serde_json = "1" diff --git a/chanlun-py/chanlun.pyi b/chanlun-py/chanlun.pyi index 159a177..b7d4adf 100644 --- a/chanlun-py/chanlun.pyi +++ b/chanlun-py/chanlun.pyi @@ -828,7 +828,8 @@ class 缠论配置: def to_dict(self) -> Dict[str, Any]: ... def to_json(self) -> str: ... def 保存配置(self, path: str = "缠论配置.json") -> None: ... - def 对比(self, other: 缠论配置) -> Dict[str, Tuple[Any, Any]]: ... + def 对比(self, other: 缠论配置) -> Dict[str, Any]: ... + def model_copy(self, update: Optional[Dict[str, Any]] = None) -> 缠论配置: ... @classmethod def 加载配置(cls, path: str = "缠论配置.json") -> 缠论配置: ... @classmethod diff --git a/chanlun-py/chanlun/__init__.pyi b/chanlun-py/chanlun/__init__.pyi index 4a7da30..d2725be 100644 --- a/chanlun-py/chanlun/__init__.pyi +++ b/chanlun-py/chanlun/__init__.pyi @@ -853,7 +853,8 @@ class 缠论配置: def to_dict(self) -> Dict[str, Any]: ... def to_json(self) -> str: ... def 保存配置(self, path: str = "缠论配置.json") -> None: ... - def 对比(self, other: 缠论配置) -> Dict[str, Tuple[Any, Any]]: ... + def 对比(self, other: 缠论配置) -> Dict[str, Any]: ... + def model_copy(self, update: Optional[Dict[str, Any]] = None) -> 缠论配置: ... @classmethod def 加载配置(cls, path: str = "缠论配置.json") -> 缠论配置: ... @classmethod diff --git a/chanlun-py/chanlun/chan.py b/chanlun-py/chanlun/chan.py index ef8f72f..c76d2cd 100644 --- a/chanlun-py/chanlun/chan.py +++ b/chanlun-py/chanlun/chan.py @@ -929,19 +929,6 @@ class 缠论配置: 买卖点_指标匹配_MACD: bool = True, # 买在负,卖在正! 买卖点_指标匹配_KDJ: bool = True, # 买在死叉之后,卖在金叉之后 买卖点_指标匹配_RSI: bool = True, # 买在均线之下,卖在均线之上 - # 以下字段酌情废弃 - 买卖点_背离率: float = float("inf"), - 买卖点_T2_回调阈值: float = 1.0, - 买卖点_T2S_最大层级: int = 3, - 买卖点_峰值条件: bool = False, - 买卖点_计算方式: str = "峰", - 买卖点_计算线段BSP1: bool = True, - 买卖点_处理BSP2: bool = True, - 买卖点_计算线段BSP3: bool = True, - 买卖点_依赖T1: bool = True, - 买卖点_中枢来源: str = "合", - 买卖点_调试输出: bool = False, - # 以上字段酌情废弃 线段内部背驰_MACD: bool = True, 线段内部背驰_斜率: bool = True, 线段内部背驰_测度: bool = True, @@ -1018,17 +1005,6 @@ class 缠论配置: self.买卖点_指标匹配_MACD = 买卖点_指标匹配_MACD self.买卖点_指标匹配_KDJ = 买卖点_指标匹配_KDJ self.买卖点_指标匹配_RSI = 买卖点_指标匹配_RSI - self.买卖点_背离率 = 买卖点_背离率 - self.买卖点_T2_回调阈值 = 买卖点_T2_回调阈值 - self.买卖点_T2S_最大层级 = 买卖点_T2S_最大层级 - self.买卖点_峰值条件 = 买卖点_峰值条件 - self.买卖点_计算方式 = 买卖点_计算方式 - self.买卖点_计算线段BSP1 = 买卖点_计算线段BSP1 - self.买卖点_处理BSP2 = 买卖点_处理BSP2 - self.买卖点_计算线段BSP3 = 买卖点_计算线段BSP3 - self.买卖点_依赖T1 = 买卖点_依赖T1 - self.买卖点_中枢来源 = 买卖点_中枢来源 - self.买卖点_调试输出 = 买卖点_调试输出 self.线段内部背驰_MACD = 线段内部背驰_MACD self.线段内部背驰_斜率 = 线段内部背驰_斜率 self.线段内部背驰_测度 = 线段内部背驰_测度 @@ -1115,17 +1091,6 @@ class 缠论配置: "买卖点_指标匹配_MACD": {"annotation": bool, "default": True}, "买卖点_指标匹配_KDJ": {"annotation": bool, "default": True}, "买卖点_指标匹配_RSI": {"annotation": bool, "default": True}, - "买卖点_背离率": {"annotation": float, "default": float("inf")}, - "买卖点_T2_回调阈值": {"annotation": float, "default": 1.0}, - "买卖点_T2S_最大层级": {"annotation": int, "default": 3}, - "买卖点_峰值条件": {"annotation": bool, "default": False}, - "买卖点_计算方式": {"annotation": str, "default": "峰"}, - "买卖点_计算线段BSP1": {"annotation": bool, "default": True}, - "买卖点_处理BSP2": {"annotation": bool, "default": True}, - "买卖点_计算线段BSP3": {"annotation": bool, "default": True}, - "买卖点_依赖T1": {"annotation": bool, "default": True}, - "买卖点_中枢来源": {"annotation": str, "default": "合"}, - "买卖点_调试输出": {"annotation": bool, "default": False}, "线段内部背驰_MACD": {"annotation": bool, "default": True}, "线段内部背驰_斜率": {"annotation": bool, "default": True}, "线段内部背驰_测度": {"annotation": bool, "default": True}, @@ -6598,7 +6563,7 @@ class 观察者: 传入差异 = 缠论配置().对比(配置) 传入差异.update(差异) 配置 = 缠论配置(**传入差异) - print("加载异常配置+传入差异", 传入差异) + logger.info(f"加载异常配置+传入差异: {传入差异}") name = Path(文件路径).name.split(".")[0] 符号, 周期, 起始时间戳, 结束时间戳 = name.split("-") diff --git a/chanlun-py/pyproject.toml b/chanlun-py/pyproject.toml index e6c51e2..e1cca15 100644 --- a/chanlun-py/pyproject.toml +++ b/chanlun-py/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "maturin" [project] name = "chanlun" -version = "2606.45" +version = "2606.47" description = "缠论技术分析库 — Rust 高性能实现" readme = { file = "README.md", content-type = "text/markdown" } license = { file = "LICENSE", content-type = "text/plain" } diff --git a/chanlun-py/src/config_py.rs b/chanlun-py/src/config_py.rs index c7930b2..328557c 100644 --- a/chanlun-py/src/config_py.rs +++ b/chanlun-py/src/config_py.rs @@ -169,8 +169,11 @@ impl 缠论配置Py { /// 将配置导出为 Python 字典。 fn to_dict(&self, py: Python<'_>) -> PyResult> { let dict = PyDict::new(py); + let valid = chanlun::config::缠论配置::model_fields(); for (k, v) in &self.fields { - dict.set_item(k, v.clone_ref(py))?; + if valid.contains(&k.as_str()) { + dict.set_item(k, v.clone_ref(py))?; + } } Ok(dict.into()) } @@ -252,26 +255,36 @@ impl 缠论配置Py { Ok(result.into()) } - /// 比较当前配置与另一个配置的差异 - #[allow(clippy::type_complexity)] - fn 对比( - &self, - py: Python<'_>, - other: &Bound<'_, 缠论配置Py>, - ) -> PyResult, Py)>> { + /// 创建当前配置的拷贝并可选择更新字段(对应 Python model_copy(update={...}, deep=True)) + #[pyo3(signature = (update = None))] + fn model_copy(&self, py: Python<'_>, update: Option<&Bound<'_, PyDict>>) -> PyResult { + let current = self.to_dict(py)?; + if let Some(updates) = update { + for (key, value) in updates.iter() { + current.bind(py).set_item(key, value)?; + } + } + Self::from_dict(&py.get_type::(), current.bind(py)) + } + + /// 比较当前配置与另一个配置的差异(对应 Python 对比 → dict[字段名, 新值]) + fn 对比(&self, py: Python<'_>, other: &Bound<'_, 缠论配置Py>) -> PyResult> { let other_ref = other.borrow(); - let mut diff = HashMap::new(); - for (key, val) in &self.fields { - if let Some(other_val) = other_ref.fields.get(key) { - let a = val.clone_ref(py); + let dict = PyDict::new(py); + let valid = chanlun::config::缠论配置::model_fields(); + for key in valid { + if let (Some(self_val), Some(other_val)) = + (self.fields.get(*key), other_ref.fields.get(*key)) + { + let a = self_val.clone_ref(py); let b = other_val.clone_ref(py); let eq = a.bind(py).eq(b.bind(py))?; if !eq { - diff.insert(key.clone(), (val.clone_ref(py), other_val.clone_ref(py))); + dict.set_item(*key, b)?; } } } - Ok(diff) + Ok(dict.into()) } } @@ -393,9 +406,9 @@ fn validate_field( ) -> Result<(), String> { use serde_json::Value; - // 输入为 null → 跳过(保留默认) + // 输入为 null → 保留(对应 Optional/Infinity 字段) if input.is_null() { - return Err("值为 null".into()); + return Ok(()); } // 字符串字段:检查有效值白名单 diff --git a/chanlun-py/tests/test_all.py b/chanlun-py/tests/test_all.py index 2c861d2..39985ff 100644 --- a/chanlun-py/tests/test_all.py +++ b/chanlun-py/tests/test_all.py @@ -2511,5 +2511,270 @@ class Test立体分析器双端一致(unittest.TestCase): self.assertTrue(eq, msg) +class Test缠论配置双端一致(unittest.TestCase): + """缠论配置 to_dict / from_dict / model_copy 双端输出一致.""" + + def _make_configs(self): + import chanlun + from chanlun import chan + + cfg_rs = chanlun.缠论配置() + cfg_py = chan.缠论配置() + return cfg_rs, cfg_py + + def test_to_dict_keys_一致(self): + """to_dict 字段名集合双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + d_rs = cfg_rs.to_dict() + d_py = cfg_py.to_dict() + + self.assertEqual(set(d_rs.keys()), set(d_py.keys()), f"to_dict 字段不一致: R extra={set(d_rs.keys()) - set(d_py.keys())} P extra={set(d_py.keys()) - set(d_rs.keys())}") + + def test_to_dict_values_一致(self): + """to_dict 值双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + d_rs = cfg_rs.to_dict() + d_py = cfg_py.to_dict() + + mismatches = [] + for k in d_rs: + v_rs = d_rs[k] + v_py = d_py.get(k) + if v_rs is None and v_py is None: + continue + if v_rs != v_py: + mismatches.append(f" {k}: R={v_rs!r} P={v_py!r}") + self.assertEqual(len(mismatches), 0, f"to_dict 值不一致 ({len(mismatches)}处):\n" + "\n".join(mismatches[:10])) + + def test_to_json_content_一致(self): + """to_json 内容一致(JSON 解析后对比).""" + import chanlun + from chanlun import chan + import json + + cfg_rs, cfg_py = self._make_configs() + j_rs = json.loads(cfg_rs.to_json()) + j_py = json.loads(cfg_py.to_json()) + + self.assertEqual(j_rs, j_py, f"to_json 内容不一致") + + def test_from_dict_roundtrip_一致(self): + """from_dict → to_dict 往返双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + d = cfg_rs.to_dict() + cfg2_rs = chanlun.缠论配置.from_dict(d) + cfg2_py = chan.缠论配置.from_dict(d) + + d2_rs = cfg2_rs.to_dict() + d2_py = cfg2_py.to_dict() + for k in d2_rs: + self.assertEqual(d2_rs[k], d2_py.get(k), f"from_dict 往返不一致: {k}") + + def test_from_json_roundtrip_一致(self): + """from_json → to_dict 往返双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + j = cfg_rs.to_json() + cfg2_rs = chanlun.缠论配置.from_json(j) + cfg2_py = chan.缠论配置.from_json(j) + + # 验证标识和关键字段一致 + self.assertEqual(cfg2_rs.标识, cfg2_py.标识) + self.assertEqual(cfg2_rs.笔内元素数量, cfg2_py.笔内元素数量) + self.assertEqual(cfg2_rs.买卖点偏移, cfg2_py.买卖点偏移) + self.assertEqual(cfg2_rs.指标计算方式, cfg2_py.指标计算方式) + + def test_custom_values_from_dict_一致(self): + """自定义字段 from_dict 双端一致.""" + import chanlun + from chanlun import chan + + data = { + "标识": "custom_test", + "缠K合并替换": True, + "笔内元素数量": 8, + "笔弱化": True, + "计算指标": False, + "指标计算方式": "高低均值", + "平滑异同移动平均线_快线周期": 12, + "买卖点偏移": 3, + "买卖点激进识别": True, + } + cfg_rs = chanlun.缠论配置.from_dict(data) + cfg_py = chan.缠论配置.from_dict(data) + + d_rs = cfg_rs.to_dict() + d_py = cfg_py.to_dict() + for k in data: + self.assertEqual(d_rs.get(k), d_py.get(k), f"自定义字段 {k}: R={d_rs.get(k)} P={d_py.get(k)}") + + def test_model_copy_一致(self): + """model_copy 双端输出一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + update = {"标识": "copied", "推送K线": False, "笔内元素数量": 10} + + copy_rs = cfg_rs.model_copy(update) + copy_py = cfg_py.model_copy(update) + + self.assertEqual(copy_rs.标识, copy_py.标识) + self.assertEqual(copy_rs.笔内元素数量, copy_py.笔内元素数量) + self.assertFalse(copy_rs.推送K线) + self.assertFalse(copy_py.推送K线) + # 未更新字段保持原值一致 + self.assertEqual(copy_rs.买卖点偏移, copy_py.买卖点偏移) + + def test_from_dict_过滤未知字段_一致(self): + """from_dict 过滤未知字段(兼容旧版本配置)双端一致.""" + import chanlun + from chanlun import chan + + data = {"标识": "test", "笔内元素数量": 7, "废弃字段_已删除": 999, "另一个旧字段": "xxx"} + cfg_rs = chanlun.缠论配置.from_dict(data) + cfg_py = chan.缠论配置.from_dict(data) + + self.assertEqual(cfg_rs.标识, cfg_py.标识) + self.assertEqual(cfg_rs.笔内元素数量, cfg_py.笔内元素数量) + # 未知字段应被忽略,不影响构造 + d_rs = cfg_rs.to_dict() + self.assertNotIn("废弃字段_已删除", d_rs) + + def test_不推送_一致(self): + """不推送 静态方法双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs = chanlun.缠论配置.不推送() + cfg_py = chan.缠论配置.不推送() + + self.assertFalse(cfg_rs.推送K线) + self.assertFalse(cfg_py.推送K线) + self.assertFalse(cfg_rs.图表展示) + self.assertFalse(cfg_py.图表展示) + self.assertEqual(cfg_rs.笔内元素数量, cfg_py.笔内元素数量) + + def test_对比_默认一致(self): + """默认配置 self 对比应无差异.""" + import chanlun + from chanlun import chan + + cfg_rs_a, _ = self._make_configs() + cfg_rs_b = chanlun.缠论配置() + diff_rs = cfg_rs_a.对比(cfg_rs_b) + self.assertIsInstance(diff_rs, dict) + self.assertEqual(len(diff_rs), 0, "默认一致配置不应有差异") + + # Python side + cfg_py_a = chan.缠论配置() + cfg_py_b = chan.缠论配置() + diff_py = cfg_py_a.对比(cfg_py_b) + self.assertEqual(len(diff_py), 0) + + def test_对比_有差异字段一致(self): + """修改字段后 对比 双端输出一致.""" + import chanlun + from chanlun import chan + + cfg_rs_a, cfg_py_a = self._make_configs() + + # 构造有差异的配置 + update = {"标识": "changed", "笔内元素数量": 99, "推送K线": False} + cfg_rs_b = cfg_rs_a.model_copy(update) + cfg_py_b = cfg_py_a.model_copy(update) + + diff_rs = cfg_rs_a.对比(cfg_rs_b) + diff_py = cfg_py_a.对比(cfg_py_b) + + self.assertEqual(set(diff_rs.keys()), set(diff_py.keys()), f"对比字段不一致: R={set(diff_rs.keys())} P={set(diff_py.keys())}") + for k in diff_rs: + self.assertEqual(diff_rs[k], diff_py[k], f"对比[{k}] 值不一致: R={diff_rs[k]!r} P={diff_py[k]!r}") + + def test_对比_往返一致(self): + """to_dict → from_dict → 对比 应无差异.""" + import chanlun + from chanlun import chan + + cfg_rs, _ = self._make_configs() + d = cfg_rs.to_dict() + cfg2_rs = chanlun.缠论配置.from_dict(d) + diff = cfg_rs.对比(cfg2_rs) + self.assertEqual(len(diff), 0, f"Rust往返后对比不应有差异: {diff}") + + # Python side + cfg_py = chan.缠论配置() + d_py = cfg_py.to_dict() + cfg2_py = chan.缠论配置.from_dict(d_py) + diff_py = cfg_py.对比(cfg2_py) + self.assertEqual(len(diff_py), 0) + + def test_对比_与chan输出一致(self): + """对比 输出与 chan.对比 逐项一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + update = {"标识": "test_x", "缠K合并替换": True, "笔内元素数量": 7, "计算指标": False, "买卖点激进识别": True, "线段_修正": True} + alt_rs = cfg_rs.model_copy(update) + alt_py = cfg_py.model_copy(update) + + # Rust binding: cfg_rs.对比(alt_rs) + diff_rs = cfg_rs.对比(alt_rs) + # chan.py: cfg_py.对比(alt_py) + diff_py = cfg_py.对比(alt_py) + + self.assertEqual(diff_rs, diff_py, f"对比输出不一致:\n R={diff_rs}\n P={diff_py}") + + def test_对比_only_model_fields(self): + """对比 仅比较 model_fields 字段.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + update = {"标识": "only_test"} + alt_rs = cfg_rs.model_copy(update) + alt_py = cfg_py.model_copy(update) + + diff_rs = cfg_rs.对比(alt_rs) + diff_py = cfg_py.对比(alt_py) + + self.assertEqual(len(diff_rs), 1) + self.assertEqual(len(diff_py), 1) + self.assertIn("标识", diff_rs) + self.assertIn("标识", diff_py) + self.assertEqual(diff_rs["标识"], "only_test") + self.assertEqual(diff_py["标识"], "only_test") + + def test_对比_不推送_一致(self): + """不推送 配置与默认配置 对比 双端一致.""" + import chanlun + from chanlun import chan + + cfg_rs, cfg_py = self._make_configs() + muted_rs = chanlun.缠论配置.不推送() + muted_py = chan.缠论配置.不推送() + + diff_rs = cfg_rs.对比(muted_rs) + diff_py = cfg_py.对比(muted_py) + + self.assertEqual(set(diff_rs.keys()), set(diff_py.keys())) + # 不推送应关闭所有推送/图表字段 + for k in diff_rs: + self.assertFalse(diff_rs[k], f"不推送差异字段 {k} 应为 False") + self.assertFalse(diff_py[k], f"不推送差异字段 {k} 应为 False") + + if __name__ == "__main__": unittest.main() diff --git a/chanlun/src/business/synthesizer.rs b/chanlun/src/business/synthesizer.rs index b248566..0d79191 100644 --- a/chanlun/src/business/synthesizer.rs +++ b/chanlun/src/business/synthesizer.rs @@ -24,6 +24,10 @@ use crate::kline::bar::K线; use std::collections::HashMap; +use tracing; + +/// 事件回调类型 — fn(信号类型, 标识, 周期, 完成K线) +type 合成器事件回调 = Box; /// K线合成器 — 将小周期K线合成为大周期K线 pub struct K线合成器 { @@ -34,15 +38,13 @@ pub struct K线合成器 { /// 事件回调 — K线完成时触发,对应 Python K线合成器.事件回调 /// 签名: fn(信号类型: str, 标识: str, 周期: i64, 完成K线: K线) /// 在 _完成K线 清空当前K线后、新K线创建前触发 - 事件回调: Option>, + 事件回调: Option<合成器事件回调>, } impl K线合成器 { /// 创建K线合成器 — 对应 Python K线合成器.__init__(标识, 周期组, 事件回调=None) pub fn new( - 标识: String, - 周期组: Vec, - 事件回调: Option>, + 标识: String, 周期组: Vec, 事件回调: Option<合成器事件回调> ) -> Self { let mut 周期组 = 周期组; 周期组.sort(); @@ -64,10 +66,7 @@ impl K线合成器 { } /// 设置事件回调 — 对应 Python `设置事件回调` - pub fn 设置事件回调( - &mut self, - 回调: Box, - ) { + pub fn 设置事件回调(&mut self, 回调: 合成器事件回调) { self.事件回调 = Some(回调); } @@ -166,9 +165,21 @@ impl K线合成器 { } /// 产生完成K线信号 — 对应 Python `_产生完成K线信号` + /// 异常安全:若回调 panic,捕获并记录错误,不中断管线 fn _产生完成K线信号(&self, 周期: i64, 完成K线: K线) { if let Some(ref cb) = self.事件回调 { - cb("K线完成".into(), self.标识.clone(), 周期, 完成K线); + let 标识 = self.标识.clone(); + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + cb("K线完成".into(), 标识, 周期, 完成K线); + })); + if let Err(e) = result { + let msg = e + .downcast_ref::<&str>() + .map(|s| s.to_string()) + .or_else(|| e.downcast_ref::().cloned()) + .unwrap_or_else(|| "未知错误".into()); + tracing::error!("K线合成器 事件回调 异常: {}", msg); + } } } diff --git a/chanlun/src/config.rs b/chanlun/src/config.rs index 51142a5..eefe08a 100644 --- a/chanlun/src/config.rs +++ b/chanlun/src/config.rs @@ -23,12 +23,9 @@ */ use serde::{Deserialize, Deserializer, Serialize}; +use std::collections::HashMap; use tracing::warn; -fn is_infinite_f64(v: &f64) -> bool { - v.is_infinite() -} - /// 缠论配置 —— 控制所有分析阶段的行为 /// /// 50+ 参数集中控制缠K合并、笔/线段划分、中枢识别、买卖点生成等所有阶段。 @@ -211,29 +208,6 @@ pub struct 缠论配置 { pub 买卖点_指标匹配_KDJ: bool, /// 买卖点指标匹配 RSI pub 买卖点_指标匹配_RSI: bool, - /// 买卖点背离率阈值(Infinity 表示不使用) - #[serde(skip_serializing_if = "is_infinite_f64")] - pub 买卖点_背离率: f64, - /// 买卖点 T2 回调阈值 - pub 买卖点_T2_回调阈值: f64, - /// 买卖点 T2S 最大层级 - pub 买卖点_T2S_最大层级: i64, - /// 买卖点峰值条件 - pub 买卖点_峰值条件: bool, - /// 买卖点计算方式(峰/谷等) - pub 买卖点_计算方式: String, - /// 是否计算线段BSP1 - pub 买卖点_计算线段BSP1: bool, - /// 是否处理BSP2 - pub 买卖点_处理BSP2: bool, - /// 是否计算线段BSP3 - pub 买卖点_计算线段BSP3: bool, - /// 是否依赖T1买卖点 - pub 买卖点_依赖T1: bool, - /// 买卖点中枢来源(实/虚/合) - pub 买卖点_中枢来源: String, - /// 买卖点调试输出 - pub 买卖点_调试输出: bool, // ---- 背驰 ---- /// 线段内部背驰使用 MACD @@ -382,17 +356,6 @@ impl Default for 缠论配置 { 买卖点_指标匹配_MACD: true, 买卖点_指标匹配_KDJ: true, 买卖点_指标匹配_RSI: true, - 买卖点_背离率: f64::INFINITY, - 买卖点_T2_回调阈值: 1.0, - 买卖点_T2S_最大层级: 3, - 买卖点_峰值条件: false, - 买卖点_计算方式: "峰".into(), - 买卖点_计算线段BSP1: true, - 买卖点_处理BSP2: true, - 买卖点_计算线段BSP3: true, - 买卖点_依赖T1: true, - 买卖点_中枢来源: "合".into(), - 买卖点_调试输出: false, 线段内部背驰_MACD: true, 线段内部背驰_斜率: true, 线段内部背驰_测度: true, @@ -445,6 +408,147 @@ impl 缠论配置 { vec![("boll".into(), self.布林带_周期, self.布林带_标准差倍数)] } + /// 序列化为 JSON 字典(对应 Python to_dict,仅返回 model_fields 中的字段) + pub fn to_dict(&self) -> serde_json::Value { + let full = serde_json::to_value(self).unwrap_or_default(); + let valid = Self::model_fields(); + if let serde_json::Value::Object(map) = full { + let filtered: serde_json::Map<_, _> = map + .into_iter() + .filter(|(k, _)| valid.contains(&k.as_str())) + .collect(); + serde_json::Value::Object(filtered) + } else { + full + } + } + + /// 从 JSON 字典反序列化(对应 Python from_dict / 兼容旧版本配置) + pub fn from_dict(value: &serde_json::Value) -> Result { + if let serde_json::Value::Object(map) = value { + let valid_fields = Self::model_fields(); + let cleaned: serde_json::Map<_, _> = map + .iter() + .filter(|(k, _)| valid_fields.contains(&k.as_str())) + .map(|(k, v)| (k.clone(), v.clone())) + .collect(); + serde_json::from_value(serde_json::Value::Object(cleaned)) + } else { + serde_json::from_value(value.clone()) + } + } + + /// 验证并修正字段值(对应 Python _validate_all_fields) + pub fn _validate_all_fields(&mut self) { + const 允许: &[&str] = &[ + "开", + "高", + "低", + "收", + "高低均值", + "高低收均值", + "开高低收均值", + ]; + if !允许.contains(&self.指标计算方式.as_str()) { + warn!( + "[指标计算方式] = {} 值不在允许范围内,使用默认值:收", + self.指标计算方式 + ); + self.指标计算方式 = "收".into(); + } + } + + /// 返回字段名列表(对应 Python model_fields().keys()) + pub fn model_fields() -> &'static [&'static str] { + &[ + "标识", + "缠K合并替换", + "笔内元素数量", + "笔内相同终点取舍", + "笔内起始分型包含整笔", + "笔内起始分型包含整笔_包括右", + "笔内原始K线包含整笔", + "笔次级成笔", + "笔弱化", + "笔弱化_原始数量", + "线段_非缺口下穿刺", + "线段_特征序列忽视老阴老阳", + "线段_缺口后紧急修正", + "线段_修正", + "线段内部中枢图显", + "扩展线段_当下分析", + "分析笔", + "分析线段", + "分析扩展线段", + "分析笔中枢", + "分析线段中枢", + "手动终止", + "计算指标", + "计算BOLL", + "指标计算方式", + "平滑异同移动平均线_快线周期", + "平滑异同移动平均线_慢线周期", + "平滑异同移动平均线_信号周期", + "相对强弱指数_周期", + "相对强弱指数_移动平均线周期", + "相对强弱指数_超买阈值", + "相对强弱指数_超卖阈值", + "随机指标_RSV周期", + "随机指标_K值平滑周期", + "随机指标_D值平滑周期", + "随机指标_超买阈值", + "随机指标_超卖阈值", + "布林带_周期", + "布林带_标准差倍数", + "MACD_参数列表", + "RSI_周期列表", + "KDJ_参数列表", + "BOLL_参数列表", + "均线_类型列表", + "均线_周期列表", + "图表展示", + "推送K线", + "推送笔", + "推送线段", + "推送中枢", + "图表展示_笔", + "图表展示_线段", + "图表展示_扩展线段", + "图表展示_扩展线段_线段", + "图表展示_线段_线段", + "图表展示_中枢_笔", + "图表展示_中枢_线段", + "图表展示_中枢_扩展线段", + "图表展示_中枢_扩展线段_线段", + "图表展示_中枢_线段_线段", + "图表展示_中枢_线段内部", + "买卖点偏移", + "买卖点激进识别", + "买卖点与MACD柱强相关", + "买卖点错过误差值", + "买卖点_指标模式", + "买卖点_指标匹配_MACD", + "买卖点_指标匹配_KDJ", + "买卖点_指标匹配_RSI", + "线段内部背驰_MACD", + "线段内部背驰_斜率", + "线段内部背驰_测度", + "线段内部背驰_模式", + "加载文件路径", + ] + } + + /// 深拷贝并更新指定字段(对应 Python model_copy(update={...}, deep=True)) + pub fn model_copy(&self, update: &HashMap) -> Self { + let mut value = serde_json::to_value(self).unwrap_or_default(); + if let serde_json::Value::Object(ref mut map) = value { + for (k, v) in update { + map.insert(k.clone(), v.clone()); + } + } + serde_json::from_value(value).unwrap_or_else(|_| self.clone()) + } + /// 序列化为 JSON 字符串 pub fn to_json(&self) -> String { serde_json::to_string_pretty(self).unwrap_or_default() @@ -467,9 +571,10 @@ impl 缠论配置 { Ok(config) } - /// 返回一个关闭所有推送/显示的新配置 + /// 返回一个关闭所有推送/显示的新配置(对应 Python 不推送) pub fn 不推送(&self) -> Self { Self { + 线段内部中枢图显: false, 图表展示: false, 推送K线: false, 推送笔: false, @@ -524,19 +629,21 @@ impl 缠论配置 { result } - /// 对比两个配置,返回差异字段 - pub fn 对比(&self, other: &Self) -> Vec { - let mut diffs = Vec::new(); - let self_json = serde_json::to_value(self).unwrap(); - let other_json = serde_json::to_value(other).unwrap(); + /// 对比两个配置,返回差异字段及新值(对应 Python 对比 → dict[字段名, 新值]) + pub fn 对比(&self, other: &Self) -> HashMap { + let mut diffs = HashMap::new(); + let self_dict = self.to_dict(); + let other_dict = other.to_dict(); if let (serde_json::Value::Object(self_map), serde_json::Value::Object(other_map)) = - (&self_json, &other_json) + (&self_dict, &other_dict) { - for (key, self_val) in self_map { - if let Some(other_val) = other_map.get(key) - && self_val != other_val + for key in Self::model_fields() { + let self_val = self_map.get(*key); + let other_val = other_map.get(*key); + if self_val != other_val + && let Some(v) = other_val { - diffs.push(key.clone()); + diffs.insert(key.to_string(), v.clone()); } } } @@ -562,7 +669,6 @@ mod tests { let config = 缠论配置::default(); assert_eq!(config.标识, "bar"); assert_eq!(config.笔内元素数量, 5); - assert!(config.买卖点_背离率.is_infinite()); assert_eq!(config.指标计算方式, "收"); } @@ -609,6 +715,66 @@ mod tests { assert_eq!(config.线段内部背驰_模式, "全量"); } + #[test] + fn test_to_dict_roundtrip() { + let config = 缠论配置::default(); + let dict = config.to_dict(); + let restored = 缠论配置::from_dict(&dict).unwrap(); + assert_eq!(config.to_json(), restored.to_json()); + } + + #[test] + fn test_from_dict_filters_unknown_fields() { + // 兼容旧版本配置 — unknown fields are silently dropped + let json = serde_json::json!({ + "标识": "test", + "不存在的字段": 42, + "另一个废弃字段": "xxx", + "笔内元素数量": 8, + }); + let config = 缠论配置::from_dict(&json).unwrap(); + assert_eq!(config.标识, "test"); + assert_eq!(config.笔内元素数量, 8); + // 未指定字段使用默认值 + assert_eq!(config.买卖点偏移, 1); + } + + #[test] + fn test_model_fields_contains_all() { + let fields = 缠论配置::model_fields(); + assert!(fields.contains(&"标识")); + assert!(fields.contains(&"笔内元素数量")); + assert!(fields.contains(&"买卖点偏移")); + assert!(fields.contains(&"线段内部背驰_MACD")); + } + + #[test] + fn test_model_copy() { + let mut update = std::collections::HashMap::new(); + update.insert("标识".into(), serde_json::json!("custom")); + update.insert("推送K线".into(), serde_json::json!(false)); + update.insert("笔内元素数量".into(), serde_json::json!(10)); + + let config = 缠论配置::default(); + let copied = config.model_copy(&update); + + assert_eq!(copied.标识, "custom"); + assert!(!copied.推送K线); + assert_eq!(copied.笔内元素数量, 10); + // 未指定字段保持不变 + assert_eq!(copied.买卖点偏移, 1); + assert!(copied.推送笔); + } + + #[test] + fn test_to_dict_to_json_consistency() { + let config = 缠论配置::default(); + let dict = config.to_dict(); + // to_dict → from_dict → to_json should equal original to_json + let restored = 缠论配置::from_dict(&dict).unwrap(); + assert_eq!(config.to_json(), restored.to_json()); + } + #[test] fn test_不推送() { let config = 缠论配置::default(); @@ -616,7 +782,98 @@ mod tests { assert!(!muted.推送K线); assert!(!muted.推送笔); assert!(!muted.图表展示); - // 其他字段不变 assert_eq!(muted.笔内元素数量, 5); } + + #[test] + fn test_对比_无差异() { + let a = 缠论配置::default(); + let b = 缠论配置::default(); + let diff = a.对比(&b); + assert!(diff.is_empty(), "identical configs should have empty diff"); + } + + #[test] + fn test_对比_有差异() { + let mut a = 缠论配置::default(); + let mut b = 缠论配置::default(); + b.标识 = "changed".into(); + b.笔内元素数量 = 99; + + let diff = a.对比(&b); + assert_eq!(diff.len(), 2); + assert_eq!(diff.get("标识").unwrap().as_str().unwrap(), "changed"); + assert_eq!(diff.get("笔内元素数量").unwrap().as_i64().unwrap(), 99); + } + + #[test] + fn test_对比_仅比较model_fields() { + // 仅比较 model_fields 中的字段(Python 一致行为) + let a = 缠论配置::default(); + let b = 缠论配置::default(); + let diff = a.对比(&b); + // 验证不包含废弃字段(如已删除的 "买卖点_背离率" 等) + assert!(!diff.contains_key("买卖点_背离率")); + assert!(diff.is_empty(), "default configs should have no diff"); + } + + #[test] + fn test_to_dict_excludes_non_model_fields() { + let config = 缠论配置::default(); + let dict = config.to_dict(); + let valid = 缠论配置::model_fields(); + if let serde_json::Value::Object(map) = &dict { + for key in map.keys() { + assert!( + valid.contains(&key.as_str()), + "{key} should not be in to_dict output" + ); + } + } + assert_eq!( + valid.len(), + dict.as_object().map(|m| m.len()).unwrap_or(0), + "to_dict should have exactly model_fields count" + ); + } + + #[test] + fn test_model_copy_then_对比() { + let config = 缠论配置::default(); + let mut update = HashMap::new(); + update.insert("标识".into(), serde_json::json!("copied")); + update.insert("笔内元素数量".into(), serde_json::json!(10)); + + let copied = config.model_copy(&update); + let diff = config.对比(&copied); + + assert_eq!(diff.len(), 2); + assert_eq!(diff["标识"].as_str().unwrap(), "copied"); + assert_eq!(diff["笔内元素数量"].as_i64().unwrap(), 10); + } + + #[test] + fn test_对比_boolean_difference() { + let mut a = 缠论配置::default(); + let mut b = 缠论配置::default(); + b.推送K线 = false; + b.图表展示 = false; + + let diff = a.对比(&b); + assert_eq!(diff.len(), 2); + assert_eq!(diff["推送K线"], serde_json::json!(false)); + assert_eq!(diff["图表展示"], serde_json::json!(false)); + } + + #[test] + fn test_to_dict_from_dict_对比_roundtrip() { + let config = 缠论配置::default(); + let dict = config.to_dict(); + let restored = 缠论配置::from_dict(&dict).unwrap(); + let diff = config.对比(&restored); + assert!( + diff.is_empty(), + "to_dict→from_dict roundtrip should produce no diff" + ); + } }