Normalize fullwidth punctuation to ASCII across codebase.

Add scripts/normalize_ambiguous_unicode.py; fix corrupted patch_instance_theme_templates.py. Preserves curly quotes in string literals; removes Git homoglyph warnings on .env.example.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
dekun
2026-07-08 23:42:26 +08:00
parent aaa72c7961
commit b733e551a0
392 changed files with 71522 additions and 71369 deletions
File diff suppressed because it is too large Load Diff
+63 -63
View File
@@ -1,63 +1,63 @@
"""AI 复盘 journal 文本格式化三所共用)。"""
from __future__ import annotations
import sqlite3
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.ai.ai_review_lib import journal_row_lines_for_ai # noqa: E402
class TestAiReviewLib(unittest.TestCase):
def test_journal_row_includes_expect_and_actual_rr(self):
text = journal_row_lines_for_ai(
1,
{
"coin": "HYPE",
"tf": "5m",
"pnl": "10.73",
"real_rr": "2.1354",
"expect_rr": "-",
"entry_reason": "趋势回调",
"exit_reason": "移动止盈",
"hold_duration": "1天 3小时",
"mood_issues": "",
"post_breakeven_stare": "",
"new_trade_while_occupied": "",
"note": "测试备注",
},
)
self.assertIn("实际RR:2.1354", text)
self.assertIn("预期RR:-", text)
self.assertIn("开仓逻辑趋势回调", text)
self.assertIn("备注测试备注", text)
self.assertNotIn("开仓类型", text)
def test_journal_row_accepts_sqlite_row(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE journal_entries (
coin TEXT, tf TEXT, pnl TEXT, real_rr TEXT, expect_rr TEXT,
entry_reason TEXT, exit_reason TEXT, hold_duration TEXT,
mood_issues TEXT, mood_score INTEGER, note TEXT
)"""
)
conn.execute(
"""INSERT INTO journal_entries VALUES (?,?,?,?,?,?,?,?,?,?,?)""",
("BTC", "15m", "5", "1.2", "2.0", "突破", "止盈", "2小时", "", None, ""),
)
row = conn.execute("SELECT * FROM journal_entries").fetchone()
conn.close()
text = journal_row_lines_for_ai(1, row)
self.assertIn("BTC 15m", text)
self.assertIn("实际RR:1.2", text)
self.assertIn("开仓逻辑突破", text)
if __name__ == "__main__":
unittest.main()
"""AI 复盘 journal 文本格式化(三所共用)."""
from __future__ import annotations
import sqlite3
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.ai.ai_review_lib import journal_row_lines_for_ai # noqa: E402
class TestAiReviewLib(unittest.TestCase):
def test_journal_row_includes_expect_and_actual_rr(self):
text = journal_row_lines_for_ai(
1,
{
"coin": "HYPE",
"tf": "5m",
"pnl": "10.73",
"real_rr": "2.1354",
"expect_rr": "-",
"entry_reason": "趋势回调",
"exit_reason": "移动止盈",
"hold_duration": "1天 3小时",
"mood_issues": "",
"post_breakeven_stare": "",
"new_trade_while_occupied": "",
"note": "测试备注",
},
)
self.assertIn("实际RR:2.1354", text)
self.assertIn("预期RR:-", text)
self.assertIn("开仓逻辑:趋势回调", text)
self.assertIn("备注:测试备注", text)
self.assertNotIn("开仓类型", text)
def test_journal_row_accepts_sqlite_row(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE journal_entries (
coin TEXT, tf TEXT, pnl TEXT, real_rr TEXT, expect_rr TEXT,
entry_reason TEXT, exit_reason TEXT, hold_duration TEXT,
mood_issues TEXT, mood_score INTEGER, note TEXT
)"""
)
conn.execute(
"""INSERT INTO journal_entries VALUES (?,?,?,?,?,?,?,?,?,?,?)""",
("BTC", "15m", "5", "1.2", "2.0", "突破", "止盈", "2小时", "", None, ""),
)
row = conn.execute("SELECT * FROM journal_entries").fetchone()
conn.close()
text = journal_row_lines_for_ai(1, row)
self.assertIn("BTC 15m", text)
self.assertIn("实际RR:1.2", text)
self.assertIn("开仓逻辑:突破", text)
if __name__ == "__main__":
unittest.main()
+182 -182
View File
@@ -1,182 +1,182 @@
import unittest
from lib.trade.entry_model_lib import (
ENTRY_MODEL_BIG_DIV_A,
ENTRY_MODEL_BIG_DIV_B,
ENTRY_MODEL_LAUNCH_A,
ENTRY_MODEL_LAUNCH_B,
ENTRY_MODEL_SMALL_DIV,
ENTRY_CATEGORY_REVERSAL,
ENTRY_CATEGORY_TREND,
build_intraday_entry_reason_options,
build_trend_div_entry_reason_options,
entry_model_categories,
entry_model_category,
entry_model_display_label,
entry_model_label,
format_entry_type_display,
hub_meta_entry_context,
intraday_entry_model_options,
is_intraday_trading_profile,
open_position_button_label,
parse_manual_order_style_fields,
resolve_trade_record_entry_reason,
trade_style_for_entry_model,
trend_manual_entry_reason_count,
)
from lib.trade.trade_policy_lib import TradePolicy, load_trade_policy
class TestEntryModelLib(unittest.TestCase):
def test_intraday_profile_btc_eth_whitelist(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
self.assertTrue(is_intraday_trading_profile(policy))
self.assertEqual(trend_manual_entry_reason_count(policy), 2)
def test_trend_div_profile_alt(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "false",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
self.assertFalse(is_intraday_trading_profile(policy))
self.assertEqual(trend_manual_entry_reason_count(policy), 7)
def test_entry_model_maps_trade_style(self):
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_LAUNCH_A), "trend")
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_BIG_DIV_A), "trend")
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_SMALL_DIV), "swing")
self.assertEqual(entry_model_label(ENTRY_MODEL_LAUNCH_B), "启动B")
self.assertEqual(entry_model_category(ENTRY_MODEL_LAUNCH_A), ENTRY_CATEGORY_REVERSAL)
self.assertEqual(entry_model_category(ENTRY_MODEL_BIG_DIV_B), ENTRY_CATEGORY_TREND)
def test_entry_model_categories_two_level(self):
cats = entry_model_categories()
keys = [c["key"] for c in cats]
self.assertEqual(keys, ["reversal", "trend", "swing"])
reversal = cats[0]["options"]
self.assertEqual([o["code"] for o in reversal], ["launch_a", "launch_b"])
self.assertEqual(len(cats[2]["options"]), 1)
def test_parse_trend_div_requires_entry_model(self):
policy = TradePolicy(False, "both", False, ())
style, code, err = parse_manual_order_style_fields(policy, {})
self.assertTrue(err)
self.assertEqual(code, None)
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": ENTRY_MODEL_SMALL_DIV}
)
self.assertIsNone(err)
self.assertEqual(code, ENTRY_MODEL_SMALL_DIV)
self.assertEqual(style, "swing")
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": ENTRY_MODEL_LAUNCH_A}
)
self.assertIsNone(err)
self.assertEqual(code, ENTRY_MODEL_LAUNCH_A)
self.assertEqual(style, "trend")
def test_hub_meta_intraday(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ctx = hub_meta_entry_context(policy)
self.assertTrue(ctx["intraday_discipline"])
self.assertEqual(ctx["order_entry_profile"], "intraday")
policy = TradePolicy(False, "both", True, ("BTC", "ETH"))
style, code, err = parse_manual_order_style_fields(policy, {"trade_style": "swing"})
self.assertIsNone(err)
self.assertIsNone(code)
self.assertEqual(style, "swing")
def test_intraday_entry_model_options(self):
opts = intraday_entry_model_options()
codes = [o.code for o in opts]
self.assertEqual(codes, ["liquidity_false_break", "structure_breakout"])
self.assertEqual(entry_model_label("liquidity_false_break"), "假破")
def test_parse_intraday_requires_entry_model(self):
policy = TradePolicy(True, "both", True, ("BTC", "ETH"))
style, code, err = parse_manual_order_style_fields(policy, {})
self.assertTrue(err)
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": "structure_breakout"}
)
self.assertIsNone(err)
self.assertEqual(code, "structure_breakout")
self.assertEqual(style, "trend")
def test_open_position_button_intraday(self):
policy = TradePolicy(True, "both", True, ("BTC", "ETH"))
self.assertEqual(
open_position_button_label(policy, "full_margin"),
"开仓日内·全仓杠杆",
)
def test_resolve_entry_reason_from_model(self):
er = resolve_trade_record_entry_reason(entry_model=ENTRY_MODEL_BIG_DIV_B)
self.assertEqual(er, "顺势/大分歧B")
er2 = resolve_trade_record_entry_reason(entry_model=ENTRY_MODEL_LAUNCH_A)
self.assertEqual(er2, "反转/启动A")
def test_entry_model_display_label(self):
self.assertEqual(entry_model_display_label(ENTRY_MODEL_LAUNCH_A), "反转/启动A")
self.assertEqual(entry_model_display_label(ENTRY_MODEL_SMALL_DIV), "波段单/小分歧")
self.assertEqual(entry_model_display_label("liquidity_false_break"), "波段单/假破")
self.assertEqual(format_entry_type_display("启动A"), "反转/启动A")
self.assertEqual(entry_model_label(ENTRY_MODEL_LAUNCH_B), "启动B")
def test_resolve_entry_reason_trade_style_fallback(self):
er = resolve_trade_record_entry_reason(trade_style="swing")
self.assertEqual(er, "波段单")
er2 = resolve_trade_record_entry_reason(trade_style="trend")
self.assertEqual(er2, "趋势单")
def test_build_trend_div_journal_options(self):
opts = build_trend_div_entry_reason_options(("趋势回调",))
self.assertEqual(opts[:5], ("反转/启动A", "反转/启动B", "顺势/大分歧A", "顺势/大分歧B", "波段单/小分歧"))
self.assertIn("趋势单", opts)
self.assertIn("波段单", opts)
self.assertIn("趋势回调", opts)
def test_build_intraday_journal_options_only_four(self):
opts = build_intraday_entry_reason_options(
(
"关键位箱体突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
),
("趋势回调", "顺势加仓"),
)
self.assertEqual(
opts,
(
"波段单/假破",
"波段单/结构突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
),
)
def test_normalize_review_entry_reason(self):
from lib.trade.entry_model_lib import normalize_review_entry_reason
allowed = build_trend_div_entry_reason_options(())
self.assertEqual(normalize_review_entry_reason("反转/启动A", allowed), "反转/启动A")
self.assertEqual(normalize_review_entry_reason("启动A", allowed), "反转/启动A")
self.assertEqual(normalize_review_entry_reason("趋势单", allowed), "趋势单")
if __name__ == "__main__":
unittest.main()
import unittest
from lib.trade.entry_model_lib import (
ENTRY_MODEL_BIG_DIV_A,
ENTRY_MODEL_BIG_DIV_B,
ENTRY_MODEL_LAUNCH_A,
ENTRY_MODEL_LAUNCH_B,
ENTRY_MODEL_SMALL_DIV,
ENTRY_CATEGORY_REVERSAL,
ENTRY_CATEGORY_TREND,
build_intraday_entry_reason_options,
build_trend_div_entry_reason_options,
entry_model_categories,
entry_model_category,
entry_model_display_label,
entry_model_label,
format_entry_type_display,
hub_meta_entry_context,
intraday_entry_model_options,
is_intraday_trading_profile,
open_position_button_label,
parse_manual_order_style_fields,
resolve_trade_record_entry_reason,
trade_style_for_entry_model,
trend_manual_entry_reason_count,
)
from lib.trade.trade_policy_lib import TradePolicy, load_trade_policy
class TestEntryModelLib(unittest.TestCase):
def test_intraday_profile_btc_eth_whitelist(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
self.assertTrue(is_intraday_trading_profile(policy))
self.assertEqual(trend_manual_entry_reason_count(policy), 2)
def test_trend_div_profile_alt(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "false",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
self.assertFalse(is_intraday_trading_profile(policy))
self.assertEqual(trend_manual_entry_reason_count(policy), 7)
def test_entry_model_maps_trade_style(self):
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_LAUNCH_A), "trend")
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_BIG_DIV_A), "trend")
self.assertEqual(trade_style_for_entry_model(ENTRY_MODEL_SMALL_DIV), "swing")
self.assertEqual(entry_model_label(ENTRY_MODEL_LAUNCH_B), "启动B")
self.assertEqual(entry_model_category(ENTRY_MODEL_LAUNCH_A), ENTRY_CATEGORY_REVERSAL)
self.assertEqual(entry_model_category(ENTRY_MODEL_BIG_DIV_B), ENTRY_CATEGORY_TREND)
def test_entry_model_categories_two_level(self):
cats = entry_model_categories()
keys = [c["key"] for c in cats]
self.assertEqual(keys, ["reversal", "trend", "swing"])
reversal = cats[0]["options"]
self.assertEqual([o["code"] for o in reversal], ["launch_a", "launch_b"])
self.assertEqual(len(cats[2]["options"]), 1)
def test_parse_trend_div_requires_entry_model(self):
policy = TradePolicy(False, "both", False, ())
style, code, err = parse_manual_order_style_fields(policy, {})
self.assertTrue(err)
self.assertEqual(code, None)
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": ENTRY_MODEL_SMALL_DIV}
)
self.assertIsNone(err)
self.assertEqual(code, ENTRY_MODEL_SMALL_DIV)
self.assertEqual(style, "swing")
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": ENTRY_MODEL_LAUNCH_A}
)
self.assertIsNone(err)
self.assertEqual(code, ENTRY_MODEL_LAUNCH_A)
self.assertEqual(style, "trend")
def test_hub_meta_intraday(self):
policy = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ctx = hub_meta_entry_context(policy)
self.assertTrue(ctx["intraday_discipline"])
self.assertEqual(ctx["order_entry_profile"], "intraday")
policy = TradePolicy(False, "both", True, ("BTC", "ETH"))
style, code, err = parse_manual_order_style_fields(policy, {"trade_style": "swing"})
self.assertIsNone(err)
self.assertIsNone(code)
self.assertEqual(style, "swing")
def test_intraday_entry_model_options(self):
opts = intraday_entry_model_options()
codes = [o.code for o in opts]
self.assertEqual(codes, ["liquidity_false_break", "structure_breakout"])
self.assertEqual(entry_model_label("liquidity_false_break"), "假破")
def test_parse_intraday_requires_entry_model(self):
policy = TradePolicy(True, "both", True, ("BTC", "ETH"))
style, code, err = parse_manual_order_style_fields(policy, {})
self.assertTrue(err)
style, code, err = parse_manual_order_style_fields(
policy, {"entry_model": "structure_breakout"}
)
self.assertIsNone(err)
self.assertEqual(code, "structure_breakout")
self.assertEqual(style, "trend")
def test_open_position_button_intraday(self):
policy = TradePolicy(True, "both", True, ("BTC", "ETH"))
self.assertEqual(
open_position_button_label(policy, "full_margin"),
"开仓(日内·全仓杠杆)",
)
def test_resolve_entry_reason_from_model(self):
er = resolve_trade_record_entry_reason(entry_model=ENTRY_MODEL_BIG_DIV_B)
self.assertEqual(er, "顺势/大分歧B")
er2 = resolve_trade_record_entry_reason(entry_model=ENTRY_MODEL_LAUNCH_A)
self.assertEqual(er2, "反转/启动A")
def test_entry_model_display_label(self):
self.assertEqual(entry_model_display_label(ENTRY_MODEL_LAUNCH_A), "反转/启动A")
self.assertEqual(entry_model_display_label(ENTRY_MODEL_SMALL_DIV), "波段单/小分歧")
self.assertEqual(entry_model_display_label("liquidity_false_break"), "波段单/假破")
self.assertEqual(format_entry_type_display("启动A"), "反转/启动A")
self.assertEqual(entry_model_label(ENTRY_MODEL_LAUNCH_B), "启动B")
def test_resolve_entry_reason_trade_style_fallback(self):
er = resolve_trade_record_entry_reason(trade_style="swing")
self.assertEqual(er, "波段单")
er2 = resolve_trade_record_entry_reason(trade_style="trend")
self.assertEqual(er2, "趋势单")
def test_build_trend_div_journal_options(self):
opts = build_trend_div_entry_reason_options(("趋势回调",))
self.assertEqual(opts[:5], ("反转/启动A", "反转/启动B", "顺势/大分歧A", "顺势/大分歧B", "波段单/小分歧"))
self.assertIn("趋势单", opts)
self.assertIn("波段单", opts)
self.assertIn("趋势回调", opts)
def test_build_intraday_journal_options_only_four(self):
opts = build_intraday_entry_reason_options(
(
"关键位箱体突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
),
("趋势回调", "顺势加仓"),
)
self.assertEqual(
opts,
(
"波段单/假破",
"波段单/结构突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
),
)
def test_normalize_review_entry_reason(self):
from lib.trade.entry_model_lib import normalize_review_entry_reason
allowed = build_trend_div_entry_reason_options(())
self.assertEqual(normalize_review_entry_reason("反转/启动A", allowed), "反转/启动A")
self.assertEqual(normalize_review_entry_reason("启动A", allowed), "反转/启动A")
self.assertEqual(normalize_review_entry_reason("趋势单", allowed), "趋势单")
if __name__ == "__main__":
unittest.main()
+44 -44
View File
@@ -1,44 +1,44 @@
"""gate_transfer_lib 单元测试"""
from __future__ import annotations
import sqlite3
import unittest
from lib.exchange.gate_transfer_lib import count_auto_transfer_blockers
class GateTransferLibTest(unittest.TestCase):
def test_counts_order_monitors_first(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO order_monitors VALUES ('active')")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 1)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 1)
self.assertEqual(n, 1)
conn.close()
def test_counts_trend_plan_when_no_order_monitors(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 1)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 0)
self.assertEqual(n, 1)
conn.close()
def test_ignores_trend_plan_without_first_order(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 0)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 0)
self.assertEqual(n, 0)
conn.close()
if __name__ == "__main__":
unittest.main()
"""gate_transfer_lib 单元测试."""
from __future__ import annotations
import sqlite3
import unittest
from lib.exchange.gate_transfer_lib import count_auto_transfer_blockers
class GateTransferLibTest(unittest.TestCase):
def test_counts_order_monitors_first(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO order_monitors VALUES ('active')")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 1)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 1)
self.assertEqual(n, 1)
conn.close()
def test_counts_trend_plan_when_no_order_monitors(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 1)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 0)
self.assertEqual(n, 1)
conn.close()
def test_ignores_trend_plan_without_first_order(self):
conn = sqlite3.connect(":memory:")
conn.execute("CREATE TABLE order_monitors (status TEXT)")
conn.execute("CREATE TABLE trend_pullback_plans (status TEXT, first_order_done INTEGER)")
conn.execute("INSERT INTO trend_pullback_plans VALUES ('active', 0)")
conn.commit()
n = count_auto_transfer_blockers(conn, count_order_monitors=lambda c: 0)
self.assertEqual(n, 0)
conn.close()
if __name__ == "__main__":
unittest.main()
+1 -1
View File
@@ -1,4 +1,4 @@
"""history_window_lib 单元测试"""
"""history_window_lib 单元测试."""
from __future__ import annotations
import unittest
+32 -32
View File
@@ -1,32 +1,32 @@
"""子代理持仓三所开仓价字段统一解析"""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "manual_trading_hub"))
from agent import _position_entry_price # noqa: E402
class TestHubAgentEntryPrice(unittest.TestCase):
def test_binance_entry_price(self):
px = _position_entry_price({"entryPrice": 65851.6, "info": {}})
self.assertAlmostEqual(px, 65851.6)
def test_okx_avg_px(self):
px = _position_entry_price({"info": {"avgPx": "72.731"}})
self.assertAlmostEqual(px, 72.731)
def test_gate_info_entry(self):
px = _position_entry_price({"info": {"entry_price": "0.2232"}})
self.assertAlmostEqual(px, 0.2232)
def test_missing_returns_none(self):
self.assertIsNone(_position_entry_price({"info": {}}))
if __name__ == "__main__":
unittest.main()
"""子代理持仓:三所开仓价字段统一解析."""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "manual_trading_hub"))
from agent import _position_entry_price # noqa: E402
class TestHubAgentEntryPrice(unittest.TestCase):
def test_binance_entry_price(self):
px = _position_entry_price({"entryPrice": 65851.6, "info": {}})
self.assertAlmostEqual(px, 65851.6)
def test_okx_avg_px(self):
px = _position_entry_price({"info": {"avgPx": "72.731"}})
self.assertAlmostEqual(px, 72.731)
def test_gate_info_entry(self):
px = _position_entry_price({"info": {"entry_price": "0.2232"}})
self.assertAlmostEqual(px, 0.2232)
def test_missing_returns_none(self):
self.assertIsNone(_position_entry_price({"info": {}}))
if __name__ == "__main__":
unittest.main()
+94 -94
View File
@@ -1,94 +1,94 @@
"""子代理持仓三所标记价字段统一解析"""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "manual_trading_hub"))
from agent import _position_mark_price, _ticker_mark_price # noqa: E402
sys.path.insert(0, str(ROOT))
from lib.hub.hub_position_metrics import ( # noqa: E402
enrich_ccxt_position_metrics_out,
estimate_linear_swap_upnl_usdt,
parse_position_unrealized_pnl,
resolve_position_display_upnl,
)
class TestHubAgentMarkPrice(unittest.TestCase):
def test_binance_mark_price(self):
px = _position_mark_price({"markPrice": 65880.1, "info": {}})
self.assertAlmostEqual(px, 65880.1)
def test_okx_mark_px(self):
px = _position_mark_price({"info": {"markPx": "72.85"}})
self.assertAlmostEqual(px, 72.85)
def test_gate_info_mark(self):
px = _position_mark_price({"info": {"mark_price": "0.2241"}})
self.assertAlmostEqual(px, 0.2241)
def test_missing_returns_none(self):
self.assertIsNone(_position_mark_price({"info": {}}))
def test_infer_from_notional_and_contracts(self):
p = {"notional": 1000, "contracts": 10, "info": {}}
px = _position_mark_price(p)
self.assertAlmostEqual(px, 100.0)
def test_ticker_fallback(self):
class _Ex:
def fetch_ticker(self, sym):
return {"mark": 99.5, "info": {}}
self.assertAlmostEqual(_ticker_mark_price(_Ex(), "BTC/USDT:USDT"), 99.5)
def test_gate_unrealised_pnl_in_info(self):
pnl = parse_position_unrealized_pnl(
{"info": {"unrealised_pnl": "6.81"}, "unrealizedPnl": None}
)
self.assertAlmostEqual(pnl, 6.81)
def test_okx_upl_signed(self):
pnl = parse_position_unrealized_pnl(
{"info": {"upl": "-2.15"}, "unrealizedPnl": None}
)
self.assertAlmostEqual(pnl, -2.15)
def test_enrich_aligns_short_gate_metrics(self):
pos = {
"side": "short",
"contracts": 11,
"entryPrice": 73.187,
"markPrice": 66.038,
"info": {"unrealised_pnl": "7.86"},
}
out = {"unrealized_pnl": 7.86, "mark_price": 66.038}
enrich_ccxt_position_metrics_out(pos, out, contract_size=1.0, funds_decimals=2)
self.assertGreater(out["unrealized_pnl"], 70.0)
def test_estimate_short_hype_contract_size(self):
upnl = estimate_linear_swap_upnl_usdt(
"short", 73.187, 66.038, 11, 0.1
)
self.assertAlmostEqual(upnl, 7.86, places=1)
def test_resolve_prefers_computed_when_exchange_off(self):
shown = resolve_position_display_upnl(
"short", 73.187, 66.038, 11, 1.0, 7.86
)
self.assertAlmostEqual(shown, 78.64, places=1)
def test_resolve_keeps_exchange_when_aligned(self):
shown = resolve_position_display_upnl(
"short", 73.187, 66.038, 11, 0.1, 7.86
)
self.assertAlmostEqual(shown, 7.86, places=2)
if __name__ == "__main__":
unittest.main()
"""子代理持仓:三所标记价字段统一解析."""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "manual_trading_hub"))
from agent import _position_mark_price, _ticker_mark_price # noqa: E402
sys.path.insert(0, str(ROOT))
from lib.hub.hub_position_metrics import ( # noqa: E402
enrich_ccxt_position_metrics_out,
estimate_linear_swap_upnl_usdt,
parse_position_unrealized_pnl,
resolve_position_display_upnl,
)
class TestHubAgentMarkPrice(unittest.TestCase):
def test_binance_mark_price(self):
px = _position_mark_price({"markPrice": 65880.1, "info": {}})
self.assertAlmostEqual(px, 65880.1)
def test_okx_mark_px(self):
px = _position_mark_price({"info": {"markPx": "72.85"}})
self.assertAlmostEqual(px, 72.85)
def test_gate_info_mark(self):
px = _position_mark_price({"info": {"mark_price": "0.2241"}})
self.assertAlmostEqual(px, 0.2241)
def test_missing_returns_none(self):
self.assertIsNone(_position_mark_price({"info": {}}))
def test_infer_from_notional_and_contracts(self):
p = {"notional": 1000, "contracts": 10, "info": {}}
px = _position_mark_price(p)
self.assertAlmostEqual(px, 100.0)
def test_ticker_fallback(self):
class _Ex:
def fetch_ticker(self, sym):
return {"mark": 99.5, "info": {}}
self.assertAlmostEqual(_ticker_mark_price(_Ex(), "BTC/USDT:USDT"), 99.5)
def test_gate_unrealised_pnl_in_info(self):
pnl = parse_position_unrealized_pnl(
{"info": {"unrealised_pnl": "6.81"}, "unrealizedPnl": None}
)
self.assertAlmostEqual(pnl, 6.81)
def test_okx_upl_signed(self):
pnl = parse_position_unrealized_pnl(
{"info": {"upl": "-2.15"}, "unrealizedPnl": None}
)
self.assertAlmostEqual(pnl, -2.15)
def test_enrich_aligns_short_gate_metrics(self):
pos = {
"side": "short",
"contracts": 11,
"entryPrice": 73.187,
"markPrice": 66.038,
"info": {"unrealised_pnl": "7.86"},
}
out = {"unrealized_pnl": 7.86, "mark_price": 66.038}
enrich_ccxt_position_metrics_out(pos, out, contract_size=1.0, funds_decimals=2)
self.assertGreater(out["unrealized_pnl"], 70.0)
def test_estimate_short_hype_contract_size(self):
upnl = estimate_linear_swap_upnl_usdt(
"short", 73.187, 66.038, 11, 0.1
)
self.assertAlmostEqual(upnl, 7.86, places=1)
def test_resolve_prefers_computed_when_exchange_off(self):
shown = resolve_position_display_upnl(
"short", 73.187, 66.038, 11, 1.0, 7.86
)
self.assertAlmostEqual(shown, 78.64, places=1)
def test_resolve_keeps_exchange_when_aligned(self):
shown = resolve_position_display_upnl(
"short", 73.187, 66.038, 11, 0.1, 7.86
)
self.assertAlmostEqual(shown, 7.86, places=2)
if __name__ == "__main__":
unittest.main()
+1 -1
View File
@@ -1,4 +1,4 @@
"""hub_backup_lib 单元测试"""
"""hub_backup_lib 单元测试."""
from __future__ import annotations
import json
+1 -1
View File
@@ -1,4 +1,4 @@
"""后台 board 缓存版本递增与快照"""
"""后台 board 缓存:版本递增与快照."""
from __future__ import annotations
import asyncio
+162 -162
View File
@@ -1,162 +1,162 @@
"""hub_calculator_lib 测算逻辑"""
import unittest
from unittest.mock import patch
from lib.hub.hub_calculator_lib import (
calc_initial_roll_qty,
calc_roll_calculator,
calc_trend_calculator,
solve_add_amount_for_total_risk,
)
MOCK_MARKET = {
"exchange_id": "0",
"exchange_key": "binance",
"exchange_name": "币安 · crypto_monitor_binance",
"exchange_label": "币安 · crypto_monitor_binance",
"base": "ETH",
"exchange_symbol": "ETH/USDT:USDT",
"display_symbol": "ETH/USDT",
"contract_size": 1.0,
"price_tick": 0.01,
"price_decimals": 2,
"amount_decimals": 3,
"min_amount": 0.001,
}
def _mock_resolve(_exchange="binance", _base="ETH"):
return MOCK_MARKET, lambda amount: round(float(amount), 3), None
class HubCalculatorLibTests(unittest.TestCase):
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_trend_calculator_long_basic(self, _mock):
data, err = calc_trend_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
leverage=5,
entry_price=100,
stop_loss=95,
add_upper=110,
take_profit=120,
dca_legs=3,
exchange_id="0",
base="ETH",
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["risk_budget_u"], 50.0)
self.assertGreaterEqual(len(data["rows"]), 2)
self.assertEqual(data["rows"][0]["label"], "首仓")
self.assertEqual(data["market"]["display_symbol"], "ETH/USDT")
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_trend_calculator_short_rejects_bad_bounds(self, _mock):
data, err = calc_trend_calculator(
direction="short",
capital_usdt=1000,
risk_percent=5,
leverage=5,
entry_price=100,
stop_loss=90,
add_upper=110,
take_profit=80,
dca_legs=3,
)
self.assertIsNone(data)
self.assertIsNotNone(err)
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_first_leg_auto(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[],
legs_done=0,
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["first_contracts"], 10.0)
self.assertEqual(len(data["rows"]), 1)
self.assertEqual(data["rows"][0]["loss_at_sl_u"], 50.0)
self.assertEqual(data["rows"][0]["profit_at_tp_u"], 200.0)
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_chain_two_legs(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[
{"add_price": 105, "new_stop_loss": 98},
{"add_price": 108, "new_stop_loss": 101},
],
legs_done=0,
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(len(data["rows"]), 3)
self.assertEqual(data["rows"][1]["label"], "滚仓1")
self.assertGreater(float(data["final_contracts"]), float(data["first_contracts"]))
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_rejects_too_many_legs(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[
{"add_price": 105, "new_stop_loss": 98},
{"add_price": 108, "new_stop_loss": 101},
{"add_price": 110, "new_stop_loss": 103},
{"add_price": 112, "new_stop_loss": 105},
],
legs_done=0,
)
self.assertIsNone(data)
self.assertIsNotNone(err)
def test_initial_roll_qty(self):
qty, err = calc_initial_roll_qty("long", 100, 95, 50, 1.0)
self.assertIsNone(err)
self.assertEqual(qty, 10.0)
def test_initial_roll_qty_with_contract_size(self):
qty, err = calc_initial_roll_qty("long", 100, 95, 50, 0.1)
self.assertIsNone(err)
self.assertEqual(qty, 100.0)
def test_solve_add_with_contract_size(self):
q2, err = solve_add_amount_for_total_risk(
"long",
qty_existing=10.0,
entry_existing=100.0,
add_price=105.0,
new_stop=98.0,
risk_budget_usdt=50.0,
contract_size=1.0,
)
self.assertIsNone(err)
self.assertIsNotNone(q2)
assert q2 is not None
self.assertGreater(q2, 0)
if __name__ == "__main__":
unittest.main()
"""hub_calculator_lib 测算逻辑."""
import unittest
from unittest.mock import patch
from lib.hub.hub_calculator_lib import (
calc_initial_roll_qty,
calc_roll_calculator,
calc_trend_calculator,
solve_add_amount_for_total_risk,
)
MOCK_MARKET = {
"exchange_id": "0",
"exchange_key": "binance",
"exchange_name": "币安 · crypto_monitor_binance",
"exchange_label": "币安 · crypto_monitor_binance",
"base": "ETH",
"exchange_symbol": "ETH/USDT:USDT",
"display_symbol": "ETH/USDT",
"contract_size": 1.0,
"price_tick": 0.01,
"price_decimals": 2,
"amount_decimals": 3,
"min_amount": 0.001,
}
def _mock_resolve(_exchange="binance", _base="ETH"):
return MOCK_MARKET, lambda amount: round(float(amount), 3), None
class HubCalculatorLibTests(unittest.TestCase):
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_trend_calculator_long_basic(self, _mock):
data, err = calc_trend_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
leverage=5,
entry_price=100,
stop_loss=95,
add_upper=110,
take_profit=120,
dca_legs=3,
exchange_id="0",
base="ETH",
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["risk_budget_u"], 50.0)
self.assertGreaterEqual(len(data["rows"]), 2)
self.assertEqual(data["rows"][0]["label"], "首仓")
self.assertEqual(data["market"]["display_symbol"], "ETH/USDT")
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_trend_calculator_short_rejects_bad_bounds(self, _mock):
data, err = calc_trend_calculator(
direction="short",
capital_usdt=1000,
risk_percent=5,
leverage=5,
entry_price=100,
stop_loss=90,
add_upper=110,
take_profit=80,
dca_legs=3,
)
self.assertIsNone(data)
self.assertIsNotNone(err)
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_first_leg_auto(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[],
legs_done=0,
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["first_contracts"], 10.0)
self.assertEqual(len(data["rows"]), 1)
self.assertEqual(data["rows"][0]["loss_at_sl_u"], 50.0)
self.assertEqual(data["rows"][0]["profit_at_tp_u"], 200.0)
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_chain_two_legs(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[
{"add_price": 105, "new_stop_loss": 98},
{"add_price": 108, "new_stop_loss": 101},
],
legs_done=0,
)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(len(data["rows"]), 3)
self.assertEqual(data["rows"][1]["label"], "滚仓1")
self.assertGreater(float(data["final_contracts"]), float(data["first_contracts"]))
@patch("lib.hub.hub_calculator_lib._resolve_market", return_value=_mock_resolve())
def test_roll_calculator_rejects_too_many_legs(self, _mock):
data, err = calc_roll_calculator(
direction="long",
capital_usdt=1000,
risk_percent=5,
entry_price=100,
stop_loss=95,
take_profit=120,
add_legs=[
{"add_price": 105, "new_stop_loss": 98},
{"add_price": 108, "new_stop_loss": 101},
{"add_price": 110, "new_stop_loss": 103},
{"add_price": 112, "new_stop_loss": 105},
],
legs_done=0,
)
self.assertIsNone(data)
self.assertIsNotNone(err)
def test_initial_roll_qty(self):
qty, err = calc_initial_roll_qty("long", 100, 95, 50, 1.0)
self.assertIsNone(err)
self.assertEqual(qty, 10.0)
def test_initial_roll_qty_with_contract_size(self):
qty, err = calc_initial_roll_qty("long", 100, 95, 50, 0.1)
self.assertIsNone(err)
self.assertEqual(qty, 100.0)
def test_solve_add_with_contract_size(self):
q2, err = solve_add_amount_for_total_risk(
"long",
qty_existing=10.0,
entry_existing=100.0,
add_price=105.0,
new_stop=98.0,
risk_budget_usdt=50.0,
contract_size=1.0,
)
self.assertIsNone(err)
self.assertIsNotNone(q2)
assert q2 is not None
self.assertGreater(q2, 0)
if __name__ == "__main__":
unittest.main()
+113 -113
View File
@@ -1,113 +1,113 @@
"""hub_calculator_market_lib 合约解析"""
import unittest
from unittest.mock import patch
from lib.hub.hub_calculator_market_lib import (
amount_decimals_from_exchange,
find_exchange,
get_calculator_market,
list_calculator_exchanges,
make_amount_precise_fn_from_market,
normalize_base_symbol,
resolve_usdt_perp_symbol,
)
class FakeExchange:
def __init__(self, markets: dict):
self.markets = markets
def market(self, symbol: str):
return self.markets[symbol]
def amount_to_precision(self, symbol: str, amount: float) -> str:
return f"{float(amount):.3f}"
class HubCalculatorMarketLibTests(unittest.TestCase):
def test_normalize_base_symbol(self):
self.assertEqual(normalize_base_symbol("eth"), "ETH")
self.assertEqual(normalize_base_symbol("ETH/USDT:USDT"), "ETH")
self.assertEqual(normalize_base_symbol("ETHUSDT"), "ETH")
def test_resolve_usdt_perp_symbol(self):
ex = FakeExchange(
{
"ETH/USDT:USDT": {
"base": "ETH",
"quote": "USDT",
"swap": True,
"active": True,
"contractSize": 1.0,
"limits": {"amount": {"min": 0.001}},
"precision": {"price": 2, "amount": 3},
}
}
)
sym, err = resolve_usdt_perp_symbol(ex, "ETH")
self.assertIsNone(err)
self.assertEqual(sym, "ETH/USDT:USDT")
def test_amount_decimals_from_exchange(self):
ex = FakeExchange({})
self.assertEqual(amount_decimals_from_exchange(ex, "ETH/USDT:USDT"), 3)
def test_make_amount_precise_fn_from_market(self):
fn = make_amount_precise_fn_from_market({"amount_decimals": 3, "min_amount": 0.001})
self.assertEqual(fn(1.23456), 1.234)
self.assertIsNone(fn(0.0001))
@patch.dict("os.environ", {"HUB_BRIDGE_TOKEN": "test-token"}, clear=False)
def test_hub_headers_use_x_hub_token(self):
from lib.hub.hub_calculator_market_lib import _hub_headers
self.assertEqual(_hub_headers(), {"X-Hub-Token": "test-token"})
@patch("lib.hub.hub_calculator_market_lib.fetch_instance_market_sync")
def test_get_calculator_market_from_instance(self, fetch_mock):
fetch_mock.return_value = {
"ok": True,
"base": "ETH",
"exchange_symbol": "ETH/USDT:USDT",
"display_symbol": "ETH/USDT",
"contract_size": 0.01,
"price_tick": 0.01,
"price_decimals": 2,
"amount_decimals": 2,
"min_amount": 0.01,
}
ex = {
"id": "0",
"key": "binance",
"name": "币安 · crypto_monitor_binance",
"enabled": True,
"flask_url": "http://127.0.0.1:5001",
}
data, err = get_calculator_market("0", "ETH", ex=ex)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["exchange_id"], "0")
self.assertEqual(data["exchange_name"], "币安 · crypto_monitor_binance")
self.assertEqual(data["contract_size"], 0.01)
@patch("lib.hub.hub_calculator_market_lib.enabled_exchanges")
def test_list_calculator_exchanges(self, enabled_mock):
enabled_mock.return_value = [
{"id": "0", "key": "binance", "name": "币安", "enabled": True},
]
rows = list_calculator_exchanges()
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["id"], "0")
def test_find_exchange_by_id(self):
with patch(
"lib.hub.hub_calculator_market_lib.load_settings",
return_value={"exchanges": [{"id": "2", "key": "gate", "name": "Gate"}]},
):
self.assertEqual(find_exchange("2")["key"], "gate")
if __name__ == "__main__":
unittest.main()
"""hub_calculator_market_lib 合约解析."""
import unittest
from unittest.mock import patch
from lib.hub.hub_calculator_market_lib import (
amount_decimals_from_exchange,
find_exchange,
get_calculator_market,
list_calculator_exchanges,
make_amount_precise_fn_from_market,
normalize_base_symbol,
resolve_usdt_perp_symbol,
)
class FakeExchange:
def __init__(self, markets: dict):
self.markets = markets
def market(self, symbol: str):
return self.markets[symbol]
def amount_to_precision(self, symbol: str, amount: float) -> str:
return f"{float(amount):.3f}"
class HubCalculatorMarketLibTests(unittest.TestCase):
def test_normalize_base_symbol(self):
self.assertEqual(normalize_base_symbol("eth"), "ETH")
self.assertEqual(normalize_base_symbol("ETH/USDT:USDT"), "ETH")
self.assertEqual(normalize_base_symbol("ETHUSDT"), "ETH")
def test_resolve_usdt_perp_symbol(self):
ex = FakeExchange(
{
"ETH/USDT:USDT": {
"base": "ETH",
"quote": "USDT",
"swap": True,
"active": True,
"contractSize": 1.0,
"limits": {"amount": {"min": 0.001}},
"precision": {"price": 2, "amount": 3},
}
}
)
sym, err = resolve_usdt_perp_symbol(ex, "ETH")
self.assertIsNone(err)
self.assertEqual(sym, "ETH/USDT:USDT")
def test_amount_decimals_from_exchange(self):
ex = FakeExchange({})
self.assertEqual(amount_decimals_from_exchange(ex, "ETH/USDT:USDT"), 3)
def test_make_amount_precise_fn_from_market(self):
fn = make_amount_precise_fn_from_market({"amount_decimals": 3, "min_amount": 0.001})
self.assertEqual(fn(1.23456), 1.234)
self.assertIsNone(fn(0.0001))
@patch.dict("os.environ", {"HUB_BRIDGE_TOKEN": "test-token"}, clear=False)
def test_hub_headers_use_x_hub_token(self):
from lib.hub.hub_calculator_market_lib import _hub_headers
self.assertEqual(_hub_headers(), {"X-Hub-Token": "test-token"})
@patch("lib.hub.hub_calculator_market_lib.fetch_instance_market_sync")
def test_get_calculator_market_from_instance(self, fetch_mock):
fetch_mock.return_value = {
"ok": True,
"base": "ETH",
"exchange_symbol": "ETH/USDT:USDT",
"display_symbol": "ETH/USDT",
"contract_size": 0.01,
"price_tick": 0.01,
"price_decimals": 2,
"amount_decimals": 2,
"min_amount": 0.01,
}
ex = {
"id": "0",
"key": "binance",
"name": "币安 · crypto_monitor_binance",
"enabled": True,
"flask_url": "http://127.0.0.1:5001",
}
data, err = get_calculator_market("0", "ETH", ex=ex)
self.assertIsNone(err)
self.assertIsNotNone(data)
assert data is not None
self.assertEqual(data["exchange_id"], "0")
self.assertEqual(data["exchange_name"], "币安 · crypto_monitor_binance")
self.assertEqual(data["contract_size"], 0.01)
@patch("lib.hub.hub_calculator_market_lib.enabled_exchanges")
def test_list_calculator_exchanges(self, enabled_mock):
enabled_mock.return_value = [
{"id": "0", "key": "binance", "name": "币安", "enabled": True},
]
rows = list_calculator_exchanges()
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["id"], "0")
def test_find_exchange_by_id(self):
with patch(
"lib.hub.hub_calculator_market_lib.load_settings",
return_value={"exchanges": [{"id": "2", "key": "gate", "name": "Gate"}]},
):
self.assertEqual(find_exchange("2")["key"], "gate")
if __name__ == "__main__":
unittest.main()
+1 -1
View File
@@ -1,4 +1,4 @@
"""行情区 chart 后台轮询订阅"""
"""行情区 chart 后台轮询订阅."""
from __future__ import annotations
import asyncio
+1 -1
View File
@@ -1,4 +1,4 @@
"""中控条件单列表子代理与 Flask exchange_tpsl 合并去重"""
"""中控条件单列表:子代理与 Flask exchange_tpsl 合并去重."""
from manual_trading_hub.hub import _merge_conditional_orders_no_dup, _merge_flask_exchange_tpsl
+2 -2
View File
@@ -10,7 +10,7 @@ from lib.hub.hub_divergence_scan_lib import (
def _synthetic_bull_div_closes(n: int = 120) -> list[float]:
"""价格双底 + MACD 抬高 → 底背离"""
"""价格双底 + MACD 抬高 → 底背离."""
closes = [100.0] * n
# 下跌
for i in range(20, 40):
@@ -84,7 +84,7 @@ class TestHubDivergenceScanLib(unittest.TestCase):
def test_detect_macd_divergence_may_hit_on_synthetic(self):
closes = _synthetic_bull_div_closes()
hit = detect_latest_macd_divergence(closes)
# 合成数据不保证必中但函数应正常返回
# 合成数据不保证必中,但函数应正常返回
self.assertIn(hit.get("direction"), (None, "bull", "bear"))
def test_analyze_ohlcv_bars_from_rows(self):
+157 -157
View File
@@ -1,157 +1,157 @@
"""开仓计划库CRUD 与胜率统计"""
from __future__ import annotations
import tempfile
from pathlib import Path
from lib.hub.hub_entry_plan_lib import (
compute_entry_plan_stats,
create_entry_plan,
delete_entry_plan,
init_db,
list_entry_plans,
normalize_plan_symbol,
resolve_stats_date_bounds,
update_entry_plan,
)
def _base_payload(**overrides):
data = {
"plan_date": "2026-06-14",
"exchange_key": "binance",
"symbol": "BTC",
"plan_type": "trend",
"trend_timeframe": "4h",
"entry_timeframe": "15m",
"direction": "long",
"target_level": "70000",
"current_range": "68000-69000",
"entry_scheme": "breakout",
"note": "test",
}
data.update(overrides)
return data
def test_normalize_plan_symbol():
assert normalize_plan_symbol("btc") == "BTC/USDT"
assert normalize_plan_symbol("ETH/USDT") == "ETH/USDT"
def test_create_without_entry_scheme():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
payload = _base_payload()
del payload["entry_scheme"]
row = create_entry_plan(payload, db_path=db)
assert row["entry_scheme"] == ""
assert row["entry_scheme_label"] == "待填写"
def test_archive_requires_entry_scheme():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
payload = _base_payload()
del payload["entry_scheme"]
row = create_entry_plan(payload, db_path=db)
try:
update_entry_plan(int(row["id"]), {"result": "win"}, db_path=db)
assert False, "expected ValueError"
except ValueError as e:
assert "入场方案" in str(e)
updated = update_entry_plan(
int(row["id"]),
{"entry_scheme": "breakout", "result": "win"},
db_path=db,
)
assert updated["status"] == "archived"
def test_create_list_delete_active_plan():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(), db_path=db)
assert row["status"] == "active"
assert row["symbol"] == "BTC/USDT"
active = list_entry_plans(status="active", db_path=db)
assert len(active) == 1
assert delete_entry_plan(int(row["id"]), db_path=db) is True
assert list_entry_plans(status="active", db_path=db) == []
def test_archive_on_result():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(symbol="SOL"), db_path=db)
updated = update_entry_plan(
int(row["id"]),
{"result": "win", "pnl_amount": 12.5},
db_path=db,
)
assert updated["status"] == "archived"
assert updated["result"] == "win"
assert updated["pnl_amount"] == 12.5
assert list_entry_plans(status="active", db_path=db) == []
archived = list_entry_plans(status="archived", db_path=db)
assert len(archived) == 1
def test_archive_without_pnl_amount():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(symbol="DOGE"), db_path=db)
updated = update_entry_plan(int(row["id"]), {"result": "loss"}, db_path=db)
assert updated["status"] == "archived"
assert updated["pnl_amount"] is None
def test_cannot_delete_archived():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(), db_path=db)
update_entry_plan(int(row["id"]), {"result": "win"}, db_path=db)
try:
delete_entry_plan(int(row["id"]), db_path=db)
assert False, "expected ValueError"
except ValueError as e:
assert "仅进行中" in str(e)
def test_compute_stats_by_symbol():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
for sym, res in (("BTC", "win"), ("BTC", "loss"), ("ETH", "win")):
row = create_entry_plan(_base_payload(symbol=sym), db_path=db)
update_entry_plan(int(row["id"]), {"result": res}, db_path=db)
stats = compute_entry_plan_stats(dimension="symbol", period="all", db_path=db)
by_sym = {it["key"]: it for it in stats["items"]}
assert by_sym["BTC/USDT"]["win_count"] == 1
assert by_sym["BTC/USDT"]["loss_count"] == 1
assert by_sym["BTC/USDT"]["win_rate"] == 50.0
assert by_sym["ETH/USDT"]["win_count"] == 1
def test_stats_period_range_filter():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row1 = create_entry_plan(_base_payload(plan_date="2026-06-01"), db_path=db)
row2 = create_entry_plan(_base_payload(plan_date="2026-06-20", symbol="ETH"), db_path=db)
update_entry_plan(int(row1["id"]), {"result": "win"}, db_path=db)
update_entry_plan(int(row2["id"]), {"result": "loss"}, db_path=db)
stats = compute_entry_plan_stats(
dimension="symbol",
period="range",
date_from="2026-06-01",
date_to="2026-06-10",
db_path=db,
)
assert len(stats["items"]) == 1
assert stats["items"][0]["key"] == "BTC/USDT"
def test_resolve_stats_date_bounds():
df, dt, label = resolve_stats_date_bounds(period="all")
assert df is None and dt is None
assert "全部" in label
"""开仓计划库:CRUD 与胜率统计."""
from __future__ import annotations
import tempfile
from pathlib import Path
from lib.hub.hub_entry_plan_lib import (
compute_entry_plan_stats,
create_entry_plan,
delete_entry_plan,
init_db,
list_entry_plans,
normalize_plan_symbol,
resolve_stats_date_bounds,
update_entry_plan,
)
def _base_payload(**overrides):
data = {
"plan_date": "2026-06-14",
"exchange_key": "binance",
"symbol": "BTC",
"plan_type": "trend",
"trend_timeframe": "4h",
"entry_timeframe": "15m",
"direction": "long",
"target_level": "70000",
"current_range": "68000-69000",
"entry_scheme": "breakout",
"note": "test",
}
data.update(overrides)
return data
def test_normalize_plan_symbol():
assert normalize_plan_symbol("btc") == "BTC/USDT"
assert normalize_plan_symbol("ETH/USDT") == "ETH/USDT"
def test_create_without_entry_scheme():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
payload = _base_payload()
del payload["entry_scheme"]
row = create_entry_plan(payload, db_path=db)
assert row["entry_scheme"] == ""
assert row["entry_scheme_label"] == "待填写"
def test_archive_requires_entry_scheme():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
payload = _base_payload()
del payload["entry_scheme"]
row = create_entry_plan(payload, db_path=db)
try:
update_entry_plan(int(row["id"]), {"result": "win"}, db_path=db)
assert False, "expected ValueError"
except ValueError as e:
assert "入场方案" in str(e)
updated = update_entry_plan(
int(row["id"]),
{"entry_scheme": "breakout", "result": "win"},
db_path=db,
)
assert updated["status"] == "archived"
def test_create_list_delete_active_plan():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(), db_path=db)
assert row["status"] == "active"
assert row["symbol"] == "BTC/USDT"
active = list_entry_plans(status="active", db_path=db)
assert len(active) == 1
assert delete_entry_plan(int(row["id"]), db_path=db) is True
assert list_entry_plans(status="active", db_path=db) == []
def test_archive_on_result():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(symbol="SOL"), db_path=db)
updated = update_entry_plan(
int(row["id"]),
{"result": "win", "pnl_amount": 12.5},
db_path=db,
)
assert updated["status"] == "archived"
assert updated["result"] == "win"
assert updated["pnl_amount"] == 12.5
assert list_entry_plans(status="active", db_path=db) == []
archived = list_entry_plans(status="archived", db_path=db)
assert len(archived) == 1
def test_archive_without_pnl_amount():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(symbol="DOGE"), db_path=db)
updated = update_entry_plan(int(row["id"]), {"result": "loss"}, db_path=db)
assert updated["status"] == "archived"
assert updated["pnl_amount"] is None
def test_cannot_delete_archived():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row = create_entry_plan(_base_payload(), db_path=db)
update_entry_plan(int(row["id"]), {"result": "win"}, db_path=db)
try:
delete_entry_plan(int(row["id"]), db_path=db)
assert False, "expected ValueError"
except ValueError as e:
assert "仅进行中" in str(e)
def test_compute_stats_by_symbol():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
for sym, res in (("BTC", "win"), ("BTC", "loss"), ("ETH", "win")):
row = create_entry_plan(_base_payload(symbol=sym), db_path=db)
update_entry_plan(int(row["id"]), {"result": res}, db_path=db)
stats = compute_entry_plan_stats(dimension="symbol", period="all", db_path=db)
by_sym = {it["key"]: it for it in stats["items"]}
assert by_sym["BTC/USDT"]["win_count"] == 1
assert by_sym["BTC/USDT"]["loss_count"] == 1
assert by_sym["BTC/USDT"]["win_rate"] == 50.0
assert by_sym["ETH/USDT"]["win_count"] == 1
def test_stats_period_range_filter():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "plans.db"
row1 = create_entry_plan(_base_payload(plan_date="2026-06-01"), db_path=db)
row2 = create_entry_plan(_base_payload(plan_date="2026-06-20", symbol="ETH"), db_path=db)
update_entry_plan(int(row1["id"]), {"result": "win"}, db_path=db)
update_entry_plan(int(row2["id"]), {"result": "loss"}, db_path=db)
stats = compute_entry_plan_stats(
dimension="symbol",
period="range",
date_from="2026-06-01",
date_to="2026-06-10",
db_path=db,
)
assert len(stats["items"]) == 1
assert stats["items"][0]["key"] == "BTC/USDT"
def test_resolve_stats_date_bounds():
df, dt, label = resolve_stats_date_bounds(period="all")
assert df is None and dt is None
assert "全部" in label
+1 -1
View File
@@ -1,4 +1,4 @@
"""OKX 中控委托须为 OCO 条件单不得带 reduceOnly 或分两笔 market"""
"""OKX 中控委托:须为 OCO 条件单,不得带 reduceOnly 或分两笔 market."""
from __future__ import annotations
import sys
+113 -113
View File
@@ -1,113 +1,113 @@
"""hub_fund_history_lib总资金回撤与日快照"""
from __future__ import annotations
from lib.hub.hub_fund_history_lib import (
account_total_usdt,
build_fund_overview,
compute_drawdown,
get_fund_history,
record_fund_snapshot,
)
def test_account_total_requires_both_sides():
assert account_total_usdt(10, 20) == 30.0
assert account_total_usdt(10, None) is None
assert account_total_usdt(None, 5) is None
def test_compute_drawdown():
dd = compute_drawdown([100, 120, 90, 110])
assert dd["peak_usdt"] == 120.0
assert dd["max_drawdown_u"] == 30.0
assert dd["max_drawdown_pct"] == 25.0
def test_build_fund_overview_skips_unmonitored(tmp_path, monkeypatch):
hist_path = tmp_path / "hub_fund_history.json"
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_PATH", hist_path)
record_fund_snapshot(
"2026-06-01",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 10,
"trading_usdt": 20,
"monitored": True,
}
],
keep_days=180,
)
record_fund_snapshot(
"2026-06-02",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 12,
"trading_usdt": 18,
"monitored": True,
}
],
keep_days=180,
)
exchanges = [
{"id": "0", "key": "binance", "name": "Binance", "enabled": True},
{"id": "2", "key": "gate", "name": "Gate", "enabled": False},
]
board_rows = [
{
"key": "binance",
"name": "Binance",
"account_ok": True,
"funding_usdt": 15,
"trading_usdt": 25,
}
]
out = build_fund_overview(
exchanges,
board_rows=board_rows,
trading_day="2026-06-02",
keep_days=180,
)
assert out["totals"]["total_usdt"] == 40.0
assert out["totals"]["monitored_count"] == 1
assert len(out["accounts"]) == 1
assert all(a["monitored"] for a in out["accounts"])
assert out["totals"]["drawdown"]["max_drawdown_u"] == 0.0
def test_history_start_day_filters_older(tmp_path, monkeypatch):
hist_path = tmp_path / "hub_fund_history.json"
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_PATH", hist_path)
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_START_DAY", "2026-06-09")
record_fund_snapshot(
"2026-06-01",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 1,
"trading_usdt": 1,
"monitored": True,
}
],
keep_days=180,
)
record_fund_snapshot(
"2026-06-09",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 10,
"trading_usdt": 20,
"monitored": True,
}
],
keep_days=180,
)
hist = get_fund_history(anchor_day="2026-06-10", keep_days=180)
assert "2026-06-01" not in hist
assert "2026-06-09" in hist
"""hub_fund_history_lib:总资金,回撤与日快照."""
from __future__ import annotations
from lib.hub.hub_fund_history_lib import (
account_total_usdt,
build_fund_overview,
compute_drawdown,
get_fund_history,
record_fund_snapshot,
)
def test_account_total_requires_both_sides():
assert account_total_usdt(10, 20) == 30.0
assert account_total_usdt(10, None) is None
assert account_total_usdt(None, 5) is None
def test_compute_drawdown():
dd = compute_drawdown([100, 120, 90, 110])
assert dd["peak_usdt"] == 120.0
assert dd["max_drawdown_u"] == 30.0
assert dd["max_drawdown_pct"] == 25.0
def test_build_fund_overview_skips_unmonitored(tmp_path, monkeypatch):
hist_path = tmp_path / "hub_fund_history.json"
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_PATH", hist_path)
record_fund_snapshot(
"2026-06-01",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 10,
"trading_usdt": 20,
"monitored": True,
}
],
keep_days=180,
)
record_fund_snapshot(
"2026-06-02",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 12,
"trading_usdt": 18,
"monitored": True,
}
],
keep_days=180,
)
exchanges = [
{"id": "0", "key": "binance", "name": "Binance", "enabled": True},
{"id": "2", "key": "gate", "name": "Gate", "enabled": False},
]
board_rows = [
{
"key": "binance",
"name": "Binance",
"account_ok": True,
"funding_usdt": 15,
"trading_usdt": 25,
}
]
out = build_fund_overview(
exchanges,
board_rows=board_rows,
trading_day="2026-06-02",
keep_days=180,
)
assert out["totals"]["total_usdt"] == 40.0
assert out["totals"]["monitored_count"] == 1
assert len(out["accounts"]) == 1
assert all(a["monitored"] for a in out["accounts"])
assert out["totals"]["drawdown"]["max_drawdown_u"] == 0.0
def test_history_start_day_filters_older(tmp_path, monkeypatch):
hist_path = tmp_path / "hub_fund_history.json"
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_PATH", hist_path)
monkeypatch.setattr("hub_fund_history_lib.FUND_HISTORY_START_DAY", "2026-06-09")
record_fund_snapshot(
"2026-06-01",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 1,
"trading_usdt": 1,
"monitored": True,
}
],
keep_days=180,
)
record_fund_snapshot(
"2026-06-09",
[
{
"key": "binance",
"name": "Binance",
"funding_usdt": 10,
"trading_usdt": 20,
"monitored": True,
}
],
keep_days=180,
)
hist = get_fund_history(anchor_day="2026-06-10", keep_days=180)
assert "2026-06-01" not in hist
assert "2026-06-09" in hist
+58 -58
View File
@@ -1,58 +1,58 @@
"""hub_host_status_lib 单元测试"""
from __future__ import annotations
import sys
import unittest
from unittest.mock import MagicMock, patch
from lib.hub.hub_host_status_lib import _disk_path, _state, get_host_status
class HubHostStatusLibTest(unittest.TestCase):
def setUp(self):
_state["primed"] = False
_state["net_ts"] = 0.0
_state["net_sent"] = 0
_state["net_recv"] = 0
def test_disk_path_env_override(self):
with patch.dict("os.environ", {"HUB_HOST_DISK_PATH": "/data"}, clear=False):
self.assertEqual(_disk_path(), "/data")
def test_get_host_status_without_psutil(self):
import builtins
real_import = builtins.__import__
def fake_import(name, globals=None, locals=None, fromlist=(), level=0):
if name == "psutil":
raise ImportError("no psutil")
return real_import(name, globals, locals, fromlist, level)
with patch("builtins.__import__", side_effect=fake_import):
out = get_host_status()
self.assertFalse(out.get("ok"))
self.assertIn("psutil", out.get("msg", ""))
def test_get_host_status_payload(self):
fake_vm = MagicMock(total=8_000_000_000, used=3_200_000_000, percent=40.0)
fake_du = MagicMock(total=100_000_000_000, used=50_000_000_000)
fake_net = MagicMock(bytes_sent=1_000_000, bytes_recv=2_000_000)
fake_psutil = MagicMock()
fake_psutil.cpu_percent.return_value = 12.5
fake_psutil.cpu_count.return_value = 4
fake_psutil.virtual_memory.return_value = fake_vm
fake_psutil.disk_usage.return_value = fake_du
fake_psutil.net_io_counters.return_value = fake_net
fake_psutil.boot_time.return_value = 1_700_000_000.0
with patch.dict(sys.modules, {"psutil": fake_psutil}):
out = get_host_status()
self.assertTrue(out.get("ok"))
self.assertEqual(out["cpu"]["percent"], 12.5)
self.assertEqual(out["memory"]["percent"], 40.0)
self.assertEqual(out["disk"]["percent"], 50.0)
self.assertIn("network", out)
if __name__ == "__main__":
unittest.main()
"""hub_host_status_lib 单元测试."""
from __future__ import annotations
import sys
import unittest
from unittest.mock import MagicMock, patch
from lib.hub.hub_host_status_lib import _disk_path, _state, get_host_status
class HubHostStatusLibTest(unittest.TestCase):
def setUp(self):
_state["primed"] = False
_state["net_ts"] = 0.0
_state["net_sent"] = 0
_state["net_recv"] = 0
def test_disk_path_env_override(self):
with patch.dict("os.environ", {"HUB_HOST_DISK_PATH": "/data"}, clear=False):
self.assertEqual(_disk_path(), "/data")
def test_get_host_status_without_psutil(self):
import builtins
real_import = builtins.__import__
def fake_import(name, globals=None, locals=None, fromlist=(), level=0):
if name == "psutil":
raise ImportError("no psutil")
return real_import(name, globals, locals, fromlist, level)
with patch("builtins.__import__", side_effect=fake_import):
out = get_host_status()
self.assertFalse(out.get("ok"))
self.assertIn("psutil", out.get("msg", ""))
def test_get_host_status_payload(self):
fake_vm = MagicMock(total=8_000_000_000, used=3_200_000_000, percent=40.0)
fake_du = MagicMock(total=100_000_000_000, used=50_000_000_000)
fake_net = MagicMock(bytes_sent=1_000_000, bytes_recv=2_000_000)
fake_psutil = MagicMock()
fake_psutil.cpu_percent.return_value = 12.5
fake_psutil.cpu_count.return_value = 4
fake_psutil.virtual_memory.return_value = fake_vm
fake_psutil.disk_usage.return_value = fake_du
fake_psutil.net_io_counters.return_value = fake_net
fake_psutil.boot_time.return_value = 1_700_000_000.0
with patch.dict(sys.modules, {"psutil": fake_psutil}):
out = get_host_status()
self.assertTrue(out.get("ok"))
self.assertEqual(out["cpu"]["percent"], 12.5)
self.assertEqual(out["memory"]["percent"], 40.0)
self.assertEqual(out["disk"]["percent"], 50.0)
self.assertIn("network", out)
if __name__ == "__main__":
unittest.main()
+466 -466
View File
@@ -1,466 +1,466 @@
"""中控 K 线库分周期保留聚合与分页读取"""
from __future__ import annotations
import tempfile
import time
import unittest
from pathlib import Path
from lib.hub.hub_kline_store import (
HUB_KLINE_REMOTE_FETCH_CAP,
_since_ms_for_span,
clear_series_bars,
init_db,
load_bars_before,
load_bars_latest,
purge_retention,
purge_timeframe_by_days,
resolve_chart_bars,
retention_days,
trim_contiguous_tail,
upsert_bars,
)
from lib.hub.hub_ohlcv_lib import (
TIMEFRAME_MS,
bar_limit_for_timeframe,
chart_fetch_start_ms,
chart_initial_limit,
last_closed_bar_open_ms,
window_start_ms,
)
class TestHubKlineStore(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.db = Path(self.tmp.name) / "test_hub_kline.db"
def tearDown(self):
self.tmp.cleanup()
def test_bar_limits(self):
self.assertEqual(bar_limit_for_timeframe("5m"), 5000)
self.assertEqual(bar_limit_for_timeframe("1h"), 1000)
self.assertEqual(bar_limit_for_timeframe("1d"), 1000)
self.assertEqual(bar_limit_for_timeframe("1w"), 500)
self.assertEqual(chart_initial_limit("5m"), 2000)
self.assertEqual(chart_initial_limit("1h"), 1000)
self.assertEqual(chart_initial_limit("1d"), 500)
def test_chart_fetch_window_exceeds_retention(self):
now = int(time.time() * 1000)
need = bar_limit_for_timeframe("1d")
fetch_start = chart_fetch_start_ms("1d", need, now)
db_start = window_start_ms("1d", need, retention_days(), now)
self.assertLess(fetch_start, db_start)
def test_purge_retention_5m_one_year(self):
init_db(self.db)
old_ms = int(time.time() * 1000) - 400 * 86400000
upsert_bars(
"okx",
"BTC/USDT",
"5m",
[
{
"open_time_ms": old_ms,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 10,
}
],
self.db,
)
n = purge_timeframe_by_days("5m", 365, self.db)
self.assertGreaterEqual(n, 1)
rows = load_bars_latest("okx", "BTC/USDT", "5m", 10, self.db)
self.assertEqual(len(rows), 0)
def test_purge_retention_keeps_1d(self):
init_db(self.db)
old_ms = int(time.time() * 1000) - 400 * 86400000
upsert_bars(
"okx",
"BTC/USDT",
"1d",
[
{
"open_time_ms": old_ms,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 10,
}
],
self.db,
)
purge_retention(self.db)
rows = load_bars_latest("okx", "BTC/USDT", "1d", 10, self.db)
self.assertEqual(len(rows), 1)
def test_resolve_uses_cache_without_remote(self):
init_db(self.db)
now = int(time.time() * 1000)
tf = "5m"
period = TIMEFRAME_MS[tf]
last_closed = last_closed_bar_open_ms(tf, now)
bars = []
for i in range(400):
oms = last_closed - (399 - i) * period
bars.append(
{
"open_time_ms": oms,
"open": 100 + i,
"high": 101 + i,
"low": 99 + i,
"close": 100.5 + i,
"volume": 1000 + i,
}
)
upsert_bars("okx", "ETH/USDT", tf, bars, self.db)
def remote_fetch(**kwargs):
self.fail("不应请求交易所")
out = resolve_chart_bars(
"okx",
"ETH/USDT",
tf,
remote_fetch,
db_path=self.db,
limit=300,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("candles") or []), 300)
def test_resolve_15m_reads_native_bars(self):
init_db(self.db)
now = int(time.time() * 1000)
period = TIMEFRAME_MS["15m"]
last_closed = last_closed_bar_open_ms("15m", now)
bars = []
for i in range(12):
oms = last_closed - (11 - i) * period
bars.append(
{
"open_time_ms": oms,
"open": 1.0 + i,
"high": 2.0 + i,
"low": 0.5 + i,
"close": 1.5 + i,
"volume": 10.0,
}
)
upsert_bars("okx", "ETH/USDT", "15m", bars, self.db)
def remote_fetch(**kwargs):
self.fail("不应请求交易所")
out = resolve_chart_bars(
"okx",
"ETH/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=10,
)
self.assertTrue(out.get("ok"))
self.assertEqual(out.get("source"), "db")
self.assertEqual(out.get("storage_timeframe"), "15m")
self.assertGreaterEqual(len(out.get("candles") or []), 10)
def test_load_bars_before(self):
init_db(self.db)
period = TIMEFRAME_MS["1h"]
base = 1_700_000_000_000
bars = []
for i in range(5):
bars.append(
{
"open_time_ms": base + i * period,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
)
upsert_bars("okx", "BTC/USDT", "1h", bars, self.db)
before = base + 3 * period
got = load_bars_before("okx", "BTC/USDT", "1h", before, 2, self.db)
self.assertEqual(len(got), 2)
self.assertEqual(got[-1]["open_time_ms"], base + 2 * period)
def test_trim_contiguous_tail_drops_orphan_prefix(self):
period = TIMEFRAME_MS["15m"]
base_old = 1_700_000_000_000
base_new = base_old + period * 500
bars = []
for i in range(3):
bars.append(
{
"open_time_ms": base_old + i * period,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
)
for i in range(5):
bars.append(
{
"open_time_ms": base_new + i * period,
"open": 2,
"high": 3,
"low": 1.5,
"close": 2.5,
"volume": 2,
}
)
trimmed, split = trim_contiguous_tail(bars, period)
self.assertEqual(split, 3)
self.assertEqual(len(trimmed), 5)
self.assertEqual(trimmed[0]["open_time_ms"], base_new)
def test_resolve_drops_discontinuous_orphans(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
old_ms = now - period * 800
upsert_bars(
"okx",
"ONDO/USDT",
"15m",
[
{
"open_time_ms": old_ms,
"open": 0.33,
"high": 0.34,
"low": 0.32,
"close": 0.335,
"volume": 100,
}
],
self.db,
)
recent = []
start = now - period * 20
for i in range(20):
recent.append(
{
"open_time_ms": start + i * period,
"open": 0.35,
"high": 0.36,
"low": 0.34,
"close": 0.355,
"volume": 50,
}
)
def remote_fetch(**kwargs):
return {"ok": True, "bars": recent, "price_tick": 0.0001}
out = resolve_chart_bars(
"okx",
"ONDO/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=50,
)
self.assertTrue(out.get("ok"))
candles = out.get("candles") or []
self.assertGreaterEqual(len(candles), 19)
if len(candles) >= 2:
for i in range(1, len(candles)):
gap = candles[i]["time"] - candles[i - 1]["time"]
self.assertLessEqual(gap, int(period / 1000 * 3.0))
def test_resolve_refetches_when_db_has_discontinuous_full_count(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
old_start = now - period * 3000
recent_start = now - period * 25
old_bars = [
{
"open_time_ms": old_start + i * period,
"open": 62000,
"high": 62100,
"low": 61900,
"close": 62050,
"volume": 10,
}
for i in range(500)
]
recent = [
{
"open_time_ms": recent_start + i * period,
"open": 104000,
"high": 104100,
"low": 103900,
"close": 104050,
"volume": 20,
}
for i in range(30)
]
upsert_bars("binance", "BTC/USDT", "15m", old_bars, self.db)
upsert_bars("binance", "BTC/USDT", "15m", recent, self.db)
fetch_calls = []
def remote_fetch(**kwargs):
fetch_calls.append(dict(kwargs))
full = []
start = now - period * 120
for i in range(120):
full.append(
{
"open_time_ms": start + i * period,
"open": 104000 + i,
"high": 104100 + i,
"low": 103900 + i,
"close": 104050 + i,
"volume": 30,
}
)
return {"ok": True, "bars": full, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=2000,
)
self.assertTrue(out.get("ok"))
self.assertGreater(len(fetch_calls), 0)
self.assertGreaterEqual(len(out.get("candles") or []), 100)
self.assertGreater(int(out.get("fetched") or 0), 0)
def test_clear_series_and_force_refetch(self):
init_db(self.db)
period = TIMEFRAME_MS["5m"]
now = int(time.time() * 1000)
stale = [
{
"open_time_ms": now - period * (i + 100),
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
for i in range(40)
]
upsert_bars("binance", "BTC/USDT", "5m", stale, self.db)
self.assertEqual(len(load_bars_latest("binance", "BTC/USDT", "5m", 100, self.db)), 40)
removed = clear_series_bars("binance", "BTC/USDT", "5m", self.db)
self.assertEqual(removed, 40)
self.assertEqual(len(load_bars_latest("binance", "BTC/USDT", "5m", 100, self.db)), 0)
fresh = [
{
"open_time_ms": now - period * (20 - i),
"open": 10,
"high": 11,
"low": 9,
"close": 10.5,
"volume": 2,
}
for i in range(20)
]
def remote_fetch(**kwargs):
return {"ok": True, "bars": fresh, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"5m",
remote_fetch,
db_path=self.db,
force_refresh=True,
clear_db=True,
limit=50,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(int(out.get("cleared") or 0), 0)
self.assertGreater(int(out.get("fetched") or 0), 0)
self.assertGreaterEqual(len(out.get("candles") or []), 19)
def test_since_span_matches_fetch_limit_not_need(self):
period = TIMEFRAME_MS["15m"]
now_ms = 1_800_000_000_000
fetch_limit = HUB_KLINE_REMOTE_FETCH_CAP
since = _since_ms_for_span(
now_ms=now_ms,
period_ms=period,
span_bars=fetch_limit,
cutoff_ms=0,
)
self.assertEqual(since, now_ms - period * fetch_limit)
wrong_since = now_ms - period * chart_initial_limit("15m")
self.assertGreater(since, wrong_since)
def test_thin_series_tail_refresh_fetches_full_window(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
last_closed = last_closed_bar_open_ms("15m", now)
bars = [
{
"open_time_ms": last_closed - period * (150 - i),
"open": 100000,
"high": 100100,
"low": 99900,
"close": 100050,
"volume": 1,
}
for i in range(150)
]
fetch_calls: list[dict] = []
def remote_fetch(**kwargs):
fetch_calls.append(dict(kwargs))
return {"ok": True, "bars": bars, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"15m",
remote_fetch,
db_path=self.db,
tail_refresh=True,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(len(out.get("candles") or []), 100)
self.assertGreater(int(out.get("fetched") or 0), 0)
self.assertTrue(any(int(c.get("limit") or 0) > 30 for c in fetch_calls))
def test_resolve_before_ms_exhausted(self):
init_db(self.db)
def remote_fetch(**kwargs):
return {"ok": False, "msg": "no remote"}
out = resolve_chart_bars(
"okx",
"BTC/USDT",
"5m",
remote_fetch,
db_path=self.db,
limit=100,
before_ms=int(time.time() * 1000),
)
self.assertTrue(out.get("ok"))
self.assertEqual(out.get("candles"), [])
self.assertTrue(out.get("exhausted"))
if __name__ == "__main__":
unittest.main()
"""中控 K 线库:分周期保留,聚合与分页读取."""
from __future__ import annotations
import tempfile
import time
import unittest
from pathlib import Path
from lib.hub.hub_kline_store import (
HUB_KLINE_REMOTE_FETCH_CAP,
_since_ms_for_span,
clear_series_bars,
init_db,
load_bars_before,
load_bars_latest,
purge_retention,
purge_timeframe_by_days,
resolve_chart_bars,
retention_days,
trim_contiguous_tail,
upsert_bars,
)
from lib.hub.hub_ohlcv_lib import (
TIMEFRAME_MS,
bar_limit_for_timeframe,
chart_fetch_start_ms,
chart_initial_limit,
last_closed_bar_open_ms,
window_start_ms,
)
class TestHubKlineStore(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.db = Path(self.tmp.name) / "test_hub_kline.db"
def tearDown(self):
self.tmp.cleanup()
def test_bar_limits(self):
self.assertEqual(bar_limit_for_timeframe("5m"), 5000)
self.assertEqual(bar_limit_for_timeframe("1h"), 1000)
self.assertEqual(bar_limit_for_timeframe("1d"), 1000)
self.assertEqual(bar_limit_for_timeframe("1w"), 500)
self.assertEqual(chart_initial_limit("5m"), 2000)
self.assertEqual(chart_initial_limit("1h"), 1000)
self.assertEqual(chart_initial_limit("1d"), 500)
def test_chart_fetch_window_exceeds_retention(self):
now = int(time.time() * 1000)
need = bar_limit_for_timeframe("1d")
fetch_start = chart_fetch_start_ms("1d", need, now)
db_start = window_start_ms("1d", need, retention_days(), now)
self.assertLess(fetch_start, db_start)
def test_purge_retention_5m_one_year(self):
init_db(self.db)
old_ms = int(time.time() * 1000) - 400 * 86400000
upsert_bars(
"okx",
"BTC/USDT",
"5m",
[
{
"open_time_ms": old_ms,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 10,
}
],
self.db,
)
n = purge_timeframe_by_days("5m", 365, self.db)
self.assertGreaterEqual(n, 1)
rows = load_bars_latest("okx", "BTC/USDT", "5m", 10, self.db)
self.assertEqual(len(rows), 0)
def test_purge_retention_keeps_1d(self):
init_db(self.db)
old_ms = int(time.time() * 1000) - 400 * 86400000
upsert_bars(
"okx",
"BTC/USDT",
"1d",
[
{
"open_time_ms": old_ms,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 10,
}
],
self.db,
)
purge_retention(self.db)
rows = load_bars_latest("okx", "BTC/USDT", "1d", 10, self.db)
self.assertEqual(len(rows), 1)
def test_resolve_uses_cache_without_remote(self):
init_db(self.db)
now = int(time.time() * 1000)
tf = "5m"
period = TIMEFRAME_MS[tf]
last_closed = last_closed_bar_open_ms(tf, now)
bars = []
for i in range(400):
oms = last_closed - (399 - i) * period
bars.append(
{
"open_time_ms": oms,
"open": 100 + i,
"high": 101 + i,
"low": 99 + i,
"close": 100.5 + i,
"volume": 1000 + i,
}
)
upsert_bars("okx", "ETH/USDT", tf, bars, self.db)
def remote_fetch(**kwargs):
self.fail("不应请求交易所")
out = resolve_chart_bars(
"okx",
"ETH/USDT",
tf,
remote_fetch,
db_path=self.db,
limit=300,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("candles") or []), 300)
def test_resolve_15m_reads_native_bars(self):
init_db(self.db)
now = int(time.time() * 1000)
period = TIMEFRAME_MS["15m"]
last_closed = last_closed_bar_open_ms("15m", now)
bars = []
for i in range(12):
oms = last_closed - (11 - i) * period
bars.append(
{
"open_time_ms": oms,
"open": 1.0 + i,
"high": 2.0 + i,
"low": 0.5 + i,
"close": 1.5 + i,
"volume": 10.0,
}
)
upsert_bars("okx", "ETH/USDT", "15m", bars, self.db)
def remote_fetch(**kwargs):
self.fail("不应请求交易所")
out = resolve_chart_bars(
"okx",
"ETH/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=10,
)
self.assertTrue(out.get("ok"))
self.assertEqual(out.get("source"), "db")
self.assertEqual(out.get("storage_timeframe"), "15m")
self.assertGreaterEqual(len(out.get("candles") or []), 10)
def test_load_bars_before(self):
init_db(self.db)
period = TIMEFRAME_MS["1h"]
base = 1_700_000_000_000
bars = []
for i in range(5):
bars.append(
{
"open_time_ms": base + i * period,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
)
upsert_bars("okx", "BTC/USDT", "1h", bars, self.db)
before = base + 3 * period
got = load_bars_before("okx", "BTC/USDT", "1h", before, 2, self.db)
self.assertEqual(len(got), 2)
self.assertEqual(got[-1]["open_time_ms"], base + 2 * period)
def test_trim_contiguous_tail_drops_orphan_prefix(self):
period = TIMEFRAME_MS["15m"]
base_old = 1_700_000_000_000
base_new = base_old + period * 500
bars = []
for i in range(3):
bars.append(
{
"open_time_ms": base_old + i * period,
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
)
for i in range(5):
bars.append(
{
"open_time_ms": base_new + i * period,
"open": 2,
"high": 3,
"low": 1.5,
"close": 2.5,
"volume": 2,
}
)
trimmed, split = trim_contiguous_tail(bars, period)
self.assertEqual(split, 3)
self.assertEqual(len(trimmed), 5)
self.assertEqual(trimmed[0]["open_time_ms"], base_new)
def test_resolve_drops_discontinuous_orphans(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
old_ms = now - period * 800
upsert_bars(
"okx",
"ONDO/USDT",
"15m",
[
{
"open_time_ms": old_ms,
"open": 0.33,
"high": 0.34,
"low": 0.32,
"close": 0.335,
"volume": 100,
}
],
self.db,
)
recent = []
start = now - period * 20
for i in range(20):
recent.append(
{
"open_time_ms": start + i * period,
"open": 0.35,
"high": 0.36,
"low": 0.34,
"close": 0.355,
"volume": 50,
}
)
def remote_fetch(**kwargs):
return {"ok": True, "bars": recent, "price_tick": 0.0001}
out = resolve_chart_bars(
"okx",
"ONDO/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=50,
)
self.assertTrue(out.get("ok"))
candles = out.get("candles") or []
self.assertGreaterEqual(len(candles), 19)
if len(candles) >= 2:
for i in range(1, len(candles)):
gap = candles[i]["time"] - candles[i - 1]["time"]
self.assertLessEqual(gap, int(period / 1000 * 3.0))
def test_resolve_refetches_when_db_has_discontinuous_full_count(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
old_start = now - period * 3000
recent_start = now - period * 25
old_bars = [
{
"open_time_ms": old_start + i * period,
"open": 62000,
"high": 62100,
"low": 61900,
"close": 62050,
"volume": 10,
}
for i in range(500)
]
recent = [
{
"open_time_ms": recent_start + i * period,
"open": 104000,
"high": 104100,
"low": 103900,
"close": 104050,
"volume": 20,
}
for i in range(30)
]
upsert_bars("binance", "BTC/USDT", "15m", old_bars, self.db)
upsert_bars("binance", "BTC/USDT", "15m", recent, self.db)
fetch_calls = []
def remote_fetch(**kwargs):
fetch_calls.append(dict(kwargs))
full = []
start = now - period * 120
for i in range(120):
full.append(
{
"open_time_ms": start + i * period,
"open": 104000 + i,
"high": 104100 + i,
"low": 103900 + i,
"close": 104050 + i,
"volume": 30,
}
)
return {"ok": True, "bars": full, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"15m",
remote_fetch,
db_path=self.db,
limit=2000,
)
self.assertTrue(out.get("ok"))
self.assertGreater(len(fetch_calls), 0)
self.assertGreaterEqual(len(out.get("candles") or []), 100)
self.assertGreater(int(out.get("fetched") or 0), 0)
def test_clear_series_and_force_refetch(self):
init_db(self.db)
period = TIMEFRAME_MS["5m"]
now = int(time.time() * 1000)
stale = [
{
"open_time_ms": now - period * (i + 100),
"open": 1,
"high": 2,
"low": 0.5,
"close": 1.5,
"volume": 1,
}
for i in range(40)
]
upsert_bars("binance", "BTC/USDT", "5m", stale, self.db)
self.assertEqual(len(load_bars_latest("binance", "BTC/USDT", "5m", 100, self.db)), 40)
removed = clear_series_bars("binance", "BTC/USDT", "5m", self.db)
self.assertEqual(removed, 40)
self.assertEqual(len(load_bars_latest("binance", "BTC/USDT", "5m", 100, self.db)), 0)
fresh = [
{
"open_time_ms": now - period * (20 - i),
"open": 10,
"high": 11,
"low": 9,
"close": 10.5,
"volume": 2,
}
for i in range(20)
]
def remote_fetch(**kwargs):
return {"ok": True, "bars": fresh, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"5m",
remote_fetch,
db_path=self.db,
force_refresh=True,
clear_db=True,
limit=50,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(int(out.get("cleared") or 0), 0)
self.assertGreater(int(out.get("fetched") or 0), 0)
self.assertGreaterEqual(len(out.get("candles") or []), 19)
def test_since_span_matches_fetch_limit_not_need(self):
period = TIMEFRAME_MS["15m"]
now_ms = 1_800_000_000_000
fetch_limit = HUB_KLINE_REMOTE_FETCH_CAP
since = _since_ms_for_span(
now_ms=now_ms,
period_ms=period,
span_bars=fetch_limit,
cutoff_ms=0,
)
self.assertEqual(since, now_ms - period * fetch_limit)
wrong_since = now_ms - period * chart_initial_limit("15m")
self.assertGreater(since, wrong_since)
def test_thin_series_tail_refresh_fetches_full_window(self):
init_db(self.db)
period = TIMEFRAME_MS["15m"]
now = int(time.time() * 1000)
last_closed = last_closed_bar_open_ms("15m", now)
bars = [
{
"open_time_ms": last_closed - period * (150 - i),
"open": 100000,
"high": 100100,
"low": 99900,
"close": 100050,
"volume": 1,
}
for i in range(150)
]
fetch_calls: list[dict] = []
def remote_fetch(**kwargs):
fetch_calls.append(dict(kwargs))
return {"ok": True, "bars": bars, "price_tick": 0.01}
out = resolve_chart_bars(
"binance",
"BTC/USDT",
"15m",
remote_fetch,
db_path=self.db,
tail_refresh=True,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(len(out.get("candles") or []), 100)
self.assertGreater(int(out.get("fetched") or 0), 0)
self.assertTrue(any(int(c.get("limit") or 0) > 30 for c in fetch_calls))
def test_resolve_before_ms_exhausted(self):
init_db(self.db)
def remote_fetch(**kwargs):
return {"ok": False, "msg": "no remote"}
out = resolve_chart_bars(
"okx",
"BTC/USDT",
"5m",
remote_fetch,
db_path=self.db,
limit=100,
before_ms=int(time.time() * 1000),
)
self.assertTrue(out.get("ok"))
self.assertEqual(out.get("candles"), [])
self.assertTrue(out.get("exhausted"))
if __name__ == "__main__":
unittest.main()
+39 -39
View File
@@ -1,39 +1,39 @@
"""hub /api/hub/monitorenrich 局部返回时须保留 keys"""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.hub.hub_bridge import build_hub_monitor_payload # noqa: E402
class TestHubMonitorPayload(unittest.TestCase):
def test_partial_enrich_keeps_keys(self):
keys = [{"id": 7, "symbol": "BTC/USDT"}]
orders = [{"id": 1}]
trends = [{"id": 9, "symbol": "ETH/USDT"}]
rolls = []
def enrich_only_trends(**_kw):
return {"trends": [{"id": 9, "add_count": 2}]}
out = build_hub_monitor_payload(
keys=keys,
orders=orders,
trends=trends,
rolls=rolls,
enrich=enrich_only_trends,
)
self.assertTrue(out["ok"])
self.assertEqual(out["keys"], keys)
self.assertEqual(out["orders"], orders)
self.assertEqual(out["rolls"], rolls)
self.assertEqual(out["trends"][0]["add_count"], 2)
if __name__ == "__main__":
unittest.main()
"""hub /api/hub/monitor:enrich 局部返回时须保留 keys."""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.hub.hub_bridge import build_hub_monitor_payload # noqa: E402
class TestHubMonitorPayload(unittest.TestCase):
def test_partial_enrich_keeps_keys(self):
keys = [{"id": 7, "symbol": "BTC/USDT"}]
orders = [{"id": 1}]
trends = [{"id": 9, "symbol": "ETH/USDT"}]
rolls = []
def enrich_only_trends(**_kw):
return {"trends": [{"id": 9, "add_count": 2}]}
out = build_hub_monitor_payload(
keys=keys,
orders=orders,
trends=trends,
rolls=rolls,
enrich=enrich_only_trends,
)
self.assertTrue(out["ok"])
self.assertEqual(out["keys"], keys)
self.assertEqual(out["orders"], orders)
self.assertEqual(out["rolls"], rolls)
self.assertEqual(out["trends"][0]["add_count"], 2)
if __name__ == "__main__":
unittest.main()
+222 -222
View File
@@ -1,222 +1,222 @@
"""hub_ohlcv_lib分页拉取Gate 等单次不足 chunk 时仍继续)。"""
from __future__ import annotations
import unittest
from lib.hub.hub_ohlcv_lib import (
aggregate_ohlcv_bars,
bars_spacing_matches_timeframe,
fetch_ohlcv_for_hub,
normalize_price_tick,
price_tick_from_market,
)
class _FakeExchange:
def __init__(self, pages, *, timeframes=None):
self.pages = list(pages)
self.calls = []
self.markets = {}
self.timeframes = timeframes if timeframes is not None else {}
def fetch_ohlcv(self, symbol, timeframe=None, since=None, limit=None):
self.calls.append(
{"symbol": symbol, "since": since, "limit": limit, "timeframe": timeframe}
)
if not self.pages:
return []
page = self.pages.pop(0)
if since is None:
return page
return [b for b in page if b[0] >= since]
class TestHubOhlcvLib(unittest.TestCase):
def test_normalize_price_tick_snaps_powers_of_ten(self):
self.assertAlmostEqual(normalize_price_tick(0.00001), 0.00001)
self.assertAlmostEqual(normalize_price_tick(0.001), 0.001)
self.assertIsNone(normalize_price_tick(0))
def test_price_tick_from_decimal_precision(self):
class _Ex:
markets = {"BTC/USDT:USDT": {"precision": {"price": 2}, "info": {}, "limits": {}}}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "12345.67"
tick = price_tick_from_market(_Ex(), "BTC/USDT:USDT")
self.assertAlmostEqual(tick, 0.01)
def test_price_tick_from_binance_price_filter(self):
class _Ex:
markets = {
"BTC/USDT:USDT": {
"precision": {"price": 2},
"info": {
"filters": [
{"filterType": "PRICE_FILTER", "tickSize": "0.10"},
{"filterType": "LOT_SIZE", "stepSize": "0.001"},
]
},
"limits": {},
}
}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "12345.6"
from lib.hub.hub_ohlcv_lib import price_tick_from_market
tick = price_tick_from_market(_Ex(), "BTC/USDT:USDT")
self.assertAlmostEqual(tick, 0.10)
def test_price_tick_from_info_tick_size(self):
class _Ex:
markets = {
"INJ/USDT:USDT": {
"precision": {"price": 4},
"info": {"tickSize": "0.001"},
"limits": {},
}
}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "7.123"
from lib.hub.hub_ohlcv_lib import price_tick_from_market
tick = price_tick_from_market(_Ex(), "INJ/USDT:USDT")
self.assertAlmostEqual(tick, 0.001)
def test_full_fetch_without_since_paginates_okx_style(self):
"""OKX 等无 since 单次约 300 根须分页至 limit"""
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
step = TIMEFRAME_MS["1h"]
want = 1000
base = max(0, int(__import__("time").time() * 1000) - want * step)
pages = [
[[base + i * step, 1.0, 1.1, 0.9, 1.05, 100.0] for i in range(300)],
[[base + (300 + i) * step, 2.0, 2.1, 1.9, 2.05, 200.0] for i in range(300)],
[[base + (600 + i) * step, 3.0, 3.1, 2.9, 3.05, 300.0] for i in range(300)],
[[base + (900 + i) * step, 4.0, 4.1, 3.9, 4.05, 400.0] for i in range(100)],
]
ex = _FakeExchange(pages)
out = fetch_ohlcv_for_hub(
symbol="ONDO/USDT",
timeframe="1h",
since_ms=None,
limit=want,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("bars") or []), 1000)
self.assertGreaterEqual(len(ex.calls), 4)
self.assertAlmostEqual(out["bars"][-1]["close"], 4.05)
def test_pagination_continues_when_page_smaller_than_chunk(self):
"""Gate 等常返回 299 根/次不应误判为已到末尾"""
base = 1_700_000_000_000
step = 4 * 60 * 60 * 1000
page1 = [
[base + i * step, 1.0, 1.1, 0.9, 1.05, 100.0] for i in range(299)
]
page2 = [
[base + (299 + i) * step, 2.0, 2.1, 1.9, 2.05, 200.0] for i in range(299)
]
page3 = [
[base + (598 + i) * step, 3.0, 3.1, 2.9, 3.05, 300.0] for i in range(50)
]
ex = _FakeExchange([page1, page2, page3])
out = fetch_ohlcv_for_hub(
symbol="INJ/USDT",
timeframe="4h",
since_ms=base,
limit=600,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("bars") or []), 600)
self.assertGreaterEqual(len(ex.calls), 3)
self.assertAlmostEqual(out["bars"][-1]["close"], 3.05)
def test_pagination_stops_when_next_since_reaches_now(self):
"""Gate 等分页 since 不得越过当前时间避免 from>to"""
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
step = TIMEFRAME_MS["1d"]
now_ms = int(__import__("time").time() * 1000)
# 最后一页最后一根 K 的 next_since 将 >= now_ms应停止不再请求
last_open = ((now_ms // step) - 2) * step
page = [
[last_open - step, 1.0, 1.1, 0.9, 1.0, 10.0],
[last_open, 1.1, 1.2, 1.0, 1.1, 11.0],
]
ex = _FakeExchange([page])
out = fetch_ohlcv_for_hub(
symbol="ONDO/USDT",
timeframe="1d",
since_ms=last_open - step * 5,
limit=10,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(len(out.get("bars") or []), 2)
self.assertLessEqual(len(ex.calls), 4)
def test_aggregate_ohlcv_bars_buckets(self):
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
h1 = TIMEFRAME_MS["1h"]
h4 = TIMEFRAME_MS["4h"]
base = (1_700_000_000_000 // h4) * h4
src = [
{
"open_time_ms": base + i * h1,
"open": 1.0,
"high": 2.0,
"low": 0.5,
"close": 1.5,
"volume": 1.0,
}
for i in range(4)
]
out = aggregate_ohlcv_bars(src, "4h")
self.assertEqual(len(out), 1)
self.assertEqual(out[0]["volume"], 4.0)
self.assertEqual(out[0]["high"], 2.0)
self.assertEqual(out[0]["low"], 0.5)
if __name__ == "__main__":
unittest.main()
"""hub_ohlcv_lib:分页拉取(Gate 等单次不足 chunk 时仍继续)."""
from __future__ import annotations
import unittest
from lib.hub.hub_ohlcv_lib import (
aggregate_ohlcv_bars,
bars_spacing_matches_timeframe,
fetch_ohlcv_for_hub,
normalize_price_tick,
price_tick_from_market,
)
class _FakeExchange:
def __init__(self, pages, *, timeframes=None):
self.pages = list(pages)
self.calls = []
self.markets = {}
self.timeframes = timeframes if timeframes is not None else {}
def fetch_ohlcv(self, symbol, timeframe=None, since=None, limit=None):
self.calls.append(
{"symbol": symbol, "since": since, "limit": limit, "timeframe": timeframe}
)
if not self.pages:
return []
page = self.pages.pop(0)
if since is None:
return page
return [b for b in page if b[0] >= since]
class TestHubOhlcvLib(unittest.TestCase):
def test_normalize_price_tick_snaps_powers_of_ten(self):
self.assertAlmostEqual(normalize_price_tick(0.00001), 0.00001)
self.assertAlmostEqual(normalize_price_tick(0.001), 0.001)
self.assertIsNone(normalize_price_tick(0))
def test_price_tick_from_decimal_precision(self):
class _Ex:
markets = {"BTC/USDT:USDT": {"precision": {"price": 2}, "info": {}, "limits": {}}}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "12345.67"
tick = price_tick_from_market(_Ex(), "BTC/USDT:USDT")
self.assertAlmostEqual(tick, 0.01)
def test_price_tick_from_binance_price_filter(self):
class _Ex:
markets = {
"BTC/USDT:USDT": {
"precision": {"price": 2},
"info": {
"filters": [
{"filterType": "PRICE_FILTER", "tickSize": "0.10"},
{"filterType": "LOT_SIZE", "stepSize": "0.001"},
]
},
"limits": {},
}
}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "12345.6"
from lib.hub.hub_ohlcv_lib import price_tick_from_market
tick = price_tick_from_market(_Ex(), "BTC/USDT:USDT")
self.assertAlmostEqual(tick, 0.10)
def test_price_tick_from_info_tick_size(self):
class _Ex:
markets = {
"INJ/USDT:USDT": {
"precision": {"price": 4},
"info": {"tickSize": "0.001"},
"limits": {},
}
}
def load_markets(self):
return self.markets
def market(self, sym):
return self.markets[sym]
def price_to_precision(self, sym, price):
return "7.123"
from lib.hub.hub_ohlcv_lib import price_tick_from_market
tick = price_tick_from_market(_Ex(), "INJ/USDT:USDT")
self.assertAlmostEqual(tick, 0.001)
def test_full_fetch_without_since_paginates_okx_style(self):
"""OKX 等无 since 单次约 300 根,须分页至 limit."""
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
step = TIMEFRAME_MS["1h"]
want = 1000
base = max(0, int(__import__("time").time() * 1000) - want * step)
pages = [
[[base + i * step, 1.0, 1.1, 0.9, 1.05, 100.0] for i in range(300)],
[[base + (300 + i) * step, 2.0, 2.1, 1.9, 2.05, 200.0] for i in range(300)],
[[base + (600 + i) * step, 3.0, 3.1, 2.9, 3.05, 300.0] for i in range(300)],
[[base + (900 + i) * step, 4.0, 4.1, 3.9, 4.05, 400.0] for i in range(100)],
]
ex = _FakeExchange(pages)
out = fetch_ohlcv_for_hub(
symbol="ONDO/USDT",
timeframe="1h",
since_ms=None,
limit=want,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("bars") or []), 1000)
self.assertGreaterEqual(len(ex.calls), 4)
self.assertAlmostEqual(out["bars"][-1]["close"], 4.05)
def test_pagination_continues_when_page_smaller_than_chunk(self):
"""Gate 等常返回 299 根/次,不应误判为已到末尾."""
base = 1_700_000_000_000
step = 4 * 60 * 60 * 1000
page1 = [
[base + i * step, 1.0, 1.1, 0.9, 1.05, 100.0] for i in range(299)
]
page2 = [
[base + (299 + i) * step, 2.0, 2.1, 1.9, 2.05, 200.0] for i in range(299)
]
page3 = [
[base + (598 + i) * step, 3.0, 3.1, 2.9, 3.05, 300.0] for i in range(50)
]
ex = _FakeExchange([page1, page2, page3])
out = fetch_ohlcv_for_hub(
symbol="INJ/USDT",
timeframe="4h",
since_ms=base,
limit=600,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertEqual(len(out.get("bars") or []), 600)
self.assertGreaterEqual(len(ex.calls), 3)
self.assertAlmostEqual(out["bars"][-1]["close"], 3.05)
def test_pagination_stops_when_next_since_reaches_now(self):
"""Gate 等:分页 since 不得越过当前时间,避免 from>to."""
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
step = TIMEFRAME_MS["1d"]
now_ms = int(__import__("time").time() * 1000)
# 最后一页最后一根 K 的 next_since 将 >= now_ms,应停止不再请求
last_open = ((now_ms // step) - 2) * step
page = [
[last_open - step, 1.0, 1.1, 0.9, 1.0, 10.0],
[last_open, 1.1, 1.2, 1.0, 1.1, 11.0],
]
ex = _FakeExchange([page])
out = fetch_ohlcv_for_hub(
symbol="ONDO/USDT",
timeframe="1d",
since_ms=last_open - step * 5,
limit=10,
normalize_symbol_input=lambda s: str(s).strip().upper(),
normalize_exchange_symbol=lambda s: f"{s}:USDT" if ":" not in s else s,
ensure_markets_loaded=lambda: None,
exchange=ex,
)
self.assertTrue(out.get("ok"))
self.assertGreaterEqual(len(out.get("bars") or []), 2)
self.assertLessEqual(len(ex.calls), 4)
def test_aggregate_ohlcv_bars_buckets(self):
from lib.hub.hub_ohlcv_lib import TIMEFRAME_MS
h1 = TIMEFRAME_MS["1h"]
h4 = TIMEFRAME_MS["4h"]
base = (1_700_000_000_000 // h4) * h4
src = [
{
"open_time_ms": base + i * h1,
"open": 1.0,
"high": 2.0,
"low": 0.5,
"close": 1.5,
"volume": 1.0,
}
for i in range(4)
]
out = aggregate_ohlcv_bars(src, "4h")
self.assertEqual(len(out), 1)
self.assertEqual(out[0]["volume"], 4.0)
self.assertEqual(out[0]["high"], 2.0)
self.assertEqual(out[0]["low"], 0.5)
if __name__ == "__main__":
unittest.main()
+1 -1
View File
@@ -1,4 +1,4 @@
"""中控改委托同步与条件单按角色去重"""
"""中控改委托同步与条件单按角色去重."""
from lib.hub.hub_order_sync_lib import (
cond_order_role,
+1 -1
View File
@@ -1,4 +1,4 @@
"""hub_supervisor_lib 单元测试"""
"""hub_supervisor_lib 单元测试."""
from __future__ import annotations
import json
+348 -348
View File
@@ -1,348 +1,348 @@
"""币种档案库5m 聚合与视窗计算"""
from __future__ import annotations
import tempfile
from pathlib import Path
from lib.hub.hub_ohlcv_lib import aggregate_ohlcv_bars
from datetime import datetime, timezone
from zoneinfo import ZoneInfo
from lib.hub.hub_symbol_archive_lib import (
CHART_DISPLAY_TZ,
_compute_period_stats,
_fill_missing_bars,
init_db,
list_daily_trades,
load_symbol_trades,
ms_to_wall_clock_str,
parse_wall_clock_ms,
resolve_archive_chart,
trading_day_bounds_ms,
upsert_bars_5m,
upsert_trade_overlay,
list_symbol_rows,
upsert_trades_cache,
)
def _seed_5m_bars(
db: Path,
start_ms: int,
count: int,
step: int = 300_000,
*,
ex: str = "gate",
sym: str = "ONDO",
) -> None:
bars = []
price = 1.0
for i in range(count):
o = start_ms + i * step
price += 0.001
bars.append(
{
"open_time_ms": o,
"open": price,
"high": price + 0.002,
"low": price - 0.001,
"close": price + 0.001,
"volume": 100 + i,
}
)
upsert_bars_5m(ex, sym, bars, db_path=db)
def test_aggregate_15m_from_5m():
start = 1_700_000_000_000
bars = []
for i in range(6):
t = start + i * 300_000
bars.append(
{
"open_time_ms": t,
"open": 1.0,
"high": 1.1,
"low": 0.9,
"close": 1.05,
"volume": 10,
}
)
agg = aggregate_ohlcv_bars(bars, "15m")
assert len(agg) >= 1
assert agg[-1]["close"] == bars[-1]["close"]
assert agg[0]["open_time_ms"] <= agg[1]["open_time_ms"]
def test_resolve_archive_chart_15m():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
anchor = 1_700_000_000_000
_seed_5m_bars(db, anchor - 50 * 300_000, 120)
out = resolve_archive_chart(
"gate",
"ONDO",
"15m",
anchor_ms=anchor,
mode="hold",
bars=40,
db_path=db,
)
assert out["ok"] is True
assert out["timeframe"] == "15m"
assert len(out["candles"]) >= 10
def test_fill_missing_bars_continuity():
period = 300_000
start = (1_700_000_000_000 // period) * period
bars = [
{
"open_time_ms": start,
"open": 1.0,
"high": 1.1,
"low": 0.9,
"close": 1.05,
"volume": 10,
},
{
"open_time_ms": start + period * 2,
"open": 1.05,
"high": 1.15,
"low": 1.0,
"close": 1.1,
"volume": 8,
},
]
filled = _fill_missing_bars(bars, period, start, start + period * 2)
assert len(filled) >= 3
assert any(b.get("filled") for b in filled)
def test_resolve_archive_chart_history_range():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
open_ms = 1_700_000_000_000
close_ms = open_ms + 6 * 3600_000
_seed_5m_bars(db, open_ms - 20 * 300_000, 200, ex="gate", sym="BNB/USDT")
out = resolve_archive_chart(
"gate",
"BNB/USDT",
"15m",
opened_ms=open_ms,
closed_ms=close_ms,
mode="hold",
range_mode="history",
db_path=db,
)
assert out["ok"] is True
assert out.get("range_mode") == "history"
assert out.get("window_end_ms") <= close_ms + 4 * 3600_000
assert len(out["candles"]) >= 40
def test_sync_prunes_missing_trades():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{"id": 1, "symbol": "BNB/USDT", "result": "止损", "pnl_amount": -1},
{"id": 2, "symbol": "BNB/USDT", "result": "止盈", "pnl_amount": 1},
],
db_path=db,
prune_missing=False,
)
stats = upsert_trades_cache(
"gate",
[{"id": 1, "symbol": "BNB/USDT", "result": "止损", "pnl_amount": -1}],
db_path=db,
prune_missing=True,
)
rows = load_symbol_trades("gate", "BNB/USDT", db_path=db)
assert len(rows) == 1
assert rows[0]["trade_id"] == 1
assert stats["removed"] == 1
def test_list_with_overlay_filters():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 1,
"symbol": "ONDO",
"direction": "long",
"result": "止盈",
"pnl_amount": 12.5,
"opened_at": "2026-01-01 10:00:00",
"closed_at": "2026-01-01 12:00:00",
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
},
{
"id": 2,
"symbol": "ONDO",
"direction": "short",
"result": "止损",
"pnl_amount": -3.2,
"opened_at": "2026-01-02 10:00:00",
"closed_at": "2026-01-02 11:00:00",
"opened_at_ms": 1_700_086_400_000,
"closed_at_ms": 1_700_090_000_000,
},
],
db_path=db,
)
upsert_trade_overlay("gate", 2, behavior_tag="sick", note="追高", db_path=db)
rows = list_symbol_rows(db_path=db)
assert len(rows) == 1
assert rows[0]["trade_count"] == 2
sick_only = list_symbol_rows(filter_sick=True, db_path=db)
assert len(sick_only) == 1
profit_only = list_symbol_rows(filter_profit=True, db_path=db)
assert len(profit_only) == 1
def test_parse_wall_clock_ms_uses_utc_plus_8():
ms = parse_wall_clock_ms("2026-06-07 20:30:00")
assert ms is not None
dt_utc = datetime.fromtimestamp(ms / 1000.0, tz=timezone.utc)
dt_bj = dt_utc.astimezone(CHART_DISPLAY_TZ)
assert dt_bj.strftime("%Y-%m-%d %H:%M:%S") == "2026-06-07 20:30:00"
assert ms_to_wall_clock_str(ms) == "2026-06-07 20:30:00"
assert parse_wall_clock_ms("2026-06-07 20:30") == ms
def test_parse_wall_clock_ms_accepts_epoch_strings():
ms = 1_700_000_000_000
assert parse_wall_clock_ms(str(ms)) == ms
assert parse_wall_clock_ms(str(ms // 1000)) == ms
def test_resolve_archive_chart_history_uses_trade_span_not_200_bars():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
opened = 1_700_000_000_000
closed = opened + 20 * 24 * 3600_000
_seed_5m_bars(db, opened - 35 * 24 * 3600_000, 40 * 24 * 12)
out = resolve_archive_chart(
"gate",
"ONDO",
"15m",
opened_ms=opened,
closed_ms=closed,
mode="hold",
bars=200,
range_mode="history",
db_path=db,
)
assert out["ok"] is True
assert out["range_mode"] == "history"
assert out["bar_count"] > 200
def test_upsert_forces_sync_exchange_key():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 77,
"exchange_key": "gate",
"account_exchange_key": "gate",
"symbol": "ETH/USDT",
"result": "止损",
"pnl_amount": -1,
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
}
],
db_path=db,
)
rows = load_symbol_trades("gate", "ETH/USDT", db_path=db)
assert len(rows) == 1
assert rows[0]["exchange_key"] == "gate"
assert "account_exchange_key" not in rows[0]
def test_compute_period_stats_win_loss_metrics():
rows = [
{"exchange_key": "binance", "pnl_amount": 10.0, "behavior_tag": ""},
{"exchange_key": "binance", "pnl_amount": 4.0, "behavior_tag": ""},
{"exchange_key": "okx", "pnl_amount": -3.0, "behavior_tag": "sick"},
{"exchange_key": "okx", "pnl_amount": -6.0, "behavior_tag": ""},
]
st = _compute_period_stats(rows)
assert st["open_count"] == 4
assert st["win_count"] == 2
assert st["loss_count"] == 2
assert st["avg_win"] == 7.0
assert st["avg_loss"] == -4.5
assert st["max_win"] == 10.0
assert st["max_loss"] == -6.0
assert st["win_rate"] == 50.0
assert st["profit_loss_ratio"] == round(7.0 / 4.5, 2)
assert st["sick_count"] == 1
assert st["pnl_total"] == 5.0
assert st["pnl_ex_sick"] == 8.0
assert st["by_exchange"]["binance"]["win_count"] == 2
assert st["by_exchange"]["binance"]["win_rate"] == 100.0
assert st["by_exchange"]["binance"]["profit_loss_ratio"] is None
def test_list_daily_trades_search_filters_stats():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
day = "2023-11-15"
start_ms, _ = trading_day_bounds_ms(day)
btc_close = start_ms + 3_600_000
eth_close = start_ms + 7_200_000
upsert_trades_cache(
"gate",
[
{
"id": 1,
"symbol": "BTC/USDT",
"result": "止盈",
"pnl_amount": 5.0,
"opened_at_ms": start_ms,
"closed_at_ms": btc_close,
},
{
"id": 2,
"symbol": "ETH/USDT",
"result": "止损",
"pnl_amount": -2.0,
"opened_at_ms": btc_close,
"closed_at_ms": eth_close,
},
],
db_path=db,
)
payload = list_daily_trades(
period="range",
date_from=day,
date_to=day,
search="btc",
db_path=db,
)
assert len(payload["trades"]) == 1
assert payload["trades"][0]["symbol"] == "BTC/USDT"
st = payload["stats"]
assert st["open_count"] == 1
assert st["win_count"] == 1
assert st["loss_count"] == 0
assert st["max_win"] == 5.0
assert st["pnl_total"] == 5.0
"""币种档案库:5m 聚合与视窗计算."""
from __future__ import annotations
import tempfile
from pathlib import Path
from lib.hub.hub_ohlcv_lib import aggregate_ohlcv_bars
from datetime import datetime, timezone
from zoneinfo import ZoneInfo
from lib.hub.hub_symbol_archive_lib import (
CHART_DISPLAY_TZ,
_compute_period_stats,
_fill_missing_bars,
init_db,
list_daily_trades,
load_symbol_trades,
ms_to_wall_clock_str,
parse_wall_clock_ms,
resolve_archive_chart,
trading_day_bounds_ms,
upsert_bars_5m,
upsert_trade_overlay,
list_symbol_rows,
upsert_trades_cache,
)
def _seed_5m_bars(
db: Path,
start_ms: int,
count: int,
step: int = 300_000,
*,
ex: str = "gate",
sym: str = "ONDO",
) -> None:
bars = []
price = 1.0
for i in range(count):
o = start_ms + i * step
price += 0.001
bars.append(
{
"open_time_ms": o,
"open": price,
"high": price + 0.002,
"low": price - 0.001,
"close": price + 0.001,
"volume": 100 + i,
}
)
upsert_bars_5m(ex, sym, bars, db_path=db)
def test_aggregate_15m_from_5m():
start = 1_700_000_000_000
bars = []
for i in range(6):
t = start + i * 300_000
bars.append(
{
"open_time_ms": t,
"open": 1.0,
"high": 1.1,
"low": 0.9,
"close": 1.05,
"volume": 10,
}
)
agg = aggregate_ohlcv_bars(bars, "15m")
assert len(agg) >= 1
assert agg[-1]["close"] == bars[-1]["close"]
assert agg[0]["open_time_ms"] <= agg[1]["open_time_ms"]
def test_resolve_archive_chart_15m():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
anchor = 1_700_000_000_000
_seed_5m_bars(db, anchor - 50 * 300_000, 120)
out = resolve_archive_chart(
"gate",
"ONDO",
"15m",
anchor_ms=anchor,
mode="hold",
bars=40,
db_path=db,
)
assert out["ok"] is True
assert out["timeframe"] == "15m"
assert len(out["candles"]) >= 10
def test_fill_missing_bars_continuity():
period = 300_000
start = (1_700_000_000_000 // period) * period
bars = [
{
"open_time_ms": start,
"open": 1.0,
"high": 1.1,
"low": 0.9,
"close": 1.05,
"volume": 10,
},
{
"open_time_ms": start + period * 2,
"open": 1.05,
"high": 1.15,
"low": 1.0,
"close": 1.1,
"volume": 8,
},
]
filled = _fill_missing_bars(bars, period, start, start + period * 2)
assert len(filled) >= 3
assert any(b.get("filled") for b in filled)
def test_resolve_archive_chart_history_range():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
open_ms = 1_700_000_000_000
close_ms = open_ms + 6 * 3600_000
_seed_5m_bars(db, open_ms - 20 * 300_000, 200, ex="gate", sym="BNB/USDT")
out = resolve_archive_chart(
"gate",
"BNB/USDT",
"15m",
opened_ms=open_ms,
closed_ms=close_ms,
mode="hold",
range_mode="history",
db_path=db,
)
assert out["ok"] is True
assert out.get("range_mode") == "history"
assert out.get("window_end_ms") <= close_ms + 4 * 3600_000
assert len(out["candles"]) >= 40
def test_sync_prunes_missing_trades():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{"id": 1, "symbol": "BNB/USDT", "result": "止损", "pnl_amount": -1},
{"id": 2, "symbol": "BNB/USDT", "result": "止盈", "pnl_amount": 1},
],
db_path=db,
prune_missing=False,
)
stats = upsert_trades_cache(
"gate",
[{"id": 1, "symbol": "BNB/USDT", "result": "止损", "pnl_amount": -1}],
db_path=db,
prune_missing=True,
)
rows = load_symbol_trades("gate", "BNB/USDT", db_path=db)
assert len(rows) == 1
assert rows[0]["trade_id"] == 1
assert stats["removed"] == 1
def test_list_with_overlay_filters():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 1,
"symbol": "ONDO",
"direction": "long",
"result": "止盈",
"pnl_amount": 12.5,
"opened_at": "2026-01-01 10:00:00",
"closed_at": "2026-01-01 12:00:00",
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
},
{
"id": 2,
"symbol": "ONDO",
"direction": "short",
"result": "止损",
"pnl_amount": -3.2,
"opened_at": "2026-01-02 10:00:00",
"closed_at": "2026-01-02 11:00:00",
"opened_at_ms": 1_700_086_400_000,
"closed_at_ms": 1_700_090_000_000,
},
],
db_path=db,
)
upsert_trade_overlay("gate", 2, behavior_tag="sick", note="追高", db_path=db)
rows = list_symbol_rows(db_path=db)
assert len(rows) == 1
assert rows[0]["trade_count"] == 2
sick_only = list_symbol_rows(filter_sick=True, db_path=db)
assert len(sick_only) == 1
profit_only = list_symbol_rows(filter_profit=True, db_path=db)
assert len(profit_only) == 1
def test_parse_wall_clock_ms_uses_utc_plus_8():
ms = parse_wall_clock_ms("2026-06-07 20:30:00")
assert ms is not None
dt_utc = datetime.fromtimestamp(ms / 1000.0, tz=timezone.utc)
dt_bj = dt_utc.astimezone(CHART_DISPLAY_TZ)
assert dt_bj.strftime("%Y-%m-%d %H:%M:%S") == "2026-06-07 20:30:00"
assert ms_to_wall_clock_str(ms) == "2026-06-07 20:30:00"
assert parse_wall_clock_ms("2026-06-07 20:30") == ms
def test_parse_wall_clock_ms_accepts_epoch_strings():
ms = 1_700_000_000_000
assert parse_wall_clock_ms(str(ms)) == ms
assert parse_wall_clock_ms(str(ms // 1000)) == ms
def test_resolve_archive_chart_history_uses_trade_span_not_200_bars():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
opened = 1_700_000_000_000
closed = opened + 20 * 24 * 3600_000
_seed_5m_bars(db, opened - 35 * 24 * 3600_000, 40 * 24 * 12)
out = resolve_archive_chart(
"gate",
"ONDO",
"15m",
opened_ms=opened,
closed_ms=closed,
mode="hold",
bars=200,
range_mode="history",
db_path=db,
)
assert out["ok"] is True
assert out["range_mode"] == "history"
assert out["bar_count"] > 200
def test_upsert_forces_sync_exchange_key():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 77,
"exchange_key": "gate",
"account_exchange_key": "gate",
"symbol": "ETH/USDT",
"result": "止损",
"pnl_amount": -1,
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
}
],
db_path=db,
)
rows = load_symbol_trades("gate", "ETH/USDT", db_path=db)
assert len(rows) == 1
assert rows[0]["exchange_key"] == "gate"
assert "account_exchange_key" not in rows[0]
def test_compute_period_stats_win_loss_metrics():
rows = [
{"exchange_key": "binance", "pnl_amount": 10.0, "behavior_tag": ""},
{"exchange_key": "binance", "pnl_amount": 4.0, "behavior_tag": ""},
{"exchange_key": "okx", "pnl_amount": -3.0, "behavior_tag": "sick"},
{"exchange_key": "okx", "pnl_amount": -6.0, "behavior_tag": ""},
]
st = _compute_period_stats(rows)
assert st["open_count"] == 4
assert st["win_count"] == 2
assert st["loss_count"] == 2
assert st["avg_win"] == 7.0
assert st["avg_loss"] == -4.5
assert st["max_win"] == 10.0
assert st["max_loss"] == -6.0
assert st["win_rate"] == 50.0
assert st["profit_loss_ratio"] == round(7.0 / 4.5, 2)
assert st["sick_count"] == 1
assert st["pnl_total"] == 5.0
assert st["pnl_ex_sick"] == 8.0
assert st["by_exchange"]["binance"]["win_count"] == 2
assert st["by_exchange"]["binance"]["win_rate"] == 100.0
assert st["by_exchange"]["binance"]["profit_loss_ratio"] is None
def test_list_daily_trades_search_filters_stats():
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
day = "2023-11-15"
start_ms, _ = trading_day_bounds_ms(day)
btc_close = start_ms + 3_600_000
eth_close = start_ms + 7_200_000
upsert_trades_cache(
"gate",
[
{
"id": 1,
"symbol": "BTC/USDT",
"result": "止盈",
"pnl_amount": 5.0,
"opened_at_ms": start_ms,
"closed_at_ms": btc_close,
},
{
"id": 2,
"symbol": "ETH/USDT",
"result": "止损",
"pnl_amount": -2.0,
"opened_at_ms": btc_close,
"closed_at_ms": eth_close,
},
],
db_path=db,
)
payload = list_daily_trades(
period="range",
date_from=day,
date_to=day,
search="btc",
db_path=db,
)
assert len(payload["trades"]) == 1
assert payload["trades"][0]["symbol"] == "BTC/USDT"
st = payload["stats"]
assert st["open_count"] == 1
assert st["win_count"] == 1
assert st["loss_count"] == 0
assert st["max_win"] == 5.0
assert st["pnl_total"] == 5.0
+102 -102
View File
@@ -1,102 +1,102 @@
"""档案交易strategy_trade_snapshots 补全 gate 漏记"""
from __future__ import annotations
import sqlite3
import tempfile
from datetime import datetime, timedelta
from pathlib import Path
from lib.hub.hub_trades_lib import fetch_trades_for_archive
def _init_db(path: Path) -> sqlite3.Connection:
conn = sqlite3.connect(str(path))
conn.row_factory = sqlite3.Row
conn.execute(
"""
CREATE TABLE trade_records (
id INTEGER PRIMARY KEY,
symbol TEXT,
direction TEXT,
result TEXT,
pnl_amount REAL,
opened_at TEXT,
closed_at TEXT,
opened_at_ms INTEGER,
closed_at_ms INTEGER,
created_at TEXT,
trend_plan_id INTEGER
)
"""
)
conn.execute(
"""
CREATE TABLE strategy_trade_snapshots (
id INTEGER PRIMARY KEY,
strategy_type TEXT,
source_id INTEGER,
symbol TEXT,
direction TEXT,
result_label TEXT,
status_at_close TEXT,
opened_at TEXT,
closed_at TEXT,
pnl_amount REAL,
snapshot_json TEXT,
created_at TEXT
)
"""
)
return conn
def test_merge_snapshot_when_trade_record_missing():
with tempfile.TemporaryDirectory() as td:
conn = _init_db(Path(td) / "t.db")
closed = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
conn.execute(
"""
INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction,
result_label, opened_at, closed_at, pnl_amount, snapshot_json, created_at
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(7, "trend_pullback", 42, "ONDO/USDT", "long", "止损", closed, closed, -1.2, "{}", closed),
)
conn.commit()
trades = fetch_trades_for_archive(conn, days=30, limit=50)
conn.close()
assert len(trades) == 1
assert trades[0]["symbol"] == "ONDO/USDT"
assert trades[0]["id"] == -7
assert trades[0].get("from_snapshot") is True
def test_skip_snapshot_when_trade_record_exists():
with tempfile.TemporaryDirectory() as td:
conn = _init_db(Path(td) / "t.db")
closed = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
conn.execute(
"""
INSERT INTO trade_records (
id, symbol, direction, result, pnl_amount,
opened_at, closed_at, opened_at_ms, closed_at_ms, created_at, trend_plan_id
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(1, "ONDO/USDT", "long", "止损", -1.2, closed, closed, 1, 2, closed, 42),
)
conn.execute(
"""
INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction,
result_label, opened_at, closed_at, pnl_amount, snapshot_json, created_at
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(7, "trend_pullback", 42, "ONDO/USDT", "long", "止损", closed, closed, -1.2, "{}", closed),
)
conn.commit()
trades = fetch_trades_for_archive(conn, days=30, limit=50)
conn.close()
assert len(trades) == 1
assert trades[0]["id"] == 1
"""档案交易:strategy_trade_snapshots 补全 gate 漏记."""
from __future__ import annotations
import sqlite3
import tempfile
from datetime import datetime, timedelta
from pathlib import Path
from lib.hub.hub_trades_lib import fetch_trades_for_archive
def _init_db(path: Path) -> sqlite3.Connection:
conn = sqlite3.connect(str(path))
conn.row_factory = sqlite3.Row
conn.execute(
"""
CREATE TABLE trade_records (
id INTEGER PRIMARY KEY,
symbol TEXT,
direction TEXT,
result TEXT,
pnl_amount REAL,
opened_at TEXT,
closed_at TEXT,
opened_at_ms INTEGER,
closed_at_ms INTEGER,
created_at TEXT,
trend_plan_id INTEGER
)
"""
)
conn.execute(
"""
CREATE TABLE strategy_trade_snapshots (
id INTEGER PRIMARY KEY,
strategy_type TEXT,
source_id INTEGER,
symbol TEXT,
direction TEXT,
result_label TEXT,
status_at_close TEXT,
opened_at TEXT,
closed_at TEXT,
pnl_amount REAL,
snapshot_json TEXT,
created_at TEXT
)
"""
)
return conn
def test_merge_snapshot_when_trade_record_missing():
with tempfile.TemporaryDirectory() as td:
conn = _init_db(Path(td) / "t.db")
closed = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
conn.execute(
"""
INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction,
result_label, opened_at, closed_at, pnl_amount, snapshot_json, created_at
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(7, "trend_pullback", 42, "ONDO/USDT", "long", "止损", closed, closed, -1.2, "{}", closed),
)
conn.commit()
trades = fetch_trades_for_archive(conn, days=30, limit=50)
conn.close()
assert len(trades) == 1
assert trades[0]["symbol"] == "ONDO/USDT"
assert trades[0]["id"] == -7
assert trades[0].get("from_snapshot") is True
def test_skip_snapshot_when_trade_record_exists():
with tempfile.TemporaryDirectory() as td:
conn = _init_db(Path(td) / "t.db")
closed = (datetime.now() - timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
conn.execute(
"""
INSERT INTO trade_records (
id, symbol, direction, result, pnl_amount,
opened_at, closed_at, opened_at_ms, closed_at_ms, created_at, trend_plan_id
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(1, "ONDO/USDT", "long", "止损", -1.2, closed, closed, 1, 2, closed, 42),
)
conn.execute(
"""
INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction,
result_label, opened_at, closed_at, pnl_amount, snapshot_json, created_at
) VALUES (?,?,?,?,?,?,?,?,?,?,?)
""",
(7, "trend_pullback", 42, "ONDO/USDT", "long", "止损", closed, closed, -1.2, "{}", closed),
)
conn.commit()
trades = fetch_trades_for_archive(conn, days=30, limit=50)
conn.close()
assert len(trades) == 1
assert trades[0]["id"] == 1
+229 -229
View File
@@ -1,229 +1,229 @@
"""hub_trades_lib 单元测试"""
from __future__ import annotations
import sqlite3
import unittest
from datetime import datetime
from lib.hub.hub_trades_lib import (
attach_journal_mood_tags,
fetch_trades_for_trading_day,
journal_trade_match_key,
summarize_trades,
trading_day_from_dt,
trading_day_window_bounds,
)
class HubTradesLibTest(unittest.TestCase):
def test_trading_day_reset(self):
dt = datetime(2026, 6, 6, 7, 30, 0)
self.assertEqual(trading_day_from_dt(dt, 8), "2026-06-05")
dt2 = datetime(2026, 6, 6, 8, 0, 0)
self.assertEqual(trading_day_from_dt(dt2, 8), "2026-06-06")
def test_trading_day_window_bounds(self):
start, end = trading_day_window_bounds("2026-06-06", 8)
self.assertEqual(start, "2026-06-06 08:00:00")
self.assertEqual(end, "2026-06-07 07:59:59")
def test_fetch_and_summarize(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"ONDO/USDT",
"short",
"止损",
None,
-0.5,
None,
None,
"2026-06-06 10:00:00",
None,
"2026-06-06 09:00:00",
None,
"2026-06-06 10:00:00",
"趋势回调",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
stats = summarize_trades(rows)
self.assertEqual(stats["closed_count"], 1)
self.assertEqual(stats["loss_count"], 1)
self.assertAlmostEqual(stats["total_pnl_u"], -0.5)
conn.close()
def test_early_morning_belongs_prev_trading_day(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"BTC/USDT",
"long",
"止盈",
None,
1.2,
None,
None,
"2026-06-07 07:30:00",
None,
"2026-06-07 06:00:00",
None,
"2026-06-07 07:30:00",
"关键位",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
self.assertEqual(len(fetch_trades_for_trading_day(conn, "2026-06-07")), 0)
self.assertEqual(len(fetch_trades_for_trading_day(conn, "2026-06-06")), 1)
conn.close()
def test_reviewed_fields_preferred(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"ETH/USDT",
"long",
"止损",
"止盈",
-0.5,
2.0,
None,
"2026-06-06 09:00:00",
"2026-06-06 11:00:00",
"2026-06-06 08:00:00",
None,
"2026-06-06 11:00:00",
"趋势回调",
None,
None,
"trend",
"",
"2026-06-06 12:00:00",
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["result"], "止盈")
self.assertAlmostEqual(rows[0]["pnl_amount"], 2.0)
self.assertTrue(rows[0]["reviewed"])
conn.close()
def test_time_close_result_included(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"BTC/USDT",
"long",
"时间平仓",
None,
1.2,
None,
None,
"2026-06-06 12:00:00",
None,
"2026-06-06 08:00:00",
None,
"2026-06-06 12:00:00",
"趋势回调",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["result"], "时间平仓")
conn.close()
def test_attach_journal_mood_tags_marks_sick(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE journal_entries (
coin TEXT, open_datetime TEXT, close_datetime TEXT, mood_issues TEXT, created_at TEXT
)"""
)
conn.execute(
"INSERT INTO journal_entries VALUES (?,?,?,?,?)",
("ETH", "2026-07-06 21:51", "2026-07-07 00:00", "报复开仓,扛单", "2026-07-07 00:05"),
)
conn.commit()
trades = [
{
"id": 42,
"symbol": "ETH/USDT",
"opened_at": "2026-07-06 21:51:00",
"closed_at": "2026-07-07 00:00:00",
}
]
attach_journal_mood_tags(conn, trades, cutoff_s="2026-01-01 00:00:00")
self.assertTrue(trades[0]["journal_mood_sick"])
self.assertEqual(trades[0]["behavior_tag"], "sick")
self.assertEqual(trades[0]["journal_mood_issues"], ["报复开仓", "扛单"])
key = journal_trade_match_key("ETH/USDT", "2026-07-06 21:51:00", "2026-07-07 00:00:00")
self.assertEqual(key, ("ETH", "2026-07-06 21:51", "2026-07-07 00:00"))
conn.close()
if __name__ == "__main__":
unittest.main()
"""hub_trades_lib 单元测试."""
from __future__ import annotations
import sqlite3
import unittest
from datetime import datetime
from lib.hub.hub_trades_lib import (
attach_journal_mood_tags,
fetch_trades_for_trading_day,
journal_trade_match_key,
summarize_trades,
trading_day_from_dt,
trading_day_window_bounds,
)
class HubTradesLibTest(unittest.TestCase):
def test_trading_day_reset(self):
dt = datetime(2026, 6, 6, 7, 30, 0)
self.assertEqual(trading_day_from_dt(dt, 8), "2026-06-05")
dt2 = datetime(2026, 6, 6, 8, 0, 0)
self.assertEqual(trading_day_from_dt(dt2, 8), "2026-06-06")
def test_trading_day_window_bounds(self):
start, end = trading_day_window_bounds("2026-06-06", 8)
self.assertEqual(start, "2026-06-06 08:00:00")
self.assertEqual(end, "2026-06-07 07:59:59")
def test_fetch_and_summarize(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"ONDO/USDT",
"short",
"止损",
None,
-0.5,
None,
None,
"2026-06-06 10:00:00",
None,
"2026-06-06 09:00:00",
None,
"2026-06-06 10:00:00",
"趋势回调",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
stats = summarize_trades(rows)
self.assertEqual(stats["closed_count"], 1)
self.assertEqual(stats["loss_count"], 1)
self.assertAlmostEqual(stats["total_pnl_u"], -0.5)
conn.close()
def test_early_morning_belongs_prev_trading_day(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"BTC/USDT",
"long",
"止盈",
None,
1.2,
None,
None,
"2026-06-07 07:30:00",
None,
"2026-06-07 06:00:00",
None,
"2026-06-07 07:30:00",
"关键位",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
self.assertEqual(len(fetch_trades_for_trading_day(conn, "2026-06-07")), 0)
self.assertEqual(len(fetch_trades_for_trading_day(conn, "2026-06-06")), 1)
conn.close()
def test_reviewed_fields_preferred(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"ETH/USDT",
"long",
"止损",
"止盈",
-0.5,
2.0,
None,
"2026-06-06 09:00:00",
"2026-06-06 11:00:00",
"2026-06-06 08:00:00",
None,
"2026-06-06 11:00:00",
"趋势回调",
None,
None,
"trend",
"",
"2026-06-06 12:00:00",
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["result"], "止盈")
self.assertAlmostEqual(rows[0]["pnl_amount"], 2.0)
self.assertTrue(rows[0]["reviewed"])
conn.close()
def test_time_close_result_included(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE trade_records (
symbol TEXT, direction TEXT, result TEXT, reviewed_result TEXT,
pnl_amount REAL, reviewed_pnl_amount REAL, exchange_realized_pnl REAL,
closed_at TEXT, reviewed_closed_at TEXT, opened_at TEXT, reviewed_opened_at TEXT,
created_at TEXT, monitor_type TEXT, actual_rr REAL, planned_rr REAL,
trade_style TEXT, entry_reason TEXT, reviewed_at TEXT
)"""
)
conn.execute(
"INSERT INTO trade_records VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
(
"BTC/USDT",
"long",
"时间平仓",
None,
1.2,
None,
None,
"2026-06-06 12:00:00",
None,
"2026-06-06 08:00:00",
None,
"2026-06-06 12:00:00",
"趋势回调",
None,
None,
"trend",
"",
None,
),
)
conn.commit()
rows = fetch_trades_for_trading_day(conn, "2026-06-06")
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["result"], "时间平仓")
conn.close()
def test_attach_journal_mood_tags_marks_sick(self):
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
conn.execute(
"""CREATE TABLE journal_entries (
coin TEXT, open_datetime TEXT, close_datetime TEXT, mood_issues TEXT, created_at TEXT
)"""
)
conn.execute(
"INSERT INTO journal_entries VALUES (?,?,?,?,?)",
("ETH", "2026-07-06 21:51", "2026-07-07 00:00", "报复开仓,扛单", "2026-07-07 00:05"),
)
conn.commit()
trades = [
{
"id": 42,
"symbol": "ETH/USDT",
"opened_at": "2026-07-06 21:51:00",
"closed_at": "2026-07-07 00:00:00",
}
]
attach_journal_mood_tags(conn, trades, cutoff_s="2026-01-01 00:00:00")
self.assertTrue(trades[0]["journal_mood_sick"])
self.assertEqual(trades[0]["behavior_tag"], "sick")
self.assertEqual(trades[0]["journal_mood_issues"], ["报复开仓", "扛单"])
key = journal_trade_match_key("ETH/USDT", "2026-07-06 21:51:00", "2026-07-07 00:00:00")
self.assertEqual(key, ("ETH", "2026-07-06 21:51", "2026-07-07 00:00"))
conn.close()
if __name__ == "__main__":
unittest.main()
+115 -115
View File
@@ -1,115 +1,115 @@
"""档案交易复盘字段优先开仓类型持仓时长开平仓时间)。"""
from __future__ import annotations
import tempfile
import unittest
from datetime import datetime, timedelta
from pathlib import Path
from lib.hub.hub_symbol_archive_lib import init_db, load_symbol_trades, upsert_trades_cache
from lib.hub.hub_trades_lib import (
_normalize_archive_trade_row,
display_entry_type_label,
effective_entry_type,
effective_hold_minutes,
)
class TestHubTradesReviewFields(unittest.TestCase):
def test_display_entry_type_for_manual_monitor_review(self):
d = {
"monitor_type": "下单监控",
"entry_reason": "",
"reviewed_entry_reason": "突破回踩",
"reviewed_at": "2026-06-08 10:00:00",
}
self.assertEqual(display_entry_type_label(d), "突破回踩")
def test_effective_entry_type_prefers_reviewed(self):
d = {
"entry_reason": "突破回踩",
"reviewed_entry_reason": "趋势回调",
"monitor_type": "下单监控",
}
self.assertEqual(effective_entry_type(d), "趋势回调")
def test_effective_hold_minutes_prefers_reviewed(self):
d = {
"hold_minutes": 30,
"reviewed_hold_minutes": 95,
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_001_800_000,
}
self.assertEqual(effective_hold_minutes(d), 95)
def test_normalize_archive_trade_row_review_fields(self):
closed = (datetime.now() - timedelta(days=2)).strftime("%Y-%m-%d %H:%M:%S")
opened = (datetime.now() - timedelta(days=2, hours=2)).strftime("%Y-%m-%d %H:%M:%S")
row = _normalize_archive_trade_row(
{
"id": 9,
"symbol": "ONDO/USDT",
"direction": "short",
"result": "止损",
"reviewed_result": "手动平仓",
"pnl_amount": -2.5,
"reviewed_pnl_amount": -2.58,
"opened_at": opened,
"reviewed_opened_at": "2026-06-07 14:30:00",
"closed_at": closed,
"reviewed_closed_at": "2026-06-08 08:44:21",
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
"entry_reason": "突破回踩",
"reviewed_entry_reason": "趋势回调",
"hold_minutes": 30,
"reviewed_hold_minutes": 1080,
"monitor_type": "趋势回调",
"reviewed_at": closed,
},
exchange_key="gate",
)
self.assertIsNotNone(row)
assert row is not None
self.assertEqual(row["entry_type"], "趋势回调")
self.assertEqual(row["hold_minutes"], 1080)
self.assertEqual(row["opened_at"], "2026-06-07 14:30:00")
self.assertEqual(row["closed_at"], "2026-06-08 08:44:21")
self.assertTrue(row["reviewed"])
def test_archive_cache_enriches_review_display_fields(self):
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 3,
"symbol": "ONDO/USDT",
"direction": "short",
"result": "手动平仓",
"pnl_amount": -2.58,
"opened_at": "2026-06-07 14:30:00",
"closed_at": "2026-06-08 08:44:21",
"opened_at_ms": 1_781_000_000_000,
"closed_at_ms": 1_781_065_000_000,
"entry_type": "趋势回调",
"hold_minutes": 1080,
"hold_minutes_text": "18小时0分钟",
"reviewed": True,
}
],
db_path=db,
)
rows = load_symbol_trades("gate", "ONDO/USDT", db_path=db)
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["entry_type"], "趋势回调")
self.assertEqual(rows[0]["hold_minutes"], 1080)
self.assertTrue(rows[0]["opened_at"].startswith("2026-06-07"))
self.assertTrue(rows[0]["closed_at"].startswith("2026-06-08"))
if __name__ == "__main__":
unittest.main()
"""档案交易:复盘字段优先(开仓类型,持仓时长,开平仓时间)."""
from __future__ import annotations
import tempfile
import unittest
from datetime import datetime, timedelta
from pathlib import Path
from lib.hub.hub_symbol_archive_lib import init_db, load_symbol_trades, upsert_trades_cache
from lib.hub.hub_trades_lib import (
_normalize_archive_trade_row,
display_entry_type_label,
effective_entry_type,
effective_hold_minutes,
)
class TestHubTradesReviewFields(unittest.TestCase):
def test_display_entry_type_for_manual_monitor_review(self):
d = {
"monitor_type": "下单监控",
"entry_reason": "",
"reviewed_entry_reason": "突破回踩",
"reviewed_at": "2026-06-08 10:00:00",
}
self.assertEqual(display_entry_type_label(d), "突破回踩")
def test_effective_entry_type_prefers_reviewed(self):
d = {
"entry_reason": "突破回踩",
"reviewed_entry_reason": "趋势回调",
"monitor_type": "下单监控",
}
self.assertEqual(effective_entry_type(d), "趋势回调")
def test_effective_hold_minutes_prefers_reviewed(self):
d = {
"hold_minutes": 30,
"reviewed_hold_minutes": 95,
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_001_800_000,
}
self.assertEqual(effective_hold_minutes(d), 95)
def test_normalize_archive_trade_row_review_fields(self):
closed = (datetime.now() - timedelta(days=2)).strftime("%Y-%m-%d %H:%M:%S")
opened = (datetime.now() - timedelta(days=2, hours=2)).strftime("%Y-%m-%d %H:%M:%S")
row = _normalize_archive_trade_row(
{
"id": 9,
"symbol": "ONDO/USDT",
"direction": "short",
"result": "止损",
"reviewed_result": "手动平仓",
"pnl_amount": -2.5,
"reviewed_pnl_amount": -2.58,
"opened_at": opened,
"reviewed_opened_at": "2026-06-07 14:30:00",
"closed_at": closed,
"reviewed_closed_at": "2026-06-08 08:44:21",
"opened_at_ms": 1_700_000_000_000,
"closed_at_ms": 1_700_007_200_000,
"entry_reason": "突破回踩",
"reviewed_entry_reason": "趋势回调",
"hold_minutes": 30,
"reviewed_hold_minutes": 1080,
"monitor_type": "趋势回调",
"reviewed_at": closed,
},
exchange_key="gate",
)
self.assertIsNotNone(row)
assert row is not None
self.assertEqual(row["entry_type"], "趋势回调")
self.assertEqual(row["hold_minutes"], 1080)
self.assertEqual(row["opened_at"], "2026-06-07 14:30:00")
self.assertEqual(row["closed_at"], "2026-06-08 08:44:21")
self.assertTrue(row["reviewed"])
def test_archive_cache_enriches_review_display_fields(self):
with tempfile.TemporaryDirectory() as td:
db = Path(td) / "archive.db"
init_db(db)
upsert_trades_cache(
"gate",
[
{
"id": 3,
"symbol": "ONDO/USDT",
"direction": "short",
"result": "手动平仓",
"pnl_amount": -2.58,
"opened_at": "2026-06-07 14:30:00",
"closed_at": "2026-06-08 08:44:21",
"opened_at_ms": 1_781_000_000_000,
"closed_at_ms": 1_781_065_000_000,
"entry_type": "趋势回调",
"hold_minutes": 1080,
"hold_minutes_text": "18小时0分钟",
"reviewed": True,
}
],
db_path=db,
)
rows = load_symbol_trades("gate", "ONDO/USDT", db_path=db)
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["entry_type"], "趋势回调")
self.assertEqual(rows[0]["hold_minutes"], 1080)
self.assertTrue(rows[0]["opened_at"].startswith("2026-06-07"))
self.assertTrue(rows[0]["closed_at"].startswith("2026-06-08"))
if __name__ == "__main__":
unittest.main()
+184 -184
View File
@@ -1,184 +1,184 @@
from datetime import datetime
from unittest.mock import MagicMock
from lib.hub.hub_volume_rank_lib import (
CACHE_VERSION,
LIQUIDITY_RANK_CACHE_VERSION,
TOP_N_DEFAULT,
_exchange_rank_row_stale,
_okx_turnover_usdt,
_scores_from_binance,
_scores_from_gate,
build_usdt_swap_volume_ranks,
cache_needs_refresh,
format_volume_quote,
merge_exchange_rank,
rank_date_label,
resolve_daily_volume_rank,
)
def test_rank_date_label_after_reset():
# 2026-06-08 09:00 北京时间 → 昨日交易日 2026-06-07
dt = datetime(2026, 6, 8, 9, 0, 0)
assert rank_date_label(now=dt, reset_hour=8) == "2026-06-07"
def test_rank_date_label_before_reset():
# 2026-06-08 07:00 → 当前交易日仍算 2026-06-07昨日为 2026-06-06
dt = datetime(2026, 6, 8, 7, 0, 0)
assert rank_date_label(now=dt, reset_hour=8) == "2026-06-06"
def test_format_volume_quote():
assert format_volume_quote(1_500_000_000) == "1.50B"
assert format_volume_quote(2_300_000) == "2.30M"
assert format_volume_quote(4500) == "4.50K"
def test_okx_turnover_usdt():
qv = _okx_turnover_usdt({"volCcy24h": "100", "last": "50"})
assert qv == 5000.0
def test_cache_needs_refresh_and_merge():
cache = {"rank_date": "2026-06-05", "exchanges": {}}
assert cache_needs_refresh(cache, expected_rank_date="2026-06-07") is True
merged = merge_exchange_rank(
cache,
"binance",
{
"ok": True,
"rank_date": "2026-06-07",
"items": [{"rank": 1, "symbol": "BTC/USDT", "volume_quote": 1.0}],
"total_symbols": 100,
},
)
assert merged["exchanges"]["binance"]["items"][0]["symbol"] == "BTC/USDT"
assert merged["rank_date"] == "2026-06-07"
def test_stale_cache_version_forces_refresh():
cache = {"version": CACHE_VERSION - 1, "rank_date": "2026-06-07", "exchanges": {"okx": {"items": [{}]}}}
assert cache_needs_refresh(cache) is True
def test_short_item_list_is_stale():
items = [{"rank": i, "symbol": f"S{i}/USDT"} for i in range(1, 13)]
row = {"items": items, "total_symbols": 12}
assert _exchange_rank_row_stale(row) is True
full = {"items": items + [{"rank": i, "symbol": f"X{i}/USDT"} for i in range(13, TOP_N_DEFAULT + 1)], "total_symbols": 300}
assert _exchange_rank_row_stale(full) is False
def test_scores_from_binance_uses_fapi_lightweight_api():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "BTCUSDT", "quoteVolume": "9000000"},
{"symbol": "ETHUSDT", "quoteVolume": "5000000"},
]
scored = _scores_from_binance(ex)
assert scored[0][1] == "BTC"
assert scored[0][2] == 9000000.0
ex.fetch_tickers.assert_not_called()
def test_scores_from_binance_skips_fetch_tickers_on_api_error():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.side_effect = RuntimeError("network")
scored = _scores_from_binance(ex)
assert scored == []
ex.fetch_tickers.assert_not_called()
def test_scores_from_gate_uses_futures_tickers_api():
ex = MagicMock()
ex.id = "gateio"
ex.publicFuturesGetSettleTickers.return_value = [
{"contract": "BTC_USDT", "volume_24h_quote": "8000000"},
{"contract": "ETH_USDT", "volume_24h_quote": "4000000"},
]
scored = _scores_from_gate(ex)
assert scored[0][1] == "BTC"
ex.fetch_tickers.assert_not_called()
def test_scores_from_gate_skips_fetch_tickers_on_api_error():
ex = MagicMock()
ex.id = "gateio"
ex.publicFuturesGetSettleTickers.side_effect = RuntimeError("network")
scored = _scores_from_gate(ex)
assert scored == []
ex.fetch_tickers.assert_not_called()
def test_resolve_daily_volume_rank_caches_result():
cache = {"version": 0, "updated_at": 0.0, "ranks": {}, "total": 0}
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "BTCUSDT", "quoteVolume": "100"},
{"symbol": "ETHUSDT", "quoteVolume": "50"},
]
rank, total = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=1000.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank == 1
assert total == 2
assert cache["version"] == LIQUIDITY_RANK_CACHE_VERSION
calls = ex.fapiPublicGetTicker24hr.call_count
rank2, _ = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=1010.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank2 == 1
assert ex.fapiPublicGetTicker24hr.call_count == calls
def test_resolve_daily_volume_rank_keeps_stale_cache_when_refresh_empty():
cache = {
"version": LIQUIDITY_RANK_CACHE_VERSION,
"updated_at": 900.0,
"ranks": {"BTC": 1},
"total": 100,
}
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = []
rank, total = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=2000.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank == 1
assert total == 100
assert cache["updated_at"] == 900.0
ex.fetch_tickers.assert_not_called()
def test_build_usdt_swap_volume_ranks():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "SOLUSDT", "quoteVolume": "200"},
]
ranks, total = build_usdt_swap_volume_ranks(ex, lambda: None)
assert ranks["SOL"] == 1
assert total == 1
from datetime import datetime
from unittest.mock import MagicMock
from lib.hub.hub_volume_rank_lib import (
CACHE_VERSION,
LIQUIDITY_RANK_CACHE_VERSION,
TOP_N_DEFAULT,
_exchange_rank_row_stale,
_okx_turnover_usdt,
_scores_from_binance,
_scores_from_gate,
build_usdt_swap_volume_ranks,
cache_needs_refresh,
format_volume_quote,
merge_exchange_rank,
rank_date_label,
resolve_daily_volume_rank,
)
def test_rank_date_label_after_reset():
# 2026-06-08 09:00 北京时间 → 昨日交易日 2026-06-07
dt = datetime(2026, 6, 8, 9, 0, 0)
assert rank_date_label(now=dt, reset_hour=8) == "2026-06-07"
def test_rank_date_label_before_reset():
# 2026-06-08 07:00 → 当前交易日仍算 2026-06-07,昨日为 2026-06-06
dt = datetime(2026, 6, 8, 7, 0, 0)
assert rank_date_label(now=dt, reset_hour=8) == "2026-06-06"
def test_format_volume_quote():
assert format_volume_quote(1_500_000_000) == "1.50B"
assert format_volume_quote(2_300_000) == "2.30M"
assert format_volume_quote(4500) == "4.50K"
def test_okx_turnover_usdt():
qv = _okx_turnover_usdt({"volCcy24h": "100", "last": "50"})
assert qv == 5000.0
def test_cache_needs_refresh_and_merge():
cache = {"rank_date": "2026-06-05", "exchanges": {}}
assert cache_needs_refresh(cache, expected_rank_date="2026-06-07") is True
merged = merge_exchange_rank(
cache,
"binance",
{
"ok": True,
"rank_date": "2026-06-07",
"items": [{"rank": 1, "symbol": "BTC/USDT", "volume_quote": 1.0}],
"total_symbols": 100,
},
)
assert merged["exchanges"]["binance"]["items"][0]["symbol"] == "BTC/USDT"
assert merged["rank_date"] == "2026-06-07"
def test_stale_cache_version_forces_refresh():
cache = {"version": CACHE_VERSION - 1, "rank_date": "2026-06-07", "exchanges": {"okx": {"items": [{}]}}}
assert cache_needs_refresh(cache) is True
def test_short_item_list_is_stale():
items = [{"rank": i, "symbol": f"S{i}/USDT"} for i in range(1, 13)]
row = {"items": items, "total_symbols": 12}
assert _exchange_rank_row_stale(row) is True
full = {"items": items + [{"rank": i, "symbol": f"X{i}/USDT"} for i in range(13, TOP_N_DEFAULT + 1)], "total_symbols": 300}
assert _exchange_rank_row_stale(full) is False
def test_scores_from_binance_uses_fapi_lightweight_api():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "BTCUSDT", "quoteVolume": "9000000"},
{"symbol": "ETHUSDT", "quoteVolume": "5000000"},
]
scored = _scores_from_binance(ex)
assert scored[0][1] == "BTC"
assert scored[0][2] == 9000000.0
ex.fetch_tickers.assert_not_called()
def test_scores_from_binance_skips_fetch_tickers_on_api_error():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.side_effect = RuntimeError("network")
scored = _scores_from_binance(ex)
assert scored == []
ex.fetch_tickers.assert_not_called()
def test_scores_from_gate_uses_futures_tickers_api():
ex = MagicMock()
ex.id = "gateio"
ex.publicFuturesGetSettleTickers.return_value = [
{"contract": "BTC_USDT", "volume_24h_quote": "8000000"},
{"contract": "ETH_USDT", "volume_24h_quote": "4000000"},
]
scored = _scores_from_gate(ex)
assert scored[0][1] == "BTC"
ex.fetch_tickers.assert_not_called()
def test_scores_from_gate_skips_fetch_tickers_on_api_error():
ex = MagicMock()
ex.id = "gateio"
ex.publicFuturesGetSettleTickers.side_effect = RuntimeError("network")
scored = _scores_from_gate(ex)
assert scored == []
ex.fetch_tickers.assert_not_called()
def test_resolve_daily_volume_rank_caches_result():
cache = {"version": 0, "updated_at": 0.0, "ranks": {}, "total": 0}
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "BTCUSDT", "quoteVolume": "100"},
{"symbol": "ETHUSDT", "quoteVolume": "50"},
]
rank, total = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=1000.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank == 1
assert total == 2
assert cache["version"] == LIQUIDITY_RANK_CACHE_VERSION
calls = ex.fapiPublicGetTicker24hr.call_count
rank2, _ = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=1010.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank2 == 1
assert ex.fapiPublicGetTicker24hr.call_count == calls
def test_resolve_daily_volume_rank_keeps_stale_cache_when_refresh_empty():
cache = {
"version": LIQUIDITY_RANK_CACHE_VERSION,
"updated_at": 900.0,
"ranks": {"BTC": 1},
"total": 100,
}
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = []
rank, total = resolve_daily_volume_rank(
"BTC",
cache,
now_ts=2000.0,
ttl_sec=60.0,
exchange=ex,
ensure_markets_loaded=lambda: None,
)
assert rank == 1
assert total == 100
assert cache["updated_at"] == 900.0
ex.fetch_tickers.assert_not_called()
def test_build_usdt_swap_volume_ranks():
ex = MagicMock()
ex.id = "binance"
ex.fapiPublicGetTicker24hr.return_value = [
{"symbol": "SOLUSDT", "quoteVolume": "200"},
]
ranks, total = build_usdt_swap_volume_ranks(ex, lambda: None)
assert ranks["SOL"] == 1
assert total == 1
+1 -1
View File
@@ -1,4 +1,4 @@
"""instance_display_prefs_lib 与 env_file_lib 单元测试"""
"""instance_display_prefs_lib 与 env_file_lib 单元测试."""
from __future__ import annotations
import os
+1 -1
View File
@@ -1,4 +1,4 @@
"""instance_embed_context_lib 顶栏统计"""
"""instance_embed_context_lib 顶栏统计."""
from __future__ import annotations
import unittest
+1 -1
View File
@@ -1,4 +1,4 @@
"""instance_live_pnl_lib 单元测试"""
"""instance_live_pnl_lib 单元测试."""
from __future__ import annotations
import unittest
+1 -1
View File
@@ -1,4 +1,4 @@
"""instance_live_push_lib 单元测试"""
"""instance_live_push_lib 单元测试."""
from __future__ import annotations
import json
+1 -1
View File
@@ -1,4 +1,4 @@
"""instance_settings_lib 单元测试"""
"""instance_settings_lib 单元测试."""
from __future__ import annotations
import os
+156 -156
View File
@@ -1,156 +1,156 @@
"""journal_images_lib / journal_upload_api_lib 单元测试"""
import json
import os
import tempfile
import unittest
from io import BytesIO
from lib.instance.journal_images_lib import (
JOURNAL_UPLOAD_TFS,
collect_journal_slot_images,
enrich_journal_api_item,
images_json_dumps,
is_valid_preuploaded_journal_file,
journal_image_paths,
journal_upload_field_name,
normalize_journal_draft_id,
parse_images_json,
primary_journal_image,
save_journal_slot_uploads,
uploaded_screenshot_field_name,
)
from lib.instance.journal_upload_api_lib import handle_journal_upload_slot
class _FakeFile:
def __init__(self, filename: str, data: bytes):
self.filename = filename
self._data = data
def save(self, path: str) -> None:
with open(path, "wb") as f:
f.write(self._data)
class _FakeFiles:
def __init__(self, mapping):
self._mapping = mapping
def get(self, key):
return self._mapping.get(key)
class _FakeForm:
def __init__(self, mapping):
self._mapping = mapping
def get(self, key, default=None):
return self._mapping.get(key, default)
class _FakeRequest:
def __init__(self, form=None, files=None):
self.form = form
self.files = files
class JournalImagesLibTest(unittest.TestCase):
def test_field_names(self):
self.assertEqual(journal_upload_field_name("5m"), "screenshot_5m")
self.assertEqual(uploaded_screenshot_field_name("5m"), "uploaded_screenshot_5m")
def test_normalize_draft_id(self):
good = "a" * 32
self.assertEqual(normalize_journal_draft_id(good), good)
self.assertIsNone(normalize_journal_draft_id("bad"))
def test_save_slot_uploads_partial(self):
with tempfile.TemporaryDirectory() as tmp:
files = _FakeFiles(
{
"screenshot_5m": _FakeFile("a.png", b"png5"),
"screenshot_1h": _FakeFile("b.jpg", b"jpg1"),
}
)
saved = save_journal_slot_uploads(
files,
"abc123" + "0" * 26,
tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(len(saved), 2)
self.assertEqual(saved[0]["tf"], "5m")
self.assertTrue(os.path.isfile(os.path.join(tmp, saved[0]["file"])))
self.assertEqual(saved[1]["tf"], "1h")
def test_collect_preuploaded(self):
entry_id = "abc123" + "0" * 26
fname = f"journal_{entry_id}_5m.png"
with tempfile.TemporaryDirectory() as tmp:
with open(os.path.join(tmp, fname), "wb") as f:
f.write(b"x")
form = _FakeForm({uploaded_screenshot_field_name("5m"): fname})
saved = collect_journal_slot_images(
form,
_FakeFiles({}),
entry_id,
tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(saved, [{"tf": "5m", "file": fname}])
def test_is_valid_preuploaded_journal_file(self):
entry_id = "abc123" + "0" * 26
fname = f"journal_{entry_id}_5m.png"
self.assertTrue(is_valid_preuploaded_journal_file(fname, entry_id, "5m"))
self.assertFalse(is_valid_preuploaded_journal_file("../evil.png", entry_id, "5m"))
self.assertFalse(is_valid_preuploaded_journal_file(fname, "b" * 32, "5m"))
def test_parse_and_enrich(self):
raw = images_json_dumps([{"tf": "5m", "file": "journal_x_5m.png"}])
item = enrich_journal_api_item({"images_json": raw, "image": "legacy.png"})
self.assertEqual(len(item["images"]), 1)
self.assertEqual(item["images"][0]["tf"], "5m")
legacy = enrich_journal_api_item({"image": "only.png"})
self.assertEqual(legacy["images"][0]["file"], "only.png")
def test_journal_image_paths_dedupe(self):
with tempfile.TemporaryDirectory() as tmp:
path = os.path.join(tmp, "same.png")
with open(path, "wb") as f:
f.write(b"x")
row = {
"image": "same.png",
"images_json": json.dumps([{"tf": "5m", "file": "same.png"}]),
}
paths = journal_image_paths(row, tmp)
self.assertEqual(len(paths), 1)
def test_primary_journal_image(self):
self.assertEqual(
primary_journal_image([{"tf": "5m", "file": "a.png"}]),
"a.png",
)
self.assertIsNone(primary_journal_image([]))
def test_handle_journal_upload_slot(self):
entry_id = "abc123" + "0" * 26
with tempfile.TemporaryDirectory() as tmp:
req = _FakeRequest(
form=_FakeForm({"journal_draft_id": entry_id, "tf": "5m"}),
files=_FakeFiles({"file": _FakeFile("local.png", b"data")}),
)
payload, code = handle_journal_upload_slot(
req,
upload_folder=tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(code, 200)
self.assertTrue(payload["ok"])
self.assertEqual(payload["tf"], "5m")
self.assertTrue(os.path.isfile(os.path.join(tmp, payload["file"])))
if __name__ == "__main__":
unittest.main()
"""journal_images_lib / journal_upload_api_lib 单元测试."""
import json
import os
import tempfile
import unittest
from io import BytesIO
from lib.instance.journal_images_lib import (
JOURNAL_UPLOAD_TFS,
collect_journal_slot_images,
enrich_journal_api_item,
images_json_dumps,
is_valid_preuploaded_journal_file,
journal_image_paths,
journal_upload_field_name,
normalize_journal_draft_id,
parse_images_json,
primary_journal_image,
save_journal_slot_uploads,
uploaded_screenshot_field_name,
)
from lib.instance.journal_upload_api_lib import handle_journal_upload_slot
class _FakeFile:
def __init__(self, filename: str, data: bytes):
self.filename = filename
self._data = data
def save(self, path: str) -> None:
with open(path, "wb") as f:
f.write(self._data)
class _FakeFiles:
def __init__(self, mapping):
self._mapping = mapping
def get(self, key):
return self._mapping.get(key)
class _FakeForm:
def __init__(self, mapping):
self._mapping = mapping
def get(self, key, default=None):
return self._mapping.get(key, default)
class _FakeRequest:
def __init__(self, form=None, files=None):
self.form = form
self.files = files
class JournalImagesLibTest(unittest.TestCase):
def test_field_names(self):
self.assertEqual(journal_upload_field_name("5m"), "screenshot_5m")
self.assertEqual(uploaded_screenshot_field_name("5m"), "uploaded_screenshot_5m")
def test_normalize_draft_id(self):
good = "a" * 32
self.assertEqual(normalize_journal_draft_id(good), good)
self.assertIsNone(normalize_journal_draft_id("bad"))
def test_save_slot_uploads_partial(self):
with tempfile.TemporaryDirectory() as tmp:
files = _FakeFiles(
{
"screenshot_5m": _FakeFile("a.png", b"png5"),
"screenshot_1h": _FakeFile("b.jpg", b"jpg1"),
}
)
saved = save_journal_slot_uploads(
files,
"abc123" + "0" * 26,
tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(len(saved), 2)
self.assertEqual(saved[0]["tf"], "5m")
self.assertTrue(os.path.isfile(os.path.join(tmp, saved[0]["file"])))
self.assertEqual(saved[1]["tf"], "1h")
def test_collect_preuploaded(self):
entry_id = "abc123" + "0" * 26
fname = f"journal_{entry_id}_5m.png"
with tempfile.TemporaryDirectory() as tmp:
with open(os.path.join(tmp, fname), "wb") as f:
f.write(b"x")
form = _FakeForm({uploaded_screenshot_field_name("5m"): fname})
saved = collect_journal_slot_images(
form,
_FakeFiles({}),
entry_id,
tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(saved, [{"tf": "5m", "file": fname}])
def test_is_valid_preuploaded_journal_file(self):
entry_id = "abc123" + "0" * 26
fname = f"journal_{entry_id}_5m.png"
self.assertTrue(is_valid_preuploaded_journal_file(fname, entry_id, "5m"))
self.assertFalse(is_valid_preuploaded_journal_file("../evil.png", entry_id, "5m"))
self.assertFalse(is_valid_preuploaded_journal_file(fname, "b" * 32, "5m"))
def test_parse_and_enrich(self):
raw = images_json_dumps([{"tf": "5m", "file": "journal_x_5m.png"}])
item = enrich_journal_api_item({"images_json": raw, "image": "legacy.png"})
self.assertEqual(len(item["images"]), 1)
self.assertEqual(item["images"][0]["tf"], "5m")
legacy = enrich_journal_api_item({"image": "only.png"})
self.assertEqual(legacy["images"][0]["file"], "only.png")
def test_journal_image_paths_dedupe(self):
with tempfile.TemporaryDirectory() as tmp:
path = os.path.join(tmp, "same.png")
with open(path, "wb") as f:
f.write(b"x")
row = {
"image": "same.png",
"images_json": json.dumps([{"tf": "5m", "file": "same.png"}]),
}
paths = journal_image_paths(row, tmp)
self.assertEqual(len(paths), 1)
def test_primary_journal_image(self):
self.assertEqual(
primary_journal_image([{"tf": "5m", "file": "a.png"}]),
"a.png",
)
self.assertIsNone(primary_journal_image([]))
def test_handle_journal_upload_slot(self):
entry_id = "abc123" + "0" * 26
with tempfile.TemporaryDirectory() as tmp:
req = _FakeRequest(
form=_FakeForm({"journal_draft_id": entry_id, "tf": "5m"}),
files=_FakeFiles({"file": _FakeFile("local.png", b"data")}),
)
payload, code = handle_journal_upload_slot(
req,
upload_folder=tmp,
secure_filename_fn=lambda x: x,
)
self.assertEqual(code, 200)
self.assertTrue(payload["ok"])
self.assertEqual(payload["tf"], "5m")
self.assertTrue(os.path.isfile(os.path.join(tmp, payload["file"])))
if __name__ == "__main__":
unittest.main()
+75 -75
View File
@@ -1,75 +1,75 @@
"""key_auto_order_lib 单元测试"""
import unittest
from lib.key_monitor.key_auto_order_lib import (
check_monitor_type_add_allowed,
effective_entry_reason_options,
effective_stats_segment_defs,
load_key_auto_order_enabled,
)
from lib.trade.position_sizing_lib import MODE_FULL_MARGIN, MODE_RISK
FULL_OPTS = (
"趋势A",
"趋势B",
"趋势C",
"趋势D",
"趋势E",
"关键位箱体突破",
"关键位收敛突破",
"关键位斐波0.618",
"关键位斐波0.786",
"关键位假突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
"趋势回调",
"顺势加仓",
)
STATS_DEFS = (
("all", "全部", {}),
("key_box", "箱体", {}),
("key_trigger", "触价", {}),
)
class KeyAutoOrderLibTest(unittest.TestCase):
def test_load_default_false(self):
self.assertFalse(load_key_auto_order_enabled({"KEY_AUTO_ORDER_ENABLED": "false"}))
self.assertFalse(load_key_auto_order_enabled({}))
self.assertTrue(load_key_auto_order_enabled({"KEY_AUTO_ORDER_ENABLED": "true"}))
def test_entry_reason_off(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_RISK, False)
self.assertNotIn("关键位箱体突破", out)
self.assertNotIn("关键位回调触价开仓", out)
self.assertIn("顺势加仓", out)
def test_entry_reason_risk_on(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_RISK, True)
self.assertIn("关键位箱体突破", out)
self.assertIn("关键位回调触价开仓", out)
def test_entry_reason_full_margin_on(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_FULL_MARGIN, True)
self.assertNotIn("关键位箱体突破", out)
self.assertIn("关键位回调触价开仓", out)
def test_stats_segments_off(self):
segs = effective_stats_segment_defs(STATS_DEFS, MODE_RISK, False)
keys = {x[0] for x in segs}
self.assertIn("all", keys)
self.assertNotIn("key_box", keys)
def test_add_key_rs_always(self):
ok, _ = check_monitor_type_add_allowed("关键支撑阻力", MODE_RISK, False)
self.assertTrue(ok)
def test_add_key_trigger_off(self):
ok, msg = check_monitor_type_add_allowed("回调触价开仓", MODE_RISK, False)
self.assertFalse(ok)
self.assertIn("KEY_AUTO_ORDER_ENABLED", msg)
if __name__ == "__main__":
unittest.main()
"""key_auto_order_lib 单元测试."""
import unittest
from lib.key_monitor.key_auto_order_lib import (
check_monitor_type_add_allowed,
effective_entry_reason_options,
effective_stats_segment_defs,
load_key_auto_order_enabled,
)
from lib.trade.position_sizing_lib import MODE_FULL_MARGIN, MODE_RISK
FULL_OPTS = (
"趋势A",
"趋势B",
"趋势C",
"趋势D",
"趋势E",
"关键位箱体突破",
"关键位收敛突破",
"关键位斐波0.618",
"关键位斐波0.786",
"关键位假突破",
"关键位回调触价开仓",
"关键位突破触价开仓",
"趋势回调",
"顺势加仓",
)
STATS_DEFS = (
("all", "全部", {}),
("key_box", "箱体", {}),
("key_trigger", "触价", {}),
)
class KeyAutoOrderLibTest(unittest.TestCase):
def test_load_default_false(self):
self.assertFalse(load_key_auto_order_enabled({"KEY_AUTO_ORDER_ENABLED": "false"}))
self.assertFalse(load_key_auto_order_enabled({}))
self.assertTrue(load_key_auto_order_enabled({"KEY_AUTO_ORDER_ENABLED": "true"}))
def test_entry_reason_off(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_RISK, False)
self.assertNotIn("关键位箱体突破", out)
self.assertNotIn("关键位回调触价开仓", out)
self.assertIn("顺势加仓", out)
def test_entry_reason_risk_on(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_RISK, True)
self.assertIn("关键位箱体突破", out)
self.assertIn("关键位回调触价开仓", out)
def test_entry_reason_full_margin_on(self):
out = effective_entry_reason_options(FULL_OPTS, MODE_FULL_MARGIN, True)
self.assertNotIn("关键位箱体突破", out)
self.assertIn("关键位回调触价开仓", out)
def test_stats_segments_off(self):
segs = effective_stats_segment_defs(STATS_DEFS, MODE_RISK, False)
keys = {x[0] for x in segs}
self.assertIn("all", keys)
self.assertNotIn("key_box", keys)
def test_add_key_rs_always(self):
ok, _ = check_monitor_type_add_allowed("关键支撑阻力", MODE_RISK, False)
self.assertTrue(ok)
def test_add_key_trigger_off(self):
ok, msg = check_monitor_type_add_allowed("回调触价开仓", MODE_RISK, False)
self.assertFalse(ok)
self.assertIn("KEY_AUTO_ORDER_ENABLED", msg)
if __name__ == "__main__":
unittest.main()
+86 -86
View File
@@ -1,86 +1,86 @@
"""阻力/支撑提醒占位与间隔防重复推送"""
from __future__ import annotations
import sqlite3
import unittest
from datetime import datetime, timedelta
from lib.key_monitor.key_monitor_lib import (
claim_rs_level_notify,
notify_interval_elapsed,
run_rs_level_alert_tick,
)
def _row(**kwargs):
base = {
"upper": 2.174,
"lower": 1.694,
"notification_count": 0,
"max_notify": 3,
"notify_interval_min": 5,
"direction": "watch",
"last_notified_at": None,
"last_rs_bar_ts": None,
}
base.update(kwargs)
return base
class TestRsLevelAlertClaim(unittest.TestCase):
def setUp(self):
self.conn = sqlite3.connect(":memory:")
self.conn.execute(
"CREATE TABLE key_monitors ("
"id INTEGER PRIMARY KEY, notification_count INTEGER DEFAULT 0, "
"direction TEXT, last_notified_at TEXT, last_rs_bar_ts INTEGER)"
)
self.conn.execute(
"INSERT INTO key_monitors (id, notification_count, direction) VALUES (1, 0, 'watch')"
)
self.conn.commit()
def test_claim_advances_once_per_index(self):
ok1 = claim_rs_level_notify(
self.conn, 1, 1, "long", "2026-06-02 00:25:00", 1000, prior_count=0
)
self.conn.commit()
self.assertTrue(ok1)
ok_dup = claim_rs_level_notify(
self.conn, 1, 1, "long", "2026-06-02 00:25:03", 1000, prior_count=0
)
self.assertFalse(ok_dup)
ok2 = claim_rs_level_notify(
self.conn, 1, 2, "long", "2026-06-02 00:30:00", 1000, prior_count=1
)
self.conn.commit()
self.assertTrue(ok2)
row = self.conn.execute(
"SELECT notification_count FROM key_monitors WHERE id=1"
).fetchone()
self.assertEqual(row[0], 2)
def test_second_push_requires_interval(self):
now = datetime(2026, 6, 2, 0, 26, 0)
row = _row(
notification_count=1,
direction="long",
last_notified_at="2026-06-02 00:25:00",
)
tick = run_rs_level_alert_tick(row, 2.18, 1000, now, default_max_notify=3, default_interval_min=5)
self.assertIsNone(tick)
later = datetime(2026, 6, 2, 0, 30, 1)
tick2 = run_rs_level_alert_tick(
row, 2.18, 1000, later, default_max_notify=3, default_interval_min=5
)
self.assertIsNotNone(tick2)
self.assertEqual(tick2["notify_index"], 2)
self.assertEqual(tick2["prior_count"], 1)
def test_notify_interval_invalid_timestamp_does_not_spam(self):
now = datetime(2026, 6, 2, 1, 0, 0)
self.assertFalse(notify_interval_elapsed("not-a-date", 5, now))
if __name__ == "__main__":
unittest.main()
"""阻力/支撑提醒:占位与间隔防重复推送."""
from __future__ import annotations
import sqlite3
import unittest
from datetime import datetime, timedelta
from lib.key_monitor.key_monitor_lib import (
claim_rs_level_notify,
notify_interval_elapsed,
run_rs_level_alert_tick,
)
def _row(**kwargs):
base = {
"upper": 2.174,
"lower": 1.694,
"notification_count": 0,
"max_notify": 3,
"notify_interval_min": 5,
"direction": "watch",
"last_notified_at": None,
"last_rs_bar_ts": None,
}
base.update(kwargs)
return base
class TestRsLevelAlertClaim(unittest.TestCase):
def setUp(self):
self.conn = sqlite3.connect(":memory:")
self.conn.execute(
"CREATE TABLE key_monitors ("
"id INTEGER PRIMARY KEY, notification_count INTEGER DEFAULT 0, "
"direction TEXT, last_notified_at TEXT, last_rs_bar_ts INTEGER)"
)
self.conn.execute(
"INSERT INTO key_monitors (id, notification_count, direction) VALUES (1, 0, 'watch')"
)
self.conn.commit()
def test_claim_advances_once_per_index(self):
ok1 = claim_rs_level_notify(
self.conn, 1, 1, "long", "2026-06-02 00:25:00", 1000, prior_count=0
)
self.conn.commit()
self.assertTrue(ok1)
ok_dup = claim_rs_level_notify(
self.conn, 1, 1, "long", "2026-06-02 00:25:03", 1000, prior_count=0
)
self.assertFalse(ok_dup)
ok2 = claim_rs_level_notify(
self.conn, 1, 2, "long", "2026-06-02 00:30:00", 1000, prior_count=1
)
self.conn.commit()
self.assertTrue(ok2)
row = self.conn.execute(
"SELECT notification_count FROM key_monitors WHERE id=1"
).fetchone()
self.assertEqual(row[0], 2)
def test_second_push_requires_interval(self):
now = datetime(2026, 6, 2, 0, 26, 0)
row = _row(
notification_count=1,
direction="long",
last_notified_at="2026-06-02 00:25:00",
)
tick = run_rs_level_alert_tick(row, 2.18, 1000, now, default_max_notify=3, default_interval_min=5)
self.assertIsNone(tick)
later = datetime(2026, 6, 2, 0, 30, 1)
tick2 = run_rs_level_alert_tick(
row, 2.18, 1000, later, default_max_notify=3, default_interval_min=5
)
self.assertIsNotNone(tick2)
self.assertEqual(tick2["notify_index"], 2)
self.assertEqual(tick2["prior_count"], 1)
def test_notify_interval_invalid_timestamp_does_not_spam(self):
now = datetime(2026, 6, 2, 1, 0, 0)
self.assertFalse(notify_interval_elapsed("not-a-date", 5, now))
if __name__ == "__main__":
unittest.main()
+2 -2
View File
@@ -1,4 +1,4 @@
"""预估盈亏比前端 manual_order_rr_preview.js公式与后端 calc_rr_ratio 口径一致"""
"""预估盈亏比(前端 manual_order_rr_preview.js)公式与后端 calc_rr_ratio 口径一致."""
def _calc_rr(direction: str, entry: float, sl: float, tp: float):
@@ -60,7 +60,7 @@ def _full_margin_risk_u(available: float, buffer: float, leverage: int, directio
def test_full_margin_risk_short_hype():
# 可用约 23.06U × 0.9 缓冲 × 5x入场 62.5止损 63.6
# 可用约 23.06U × 0.9 缓冲 × 5x,入场 62.5,止损 63.6
risk = _full_margin_risk_u(23.06, 0.9, 5, "short", 62.5, 63.6)
assert risk is not None
assert 1.5 <= risk <= 2.5
+1 -1
View File
@@ -1,4 +1,4 @@
"""OKX 持仓指标解析未实现盈亏须支持负数"""
"""OKX 持仓指标解析:未实现盈亏须支持负数."""
from __future__ import annotations
import unittest
+1 -1
View File
@@ -1,4 +1,4 @@
"""期权定价单测"""
"""期权定价单测."""
from lib.options.options_pricing_lib import (
calc_order_size,
premium_per_sheet,
+34 -34
View File
@@ -1,34 +1,34 @@
"""全仓 / 以损定仓 风险展示文案"""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.trade.position_sizing_lib import ( # noqa: E402
format_risk_display_text,
risk_percent_for_storage,
)
class TestPositionSizingRiskDisplay(unittest.TestCase):
def test_full_margin_shows_amount_only(self):
self.assertEqual(
format_risk_display_text("full_margin", 1.0, 2.58, decimals=2),
"2.58U",
)
self.assertIsNone(risk_percent_for_storage("full_margin", 1.0))
def test_risk_mode_shows_percent_and_amount(self):
self.assertEqual(
format_risk_display_text("risk", 2.0, 10.5, decimals=2),
"2%≈10.5U",
)
self.assertEqual(risk_percent_for_storage("risk", 2.0), 2.0)
if __name__ == "__main__":
unittest.main()
"""全仓 / 以损定仓 风险展示文案."""
from __future__ import annotations
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.trade.position_sizing_lib import ( # noqa: E402
format_risk_display_text,
risk_percent_for_storage,
)
class TestPositionSizingRiskDisplay(unittest.TestCase):
def test_full_margin_shows_amount_only(self):
self.assertEqual(
format_risk_display_text("full_margin", 1.0, 2.58, decimals=2),
"2.58U",
)
self.assertIsNone(risk_percent_for_storage("full_margin", 1.0))
def test_risk_mode_shows_percent_and_amount(self):
self.assertEqual(
format_risk_display_text("risk", 2.0, 10.5, decimals=2),
"2%≈10.5U",
)
self.assertEqual(risk_percent_for_storage("risk", 2.0), 2.0)
if __name__ == "__main__":
unittest.main()
+41 -41
View File
@@ -1,41 +1,41 @@
"""deploy/sanitize_hub_settings.py 单元测试"""
from __future__ import annotations
import json
import sys
from pathlib import Path
REPO = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO / "deploy"))
from sanitize_hub_settings import sanitize_settings # noqa: E402
def test_drops_gate_bot_and_keeps_gate():
raw = {
"exchanges": [
{"id": "0", "key": "binance", "name": "币安", "agent_url": "http://127.0.0.1:15200"},
{"id": "3", "key": "gate_bot", "name": "Gate bot", "agent_url": "http://127.0.0.1:15203"},
{"id": "2", "key": "gate", "name": "Gate", "flask_url": "http://127.0.0.1:5000"},
]
}
cleaned, removed = sanitize_settings(raw)
keys = [x["key"] for x in cleaned["exchanges"]]
assert keys == ["binance", "gate"]
assert len(removed) == 1
def test_drops_port_5002_legacy():
raw = {
"exchanges": [
{
"id": "3",
"key": "legacy",
"name": "crypto_monitor_gate_bot",
"flask_url": "http://127.0.0.1:5002",
},
]
}
cleaned, removed = sanitize_settings(raw)
assert cleaned["exchanges"] == []
assert removed
"""deploy/sanitize_hub_settings.py 单元测试."""
from __future__ import annotations
import json
import sys
from pathlib import Path
REPO = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO / "deploy"))
from sanitize_hub_settings import sanitize_settings # noqa: E402
def test_drops_gate_bot_and_keeps_gate():
raw = {
"exchanges": [
{"id": "0", "key": "binance", "name": "币安", "agent_url": "http://127.0.0.1:15200"},
{"id": "3", "key": "gate_bot", "name": "Gate bot", "agent_url": "http://127.0.0.1:15203"},
{"id": "2", "key": "gate", "name": "Gate", "flask_url": "http://127.0.0.1:5000"},
]
}
cleaned, removed = sanitize_settings(raw)
keys = [x["key"] for x in cleaned["exchanges"]]
assert keys == ["binance", "gate"]
assert len(removed) == 1
def test_drops_port_5002_legacy():
raw = {
"exchanges": [
{
"id": "3",
"key": "legacy",
"name": "crypto_monitor_gate_bot",
"flask_url": "http://127.0.0.1:5002",
},
]
}
cleaned, removed = sanitize_settings(raw)
assert cleaned["exchanges"] == []
assert removed
+1 -1
View File
@@ -1,4 +1,4 @@
"""shared_env_libAI 字段与四文件同步"""
"""shared_env_lib:AI 字段与四文件同步."""
from __future__ import annotations
import os
+46 -46
View File
@@ -1,46 +1,46 @@
"""strategy_roll_ui_lib 单元测试"""
from __future__ import annotations
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
import lib.strategy.strategy_roll_ui_lib as roll_ui
def test_compute_roll_chain_metrics_short():
group = {
"id": 1,
"direction": "short",
"initial_take_profit": 60.0,
}
legs = [
{"id": 10, "leg_index": 1, "amount": 3.0, "fill_price": 65.0, "status": "filled"},
{"id": 11, "leg_index": 2, "amount": 5.0, "fill_price": 64.0, "status": "filled"},
]
per_leg, group_metrics = roll_ui.compute_roll_chain_metrics(
group,
legs,
qty_live=8.0,
entry_live=63.5,
monitor={"trigger_price": 66.0, "order_amount": 3.0},
)
assert per_leg[10]["avg_entry_after"] is not None
assert per_leg[11]["avg_entry_after"] is not None
assert group_metrics["reward_at_tp_usdt"] is not None
assert group_metrics["initial_qty"] == 3.0
assert group_metrics["current_qty"] == 8.0
assert per_leg[11]["reward_at_tp_usdt"] >= per_leg[10]["reward_at_tp_usdt"]
def test_infer_initial_position_from_live():
legs = [{"amount": 2.0, "fill_price": 64.0, "status": "filled"}]
q0, e0 = roll_ui.infer_initial_position(5.0, 63.0, legs)
assert q0 == 3.0
assert abs(e0 - 62.3333333333) < 0.001
def test_reward_at_tp_long():
assert roll_ui.reward_at_tp_usdt("long", 100.0, 110.0, 2.0) == 20.0
"""strategy_roll_ui_lib 单元测试."""
from __future__ import annotations
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
import lib.strategy.strategy_roll_ui_lib as roll_ui
def test_compute_roll_chain_metrics_short():
group = {
"id": 1,
"direction": "short",
"initial_take_profit": 60.0,
}
legs = [
{"id": 10, "leg_index": 1, "amount": 3.0, "fill_price": 65.0, "status": "filled"},
{"id": 11, "leg_index": 2, "amount": 5.0, "fill_price": 64.0, "status": "filled"},
]
per_leg, group_metrics = roll_ui.compute_roll_chain_metrics(
group,
legs,
qty_live=8.0,
entry_live=63.5,
monitor={"trigger_price": 66.0, "order_amount": 3.0},
)
assert per_leg[10]["avg_entry_after"] is not None
assert per_leg[11]["avg_entry_after"] is not None
assert group_metrics["reward_at_tp_usdt"] is not None
assert group_metrics["initial_qty"] == 3.0
assert group_metrics["current_qty"] == 8.0
assert per_leg[11]["reward_at_tp_usdt"] >= per_leg[10]["reward_at_tp_usdt"]
def test_infer_initial_position_from_live():
legs = [{"amount": 2.0, "fill_price": 64.0, "status": "filled"}]
q0, e0 = roll_ui.infer_initial_position(5.0, 63.0, legs)
assert q0 == 3.0
assert abs(e0 - 62.3333333333) < 0.001
def test_reward_at_tp_long():
assert roll_ui.reward_at_tp_usdt("long", 100.0, 110.0, 2.0) == 20.0
+183 -183
View File
@@ -1,183 +1,183 @@
"""策略快照同一计划同结果不重复写入"""
from __future__ import annotations
import json
import sqlite3
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_snapshot_lib import ( # noqa: E402
STRATEGY_TREND,
dedupe_strategy_snapshots,
init_strategy_snapshot_table,
list_strategy_snapshots,
save_trend_plan_snapshot,
)
def _mem_conn() -> sqlite3.Connection:
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
init_strategy_snapshot_table(conn)
return conn
def test_save_trend_plan_snapshot_skips_duplicate_result():
conn = _mem_conn()
plan = {
"id": 42,
"symbol": "ONDO/USDT",
"exchange_symbol": "ONDO/USDT:USDT",
"direction": "short",
"status": "active",
"opened_at": "2026-06-08 08:00:00",
"legs_done": 4,
"dca_legs": 4,
"first_order_done": 1,
"grid_prices_json": "[]",
"leg_amounts_json": "[]",
}
cfg = {"app_module": type("M", (), {"app_now_str": staticmethod(lambda: "2026-06-08 08:41:00")})()}
save_trend_plan_snapshot(cfg, conn, plan, result_label="止损", pnl_amount=-2.3)
save_trend_plan_snapshot(cfg, conn, plan, result_label="止损", pnl_amount=-2.4)
conn.commit()
rows = conn.execute(
"SELECT COUNT(*) AS c FROM strategy_trade_snapshots WHERE source_id=? AND result_label=?",
(42, "止损"),
).fetchone()
assert int(rows["c"]) == 1
def test_dedupe_strategy_snapshots_handles_many_duplicates():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id in range(1, 46):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
99,
"ONDO/USDT",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
-2.2,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 44
row = conn.execute(
"SELECT COUNT(*) AS c FROM strategy_trade_snapshots WHERE source_id=?",
(99,),
).fetchone()
assert int(row["c"]) == 1
def test_dedupe_strategy_snapshots_keeps_latest_id():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id, pnl in ((1, -2.23), (2, -2.31), (3, -2.38)):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
5,
"ONDO/USDT",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
pnl,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 2
row = conn.execute(
"SELECT id, pnl_amount FROM strategy_trade_snapshots WHERE source_id=?",
(5,),
).fetchone()
assert int(row["id"]) == 3
assert abs(float(row["pnl_amount"]) - (-2.38)) < 1e-6
def test_list_strategy_snapshots_hides_duplicate_keys():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT", "dca_levels": []}, ensure_ascii=False)
for snap_id in (10, 11, 12):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction, result_label,
snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
7,
"ONDO/USDT",
"short",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
-2.2,
),
)
conn.commit()
rows = list_strategy_snapshots(conn, limit=50)
stop_rows = [r for r in rows if int(r.get("source_id") or 0) == 7]
assert len(stop_rows) == 1
assert int(stop_rows[0]["id"]) == 12
def test_dedupe_keeps_manual_over_stop_loss():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id, label in ((10, "止损"), (11, "手动平仓")):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
7,
"ONDO/USDT",
label,
payload,
"2026-06-08 08:44:00",
"2026-06-08 08:44:00",
-2.23,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 1
row = conn.execute(
"SELECT result_label FROM strategy_trade_snapshots WHERE source_id=?",
(7,),
).fetchone()
assert row["result_label"] == "手动平仓"
if __name__ == "__main__":
test_save_trend_plan_snapshot_skips_duplicate_result()
test_dedupe_strategy_snapshots_handles_many_duplicates()
test_dedupe_strategy_snapshots_keeps_latest_id()
test_list_strategy_snapshots_hides_duplicate_keys()
test_dedupe_keeps_manual_over_stop_loss()
print("all ok")
"""策略快照:同一计划同结果不重复写入."""
from __future__ import annotations
import json
import sqlite3
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_snapshot_lib import ( # noqa: E402
STRATEGY_TREND,
dedupe_strategy_snapshots,
init_strategy_snapshot_table,
list_strategy_snapshots,
save_trend_plan_snapshot,
)
def _mem_conn() -> sqlite3.Connection:
conn = sqlite3.connect(":memory:")
conn.row_factory = sqlite3.Row
init_strategy_snapshot_table(conn)
return conn
def test_save_trend_plan_snapshot_skips_duplicate_result():
conn = _mem_conn()
plan = {
"id": 42,
"symbol": "ONDO/USDT",
"exchange_symbol": "ONDO/USDT:USDT",
"direction": "short",
"status": "active",
"opened_at": "2026-06-08 08:00:00",
"legs_done": 4,
"dca_legs": 4,
"first_order_done": 1,
"grid_prices_json": "[]",
"leg_amounts_json": "[]",
}
cfg = {"app_module": type("M", (), {"app_now_str": staticmethod(lambda: "2026-06-08 08:41:00")})()}
save_trend_plan_snapshot(cfg, conn, plan, result_label="止损", pnl_amount=-2.3)
save_trend_plan_snapshot(cfg, conn, plan, result_label="止损", pnl_amount=-2.4)
conn.commit()
rows = conn.execute(
"SELECT COUNT(*) AS c FROM strategy_trade_snapshots WHERE source_id=? AND result_label=?",
(42, "止损"),
).fetchone()
assert int(rows["c"]) == 1
def test_dedupe_strategy_snapshots_handles_many_duplicates():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id in range(1, 46):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
99,
"ONDO/USDT",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
-2.2,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 44
row = conn.execute(
"SELECT COUNT(*) AS c FROM strategy_trade_snapshots WHERE source_id=?",
(99,),
).fetchone()
assert int(row["c"]) == 1
def test_dedupe_strategy_snapshots_keeps_latest_id():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id, pnl in ((1, -2.23), (2, -2.31), (3, -2.38)):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
5,
"ONDO/USDT",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
pnl,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 2
row = conn.execute(
"SELECT id, pnl_amount FROM strategy_trade_snapshots WHERE source_id=?",
(5,),
).fetchone()
assert int(row["id"]) == 3
assert abs(float(row["pnl_amount"]) - (-2.38)) < 1e-6
def test_list_strategy_snapshots_hides_duplicate_keys():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT", "dca_levels": []}, ensure_ascii=False)
for snap_id in (10, 11, 12):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, direction, result_label,
snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
7,
"ONDO/USDT",
"short",
"止损",
payload,
"2026-06-08 08:41:00",
"2026-06-08 08:41:00",
-2.2,
),
)
conn.commit()
rows = list_strategy_snapshots(conn, limit=50)
stop_rows = [r for r in rows if int(r.get("source_id") or 0) == 7]
assert len(stop_rows) == 1
assert int(stop_rows[0]["id"]) == 12
def test_dedupe_keeps_manual_over_stop_loss():
conn = _mem_conn()
payload = json.dumps({"symbol": "ONDO/USDT"}, ensure_ascii=False)
for snap_id, label in ((10, "止损"), (11, "手动平仓")):
conn.execute(
"""INSERT INTO strategy_trade_snapshots (
id, strategy_type, source_id, symbol, result_label, snapshot_json, closed_at, created_at, pnl_amount
) VALUES (?,?,?,?,?,?,?,?,?)""",
(
snap_id,
STRATEGY_TREND,
7,
"ONDO/USDT",
label,
payload,
"2026-06-08 08:44:00",
"2026-06-08 08:44:00",
-2.23,
),
)
conn.commit()
removed = dedupe_strategy_snapshots(conn)
conn.commit()
assert removed == 1
row = conn.execute(
"SELECT result_label FROM strategy_trade_snapshots WHERE source_id=?",
(7,),
).fetchone()
assert row["result_label"] == "手动平仓"
if __name__ == "__main__":
test_save_trend_plan_snapshot_skips_duplicate_result()
test_dedupe_strategy_snapshots_handles_many_duplicates()
test_dedupe_strategy_snapshots_keeps_latest_id()
test_list_strategy_snapshots_hides_duplicate_keys()
test_dedupe_keeps_manual_over_stop_loss()
print("all ok")
+90 -90
View File
@@ -1,90 +1,90 @@
"""账户方向 / 币种白名单 env 策略"""
from lib.trade.trade_policy_lib import (
assert_direction_allowed,
assert_symbol_allowed,
assert_trade_policy_open,
load_trade_policy,
parse_symbol_whitelist,
symbol_base_coin,
trade_policy_badge_parts,
)
def test_default_policy_unrestricted():
p = load_trade_policy({})
assert not p.direction_restrict_enabled
assert not p.symbol_restrict_enabled
assert p.allows_long and p.allows_short
def test_long_only_blocks_short():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "true",
"TRADE_DIRECTION": "long_only",
}
)
ok, msg = assert_direction_allowed(p, "short")
assert not ok
assert "仅做多" in msg
ok2, _ = assert_direction_allowed(p, "long")
assert ok2
def test_symbol_whitelist_btc_eth():
p = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ok, _ = assert_symbol_allowed(p, "BTC/USDT")
assert ok
ok2, msg = assert_symbol_allowed(p, "SOL")
assert not ok2
assert "SOL" in msg
def test_symbol_whitelist_without_list_disables_restrict():
p = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "",
}
)
assert not p.symbol_restrict_enabled
def test_combined_open_validation():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "1",
"TRADE_DIRECTION": "",
"TRADE_SYMBOL_RESTRICT_ENABLED": "yes",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ok, _ = assert_trade_policy_open(p, "ETH", "long")
assert ok
ok2, msg = assert_trade_policy_open(p, "ETH", "short")
assert not ok2
ok3, msg3 = assert_trade_policy_open(p, "BNB", "long")
assert not ok3
assert "BNB" in msg3
def test_parse_whitelist_and_base_coin():
assert parse_symbol_whitelist("btc, eth") == ("BTC", "ETH")
assert symbol_base_coin("btc/usdt:usdt") == "BTC"
def test_badge_parts():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "true",
"TRADE_DIRECTION": "long_only",
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
assert trade_policy_badge_parts(p) == ("仅多", "BTC/ETH")
"""账户方向 / 币种白名单 env 策略."""
from lib.trade.trade_policy_lib import (
assert_direction_allowed,
assert_symbol_allowed,
assert_trade_policy_open,
load_trade_policy,
parse_symbol_whitelist,
symbol_base_coin,
trade_policy_badge_parts,
)
def test_default_policy_unrestricted():
p = load_trade_policy({})
assert not p.direction_restrict_enabled
assert not p.symbol_restrict_enabled
assert p.allows_long and p.allows_short
def test_long_only_blocks_short():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "true",
"TRADE_DIRECTION": "long_only",
}
)
ok, msg = assert_direction_allowed(p, "short")
assert not ok
assert "仅做多" in msg
ok2, _ = assert_direction_allowed(p, "long")
assert ok2
def test_symbol_whitelist_btc_eth():
p = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ok, _ = assert_symbol_allowed(p, "BTC/USDT")
assert ok
ok2, msg = assert_symbol_allowed(p, "SOL")
assert not ok2
assert "SOL" in msg
def test_symbol_whitelist_without_list_disables_restrict():
p = load_trade_policy(
{
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "",
}
)
assert not p.symbol_restrict_enabled
def test_combined_open_validation():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "1",
"TRADE_DIRECTION": "",
"TRADE_SYMBOL_RESTRICT_ENABLED": "yes",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
ok, _ = assert_trade_policy_open(p, "ETH", "long")
assert ok
ok2, msg = assert_trade_policy_open(p, "ETH", "short")
assert not ok2
ok3, msg3 = assert_trade_policy_open(p, "BNB", "long")
assert not ok3
assert "BNB" in msg3
def test_parse_whitelist_and_base_coin():
assert parse_symbol_whitelist("btc, eth") == ("BTC", "ETH")
assert symbol_base_coin("btc/usdt:usdt") == "BTC"
def test_badge_parts():
p = load_trade_policy(
{
"TRADE_DIRECTION_RESTRICT_ENABLED": "true",
"TRADE_DIRECTION": "long_only",
"TRADE_SYMBOL_RESTRICT_ENABLED": "true",
"TRADE_SYMBOL_WHITELIST": "BTC,ETH",
}
)
assert trade_policy_badge_parts(p) == ("仅多", "BTC/ETH")
+30 -30
View File
@@ -1,30 +1,30 @@
from lib.trade.trade_result_lib import normalize_result_with_pnl, normalize_display_result, is_winning_pnl
def test_stop_loss_with_profit_becomes_trailing_tp():
assert normalize_result_with_pnl("止损", 4.33) == "移动止盈"
def test_manual_close_unchanged_even_with_profit():
assert normalize_result_with_pnl("手动平仓", 10) == "手动平仓"
def test_stop_loss_with_loss_unchanged():
assert normalize_result_with_pnl("止损", -2.5) == "止损"
def test_take_profit_unchanged():
assert normalize_result_with_pnl("止盈", 5) == "止盈"
def test_external_close_becomes_manual_close():
assert normalize_display_result("外部平仓") == "手动平仓"
assert normalize_result_with_pnl("外部平仓", 2.5) == "手动平仓"
assert normalize_result_with_pnl("外部平仓自动同步", -1) == "手动平仓"
def test_winning_pnl_positive_only():
assert is_winning_pnl(2.96) is True
assert is_winning_pnl(0) is False
assert is_winning_pnl(-1.05) is False
assert is_winning_pnl(None) is False
from lib.trade.trade_result_lib import normalize_result_with_pnl, normalize_display_result, is_winning_pnl
def test_stop_loss_with_profit_becomes_trailing_tp():
assert normalize_result_with_pnl("止损", 4.33) == "移动止盈"
def test_manual_close_unchanged_even_with_profit():
assert normalize_result_with_pnl("手动平仓", 10) == "手动平仓"
def test_stop_loss_with_loss_unchanged():
assert normalize_result_with_pnl("止损", -2.5) == "止损"
def test_take_profit_unchanged():
assert normalize_result_with_pnl("止盈", 5) == "止盈"
def test_external_close_becomes_manual_close():
assert normalize_display_result("外部平仓") == "手动平仓"
assert normalize_result_with_pnl("外部平仓", 2.5) == "手动平仓"
assert normalize_result_with_pnl("外部平仓(自动同步)", -1) == "手动平仓"
def test_winning_pnl_positive_only():
assert is_winning_pnl(2.96) is True
assert is_winning_pnl(0) is False
assert is_winning_pnl(-1.05) is False
assert is_winning_pnl(None) is False
+26 -26
View File
@@ -1,26 +1,26 @@
"""trade_result_lib过滤「错过」记录"""
import unittest
from lib.trade.trade_result_lib import (
filter_trade_records_excluding_miss,
is_miss_trade_result,
)
class TradeResultMissFilterTest(unittest.TestCase):
def test_is_miss_trade_result(self):
self.assertTrue(is_miss_trade_result("错过"))
self.assertFalse(is_miss_trade_result("止盈"))
def test_filter_excludes_miss(self):
rows = [
{"effective_result": "止盈", "id": 1},
{"effective_result": "错过", "id": 2},
{"result": "错过", "id": 3},
]
out = filter_trade_records_excluding_miss(rows)
self.assertEqual([r["id"] for r in out], [1])
if __name__ == "__main__":
unittest.main()
"""trade_result_lib:过滤「错过」记录."""
import unittest
from lib.trade.trade_result_lib import (
filter_trade_records_excluding_miss,
is_miss_trade_result,
)
class TradeResultMissFilterTest(unittest.TestCase):
def test_is_miss_trade_result(self):
self.assertTrue(is_miss_trade_result("错过"))
self.assertFalse(is_miss_trade_result("止盈"))
def test_filter_excludes_miss(self):
rows = [
{"effective_result": "止盈", "id": 1},
{"effective_result": "错过", "id": 2},
{"result": "错过", "id": 3},
]
out = filter_trade_records_excluding_miss(rows)
self.assertEqual([r["id"] for r in out], [1])
if __name__ == "__main__":
unittest.main()
+101 -101
View File
@@ -1,101 +1,101 @@
"""趋势回调运行中计划实际成交价重算补仓表与金额盈亏比"""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_snapshot_lib import attach_trend_dca_levels # noqa: E402
from lib.strategy.strategy_trend_lib import ( # noqa: E402
calc_trend_plan_money_metrics,
trend_leg_display_price,
)
class TestTrendDcaEnrichFills(unittest.TestCase):
def _base_plan(self, **overrides):
plan = {
"direction": "long",
"stop_loss": 0.329,
"take_profit": 0.476,
"first_order_amount": 115,
"snapshot_available_usdt": 97.98,
"risk_percent": 5,
"contract_size": 1.0,
"grid_prices_json": json.dumps([0.3465, 0.343, 0.3395, 0.336, 0.3325]),
"leg_amounts_json": json.dumps([23, 23, 23, 23, 23]),
"dca_legs": 5,
"first_order_done": 1,
"legs_done": 0,
"avg_entry_price": 0.3537,
"order_amount_open": 115,
"target_order_amount": 230,
"leg_fill_prices_json": json.dumps([0.3537]),
}
plan.update(overrides)
return plan
def test_header_money_rr_not_price_rr(self):
plan = self._base_plan()
metrics = calc_trend_plan_money_metrics(plan)
self.assertAlmostEqual(metrics["risk_amount_u"], 4.899, places=2)
self.assertIsNotNone(metrics["money_rr"])
self.assertLess(metrics["money_rr"], 4.0)
def test_done_dca_uses_actual_fill_price(self):
plan = self._base_plan(
legs_done=1,
avg_entry_price=0.3512,
order_amount_open=138,
leg_fill_prices_json=json.dumps([0.3537, 0.3458]),
)
enriched = attach_trend_dca_levels(plan)
levels = enriched["dca_levels"]
self.assertEqual(len(levels), 6)
dca1 = levels[1]
self.assertEqual(dca1["status"], "done")
self.assertAlmostEqual(dca1["price"], 0.3458, places=4)
self.assertIsNotNone(dca1["avg_entry"])
self.assertIsNotNone(dca1["rr"])
dca2 = levels[2]
self.assertEqual(dca2["status"], "pending")
self.assertAlmostEqual(dca2["price"], 0.343, places=4)
def test_missing_dca_fills_use_grid_trigger_not_inferred_price(self):
"""缺补仓成交价时触发价用计划网格末档均价对齐头部禁止反推离谱成交价"""
plan = self._base_plan(
legs_done=2,
avg_entry_price=0.3507,
order_amount_open=161,
leg_fill_prices_json=json.dumps([0.3436]),
grid_prices_json=json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
)
enriched = attach_trend_dca_levels(plan)
levels = enriched["dca_levels"]
dca1 = levels[1]
dca2 = levels[2]
self.assertEqual(dca1["status"], "done")
self.assertAlmostEqual(dca1["price"], 0.343, places=4)
self.assertEqual(dca2["status"], "done")
self.assertAlmostEqual(dca2["price"], 0.343, places=4)
self.assertAlmostEqual(dca2["avg_entry"], 0.3507, places=4)
self.assertLess(dca2["price"], 0.36)
def test_display_price_never_infers_from_target_avg(self):
"""三所共用缺记录时只用网格不因均价反推离谱触发价"""
plan = self._base_plan(
legs_done=2,
avg_entry_price=0.3507,
leg_fill_prices_json=json.dumps([0.3436]),
grid_prices_json=json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
)
self.assertAlmostEqual(trend_leg_display_price(plan, 2), 0.343, places=4)
self.assertLess(trend_leg_display_price(plan, 2), 0.36)
if __name__ == "__main__":
unittest.main()
"""趋势回调运行中计划:实际成交价重算补仓表与金额盈亏比."""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_snapshot_lib import attach_trend_dca_levels # noqa: E402
from lib.strategy.strategy_trend_lib import ( # noqa: E402
calc_trend_plan_money_metrics,
trend_leg_display_price,
)
class TestTrendDcaEnrichFills(unittest.TestCase):
def _base_plan(self, **overrides):
plan = {
"direction": "long",
"stop_loss": 0.329,
"take_profit": 0.476,
"first_order_amount": 115,
"snapshot_available_usdt": 97.98,
"risk_percent": 5,
"contract_size": 1.0,
"grid_prices_json": json.dumps([0.3465, 0.343, 0.3395, 0.336, 0.3325]),
"leg_amounts_json": json.dumps([23, 23, 23, 23, 23]),
"dca_legs": 5,
"first_order_done": 1,
"legs_done": 0,
"avg_entry_price": 0.3537,
"order_amount_open": 115,
"target_order_amount": 230,
"leg_fill_prices_json": json.dumps([0.3537]),
}
plan.update(overrides)
return plan
def test_header_money_rr_not_price_rr(self):
plan = self._base_plan()
metrics = calc_trend_plan_money_metrics(plan)
self.assertAlmostEqual(metrics["risk_amount_u"], 4.899, places=2)
self.assertIsNotNone(metrics["money_rr"])
self.assertLess(metrics["money_rr"], 4.0)
def test_done_dca_uses_actual_fill_price(self):
plan = self._base_plan(
legs_done=1,
avg_entry_price=0.3512,
order_amount_open=138,
leg_fill_prices_json=json.dumps([0.3537, 0.3458]),
)
enriched = attach_trend_dca_levels(plan)
levels = enriched["dca_levels"]
self.assertEqual(len(levels), 6)
dca1 = levels[1]
self.assertEqual(dca1["status"], "done")
self.assertAlmostEqual(dca1["price"], 0.3458, places=4)
self.assertIsNotNone(dca1["avg_entry"])
self.assertIsNotNone(dca1["rr"])
dca2 = levels[2]
self.assertEqual(dca2["status"], "pending")
self.assertAlmostEqual(dca2["price"], 0.343, places=4)
def test_missing_dca_fills_use_grid_trigger_not_inferred_price(self):
"""缺补仓成交价时:触发价用计划网格,末档均价对齐头部,禁止反推离谱成交价."""
plan = self._base_plan(
legs_done=2,
avg_entry_price=0.3507,
order_amount_open=161,
leg_fill_prices_json=json.dumps([0.3436]),
grid_prices_json=json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
)
enriched = attach_trend_dca_levels(plan)
levels = enriched["dca_levels"]
dca1 = levels[1]
dca2 = levels[2]
self.assertEqual(dca1["status"], "done")
self.assertAlmostEqual(dca1["price"], 0.343, places=4)
self.assertEqual(dca2["status"], "done")
self.assertAlmostEqual(dca2["price"], 0.343, places=4)
self.assertAlmostEqual(dca2["avg_entry"], 0.3507, places=4)
self.assertLess(dca2["price"], 0.36)
def test_display_price_never_infers_from_target_avg(self):
"""三所共用:缺记录时只用网格,不因均价反推离谱触发价."""
plan = self._base_plan(
legs_done=2,
avg_entry_price=0.3507,
leg_fill_prices_json=json.dumps([0.3436]),
grid_prices_json=json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
)
self.assertAlmostEqual(trend_leg_display_price(plan, 2), 0.343, places=4)
self.assertLess(trend_leg_display_price(plan, 2), 0.36)
if __name__ == "__main__":
unittest.main()
+43 -43
View File
@@ -1,43 +1,43 @@
"""趋势回调补仓触达与有效保证金估算"""
from lib.strategy.strategy_trend_lib import trend_dca_level_reached, trend_effective_margin_capital
def test_trend_dca_short_monotonic_up_fills_missed_legs():
"""做空价升旧逻辑需 last<level价越过 0.3437 后 last 已高于该档则永不补仓"""
direction = "short"
levels = [0.3413, 0.3437, 0.346, 0.3483, 0.3507]
pf = 0.353
filled = [lv for lv in levels if trend_dca_level_reached(direction, pf, lv)]
assert filled == levels
def test_trend_dca_short_not_before_first_level():
direction = "short"
assert not trend_dca_level_reached(direction, 0.336, 0.3413)
assert trend_dca_level_reached(direction, 0.3413, 0.3413)
def test_trend_dca_long_mark_below_trigger():
direction = "long"
assert trend_dca_level_reached(direction, 0.344, 0.3465)
assert not trend_dca_level_reached(direction, 0.347, 0.3465)
def test_trend_effective_margin_first_leg_only():
plan = {
"plan_margin_capital": 12.11,
"target_order_amount": 359.0,
"order_amount_open": 179.0,
"first_order_amount": 179.0,
}
m = trend_effective_margin_capital(plan)
assert abs(m - 12.11 * 179 / 359) < 0.01
def test_trend_effective_margin_full_position():
plan = {
"plan_margin_capital": 12.11,
"target_order_amount": 359.0,
"order_amount_open": 359.0,
}
assert trend_effective_margin_capital(plan) == 12.11
"""趋势回调:补仓触达与有效保证金估算."""
from lib.strategy.strategy_trend_lib import trend_dca_level_reached, trend_effective_margin_capital
def test_trend_dca_short_monotonic_up_fills_missed_legs():
"""做空价升:旧逻辑需 last<level,价越过 0.3437 后 last 已高于该档则永不补仓."""
direction = "short"
levels = [0.3413, 0.3437, 0.346, 0.3483, 0.3507]
pf = 0.353
filled = [lv for lv in levels if trend_dca_level_reached(direction, pf, lv)]
assert filled == levels
def test_trend_dca_short_not_before_first_level():
direction = "short"
assert not trend_dca_level_reached(direction, 0.336, 0.3413)
assert trend_dca_level_reached(direction, 0.3413, 0.3413)
def test_trend_dca_long_mark_below_trigger():
direction = "long"
assert trend_dca_level_reached(direction, 0.344, 0.3465)
assert not trend_dca_level_reached(direction, 0.347, 0.3465)
def test_trend_effective_margin_first_leg_only():
plan = {
"plan_margin_capital": 12.11,
"target_order_amount": 359.0,
"order_amount_open": 179.0,
"first_order_amount": 179.0,
}
m = trend_effective_margin_capital(plan)
assert abs(m - 12.11 * 179 / 359) < 0.01
def test_trend_effective_margin_full_position():
plan = {
"plan_margin_capital": 12.11,
"target_order_amount": 359.0,
"order_amount_open": 359.0,
}
assert trend_effective_margin_capital(plan) == 12.11
+92 -92
View File
@@ -1,92 +1,92 @@
"""趋势计划结束须写入 trade_records三所统一)。"""
from __future__ import annotations
import inspect
import sqlite3
import sys
import unittest
from pathlib import Path
from unittest.mock import MagicMock
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import _call_insert_trade_record # noqa: E402
class _GateBotLikeModule:
"""模拟 gate曾有 trend_plan_id 但缺 entry_reason 参数"""
@staticmethod
def insert_trade_record(
conn,
symbol,
monitor_type,
direction,
trigger_price,
stop_loss,
initial_stop_loss=None,
take_profit=None,
margin_capital=None,
leverage=None,
pnl_amount=0,
hold_seconds=0,
trade_style=None,
risk_amount=None,
planned_rr=None,
actual_rr=None,
result="",
miss_reason=None,
opened_at=None,
opened_at_ms=None,
closed_at=None,
closed_at_ms=None,
exchange_trade_id=None,
trend_plan_id=None,
):
conn.execute(
"INSERT INTO trade_records (symbol, monitor_type, direction, result, trend_plan_id) "
"VALUES (?,?,?,?,?)",
(symbol, monitor_type, direction, result, trend_plan_id),
)
class TestTrendFinalizeTradeRecord(unittest.TestCase):
def test_call_insert_filters_unknown_entry_reason(self):
conn = sqlite3.connect(":memory:")
conn.execute(
"CREATE TABLE trade_records (symbol TEXT, monitor_type TEXT, direction TEXT, "
"result TEXT, trend_plan_id INTEGER)"
)
m = _GateBotLikeModule()
_call_insert_trade_record(
m,
4,
dict(
conn=conn,
symbol="ONDO/USDT",
monitor_type="趋势回调",
direction="long",
trigger_price=0.35,
stop_loss=0.329,
result="止损",
entry_reason="趋势回调",
),
)
row = conn.execute(
"SELECT symbol, monitor_type, trend_plan_id FROM trade_records"
).fetchone()
self.assertEqual(row[0], "ONDO/USDT")
self.assertEqual(row[1], "趋势回调")
self.assertEqual(row[2], 4)
def test_gate_insert_accepts_entry_reason(self):
from crypto_monitor_gate import app as gate_app # noqa: E402
sig = inspect.signature(gate_app.insert_trade_record)
self.assertIn("entry_reason", sig.parameters)
self.assertIn("trend_plan_id", sig.parameters)
if __name__ == "__main__":
unittest.main()
"""趋势计划结束:须写入 trade_records(三所统一)."""
from __future__ import annotations
import inspect
import sqlite3
import sys
import unittest
from pathlib import Path
from unittest.mock import MagicMock
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import _call_insert_trade_record # noqa: E402
class _GateBotLikeModule:
"""模拟 gate:曾有 trend_plan_id 但缺 entry_reason 参数."""
@staticmethod
def insert_trade_record(
conn,
symbol,
monitor_type,
direction,
trigger_price,
stop_loss,
initial_stop_loss=None,
take_profit=None,
margin_capital=None,
leverage=None,
pnl_amount=0,
hold_seconds=0,
trade_style=None,
risk_amount=None,
planned_rr=None,
actual_rr=None,
result="",
miss_reason=None,
opened_at=None,
opened_at_ms=None,
closed_at=None,
closed_at_ms=None,
exchange_trade_id=None,
trend_plan_id=None,
):
conn.execute(
"INSERT INTO trade_records (symbol, monitor_type, direction, result, trend_plan_id) "
"VALUES (?,?,?,?,?)",
(symbol, monitor_type, direction, result, trend_plan_id),
)
class TestTrendFinalizeTradeRecord(unittest.TestCase):
def test_call_insert_filters_unknown_entry_reason(self):
conn = sqlite3.connect(":memory:")
conn.execute(
"CREATE TABLE trade_records (symbol TEXT, monitor_type TEXT, direction TEXT, "
"result TEXT, trend_plan_id INTEGER)"
)
m = _GateBotLikeModule()
_call_insert_trade_record(
m,
4,
dict(
conn=conn,
symbol="ONDO/USDT",
monitor_type="趋势回调",
direction="long",
trigger_price=0.35,
stop_loss=0.329,
result="止损",
entry_reason="趋势回调",
),
)
row = conn.execute(
"SELECT symbol, monitor_type, trend_plan_id FROM trade_records"
).fetchone()
self.assertEqual(row[0], "ONDO/USDT")
self.assertEqual(row[1], "趋势回调")
self.assertEqual(row[2], 4)
def test_gate_insert_accepts_entry_reason(self):
from crypto_monitor_gate import app as gate_app # noqa: E402
sig = inspect.signature(gate_app.insert_trade_record)
self.assertIn("entry_reason", sig.parameters)
self.assertIn("trend_plan_id", sig.parameters)
if __name__ == "__main__":
unittest.main()
+44 -44
View File
@@ -1,44 +1,44 @@
"""趋势回调中控 enrich补仓次数与加仓价"""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
from unittest.mock import MagicMock
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import _trend_add_leg_fields # noqa: E402
class TestTrendHubEnrich(unittest.TestCase):
def test_add_count_and_prices(self):
mock_ex = MagicMock()
mock_ex.price_to_precision = lambda sym, px: f"{float(px):.4f}"
app_mod = MagicMock()
app_mod.exchange = mock_ex
app_mod.ensure_markets_loaded = MagicMock()
app_mod.normalize_exchange_symbol = lambda s: s
cfg = {"app_module": app_mod}
raw = {
"symbol": "ETH/USDT",
"exchange_symbol": "ETH/USDT:USDT",
"legs_done": 2,
"dca_legs": 5,
"grid_prices_json": json.dumps([1800.1, 1750.2, 1700.3]),
"stop_loss": 1600,
"take_profit": 2000,
"avg_entry_price": 1820.5,
}
out = _trend_add_leg_fields(cfg, raw)
self.assertEqual(out["add_count"], 2)
self.assertEqual(out["add_count_total"], 5)
self.assertEqual(out["add_prices"], [1800.1, 1750.2])
self.assertEqual(len(out["add_prices_display"]), 2)
self.assertEqual(out["stop_loss_display"], "1600.0000")
if __name__ == "__main__":
unittest.main()
"""趋势回调中控 enrich:补仓次数与加仓价."""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
from unittest.mock import MagicMock
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import _trend_add_leg_fields # noqa: E402
class TestTrendHubEnrich(unittest.TestCase):
def test_add_count_and_prices(self):
mock_ex = MagicMock()
mock_ex.price_to_precision = lambda sym, px: f"{float(px):.4f}"
app_mod = MagicMock()
app_mod.exchange = mock_ex
app_mod.ensure_markets_loaded = MagicMock()
app_mod.normalize_exchange_symbol = lambda s: s
cfg = {"app_module": app_mod}
raw = {
"symbol": "ETH/USDT",
"exchange_symbol": "ETH/USDT:USDT",
"legs_done": 2,
"dca_legs": 5,
"grid_prices_json": json.dumps([1800.1, 1750.2, 1700.3]),
"stop_loss": 1600,
"take_profit": 2000,
"avg_entry_price": 1820.5,
}
out = _trend_add_leg_fields(cfg, raw)
self.assertEqual(out["add_count"], 2)
self.assertEqual(out["add_count_total"], 5)
self.assertEqual(out["add_prices"], [1800.1, 1750.2])
self.assertEqual(len(out["add_prices_display"]), 2)
self.assertEqual(out["stop_loss_display"], "1600.0000")
if __name__ == "__main__":
unittest.main()
+92 -92
View File
@@ -1,92 +1,92 @@
"""三所趋势 enrich实例与中控 monitor 字段一致"""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import ( # noqa: E402
enrich_trend_plan,
enrich_trend_plan_for_hub,
)
class _FakeModule:
@staticmethod
def normalize_exchange_symbol(sym):
return sym
@staticmethod
def ensure_markets_loaded():
return None
@staticmethod
def get_live_position_exchange_metrics(ex_sym, direction, order_leverage=None):
return {
"entry_price": 0.3507,
"mark_price": 0.3431,
"unrealized_pnl": -1.23,
}
@staticmethod
def get_contract_size(_ex_sym):
return 1.0
class TestTrendHubEnrichUnified(unittest.TestCase):
def _cfg(self):
return {
"app_module": _FakeModule(),
"breakeven_offset_pct": 0.3,
"row_to_dict": lambda row: dict(row) if not isinstance(row, dict) else row,
}
def _plan_row(self):
return {
"id": 4,
"symbol": "ONDO/USDT:USDT",
"exchange_symbol": "ONDO/USDT:USDT",
"direction": "long",
"stop_loss": 0.329,
"take_profit": 0.476,
"add_upper": 0.35,
"first_order_amount": 115,
"snapshot_available_usdt": 97.98,
"risk_percent": 5,
"contract_size": 1.0,
"grid_prices_json": json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
"leg_amounts_json": json.dumps([23, 23, 23, 23, 23]),
"dca_legs": 5,
"first_order_done": 1,
"legs_done": 2,
"avg_entry_price": 0.3434,
"order_amount_open": 161,
"leg_fill_prices_json": json.dumps([0.3436]),
"leverage": 10,
"plan_margin_capital": 8.17,
}
def test_hub_and_page_share_live_avg_and_dca_levels(self):
cfg = self._cfg()
row = self._plan_row()
page = enrich_trend_plan(cfg, row)
hub = enrich_trend_plan_for_hub(cfg, row)
self.assertAlmostEqual(page["avg_entry_price"], 0.3507, places=4)
self.assertAlmostEqual(hub["avg_entry_price"], 0.3507, places=4)
self.assertIn("dca_levels", page)
self.assertIn("dca_levels", hub)
last_done = hub["dca_levels"][2]
self.assertEqual(last_done["status"], "done")
self.assertAlmostEqual(last_done["price"], 0.343, places=4)
self.assertAlmostEqual(last_done["avg_entry"], 0.3507, places=4)
self.assertLess(last_done["price"], 0.36)
self.assertEqual(hub.get("monitor_source"), "趋势回调计划")
self.assertEqual(hub.get("add_count"), 2)
if __name__ == "__main__":
unittest.main()
"""三所趋势 enrich:实例与中控 monitor 字段一致."""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_register import ( # noqa: E402
enrich_trend_plan,
enrich_trend_plan_for_hub,
)
class _FakeModule:
@staticmethod
def normalize_exchange_symbol(sym):
return sym
@staticmethod
def ensure_markets_loaded():
return None
@staticmethod
def get_live_position_exchange_metrics(ex_sym, direction, order_leverage=None):
return {
"entry_price": 0.3507,
"mark_price": 0.3431,
"unrealized_pnl": -1.23,
}
@staticmethod
def get_contract_size(_ex_sym):
return 1.0
class TestTrendHubEnrichUnified(unittest.TestCase):
def _cfg(self):
return {
"app_module": _FakeModule(),
"breakeven_offset_pct": 0.3,
"row_to_dict": lambda row: dict(row) if not isinstance(row, dict) else row,
}
def _plan_row(self):
return {
"id": 4,
"symbol": "ONDO/USDT:USDT",
"exchange_symbol": "ONDO/USDT:USDT",
"direction": "long",
"stop_loss": 0.329,
"take_profit": 0.476,
"add_upper": 0.35,
"first_order_amount": 115,
"snapshot_available_usdt": 97.98,
"risk_percent": 5,
"contract_size": 1.0,
"grid_prices_json": json.dumps([0.343, 0.343, 0.3395, 0.336, 0.3325]),
"leg_amounts_json": json.dumps([23, 23, 23, 23, 23]),
"dca_legs": 5,
"first_order_done": 1,
"legs_done": 2,
"avg_entry_price": 0.3434,
"order_amount_open": 161,
"leg_fill_prices_json": json.dumps([0.3436]),
"leverage": 10,
"plan_margin_capital": 8.17,
}
def test_hub_and_page_share_live_avg_and_dca_levels(self):
cfg = self._cfg()
row = self._plan_row()
page = enrich_trend_plan(cfg, row)
hub = enrich_trend_plan_for_hub(cfg, row)
self.assertAlmostEqual(page["avg_entry_price"], 0.3507, places=4)
self.assertAlmostEqual(hub["avg_entry_price"], 0.3507, places=4)
self.assertIn("dca_levels", page)
self.assertIn("dca_levels", hub)
last_done = hub["dca_levels"][2]
self.assertEqual(last_done["status"], "done")
self.assertAlmostEqual(last_done["price"], 0.343, places=4)
self.assertAlmostEqual(last_done["avg_entry"], 0.3507, places=4)
self.assertLess(last_done["price"], 0.36)
self.assertEqual(hub.get("monitor_source"), "趋势回调计划")
self.assertEqual(hub.get("add_count"), 2)
if __name__ == "__main__":
unittest.main()
+41 -41
View File
@@ -1,41 +1,41 @@
"""趋势补仓下单空 params 不得变成 Noneccxt 会报 not iterable)。"""
from __future__ import annotations
import unittest
from unittest.mock import MagicMock
from lib.strategy.strategy_trend_exchange import trend_market_add
class TestTrendMarketAddParams(unittest.TestCase):
def test_empty_gate_params_not_passed_as_none(self):
ex = MagicMock()
ex.create_order.return_value = {"id": "1", "average": 0.34}
app = MagicMock()
app.exchange = ex
app.ensure_markets_loaded = MagicMock()
app.build_gate_order_params = MagicMock(return_value={})
cfg = {"app_module": app}
trend_market_add(cfg, "ONDO/USDT:USDT", "long", 23, 10)
args = ex.create_order.call_args
self.assertEqual(args[0][5], {})
self.assertIsNotNone(args[0][5])
def test_binance_oneway_empty_params_not_passed_as_none(self):
ex = MagicMock()
ex.create_order.return_value = {"id": "1"}
app = MagicMock(spec=["exchange", "ensure_markets_loaded", "build_binance_order_params"])
app.exchange = ex
app.ensure_markets_loaded = MagicMock()
app.build_binance_order_params = MagicMock(return_value={})
cfg = {"app_module": app}
trend_market_add(cfg, "BTC/USDT:USDT", "long", 1, 10)
self.assertEqual(ex.create_order.call_args[0][5], {})
if __name__ == "__main__":
unittest.main()
"""趋势补仓下单:空 params 不得变成 None(ccxt 会报 not iterable)."""
from __future__ import annotations
import unittest
from unittest.mock import MagicMock
from lib.strategy.strategy_trend_exchange import trend_market_add
class TestTrendMarketAddParams(unittest.TestCase):
def test_empty_gate_params_not_passed_as_none(self):
ex = MagicMock()
ex.create_order.return_value = {"id": "1", "average": 0.34}
app = MagicMock()
app.exchange = ex
app.ensure_markets_loaded = MagicMock()
app.build_gate_order_params = MagicMock(return_value={})
cfg = {"app_module": app}
trend_market_add(cfg, "ONDO/USDT:USDT", "long", 23, 10)
args = ex.create_order.call_args
self.assertEqual(args[0][5], {})
self.assertIsNotNone(args[0][5])
def test_binance_oneway_empty_params_not_passed_as_none(self):
ex = MagicMock()
ex.create_order.return_value = {"id": "1"}
app = MagicMock(spec=["exchange", "ensure_markets_loaded", "build_binance_order_params"])
app.exchange = ex
app.ensure_markets_loaded = MagicMock()
app.build_binance_order_params = MagicMock(return_value={})
cfg = {"app_module": app}
trend_market_add(cfg, "BTC/USDT:USDT", "long", 1, 10)
self.assertEqual(ex.create_order.call_args[0][5], {})
if __name__ == "__main__":
unittest.main()
+58 -58
View File
@@ -1,58 +1,58 @@
"""趋势回调预览止盈盈利 U止损金额 U金额盈亏比"""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_lib import ( # noqa: E402
build_trend_preview_level_rows,
calc_money_reward_risk_ratio,
calc_risk_budget_usdt,
calc_tp_profit_usdt,
)
class TestTrendPreviewTp(unittest.TestCase):
def test_risk_budget_from_snapshot(self):
self.assertAlmostEqual(calc_risk_budget_usdt(110.73, 5), 5.5365, places=2)
def test_short_profit_at_form_take_profit(self):
profit = calc_tp_profit_usdt("short", 72.53, 66.0, 1114, 0.00167)
self.assertIsNotNone(profit)
self.assertGreater(profit, 0)
rr = calc_money_reward_risk_ratio(profit, 5.5365)
self.assertIsNotNone(rr)
self.assertGreater(rr, 1.5)
def test_preview_levels_use_money_rr(self):
preview = {
"direction": "short",
"live_price_ref": 72.53,
"stop_loss": 75.5,
"take_profit": 66.0,
"first_order_amount": 1114,
"snapshot_available_usdt": 110.73,
"risk_percent": 5,
"contract_size": 0.00167,
"grid_prices_json": json.dumps([73.42, 73.83]),
"leg_amounts_json": json.dumps([222, 222]),
}
enriched, rows = build_trend_preview_level_rows(preview)
self.assertAlmostEqual(enriched["preview_risk_amount_u"], 5.5365, places=2)
self.assertEqual(enriched["preview_take_profit_price"], 66.0)
self.assertEqual(len(rows), 3)
self.assertEqual(rows[0]["label"], "首仓")
self.assertEqual(rows[0]["risk_u"], enriched["preview_risk_amount_u"])
self.assertIsNotNone(rows[0]["profit_u"])
self.assertAlmostEqual(rows[0]["rr"], rows[0]["profit_u"] / 5.5365, places=2)
self.assertEqual(rows[1]["risk_u"], enriched["preview_risk_amount_u"])
self.assertGreater(rows[2]["profit_u"], rows[1]["profit_u"])
if __name__ == "__main__":
unittest.main()
"""趋势回调预览:止盈盈利 U,止损金额 U,金额盈亏比."""
from __future__ import annotations
import json
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from lib.strategy.strategy_trend_lib import ( # noqa: E402
build_trend_preview_level_rows,
calc_money_reward_risk_ratio,
calc_risk_budget_usdt,
calc_tp_profit_usdt,
)
class TestTrendPreviewTp(unittest.TestCase):
def test_risk_budget_from_snapshot(self):
self.assertAlmostEqual(calc_risk_budget_usdt(110.73, 5), 5.5365, places=2)
def test_short_profit_at_form_take_profit(self):
profit = calc_tp_profit_usdt("short", 72.53, 66.0, 1114, 0.00167)
self.assertIsNotNone(profit)
self.assertGreater(profit, 0)
rr = calc_money_reward_risk_ratio(profit, 5.5365)
self.assertIsNotNone(rr)
self.assertGreater(rr, 1.5)
def test_preview_levels_use_money_rr(self):
preview = {
"direction": "short",
"live_price_ref": 72.53,
"stop_loss": 75.5,
"take_profit": 66.0,
"first_order_amount": 1114,
"snapshot_available_usdt": 110.73,
"risk_percent": 5,
"contract_size": 0.00167,
"grid_prices_json": json.dumps([73.42, 73.83]),
"leg_amounts_json": json.dumps([222, 222]),
}
enriched, rows = build_trend_preview_level_rows(preview)
self.assertAlmostEqual(enriched["preview_risk_amount_u"], 5.5365, places=2)
self.assertEqual(enriched["preview_take_profit_price"], 66.0)
self.assertEqual(len(rows), 3)
self.assertEqual(rows[0]["label"], "首仓")
self.assertEqual(rows[0]["risk_u"], enriched["preview_risk_amount_u"])
self.assertIsNotNone(rows[0]["profit_u"])
self.assertAlmostEqual(rows[0]["rr"], rows[0]["profit_u"] / 5.5365, places=2)
self.assertEqual(rows[1]["risk_u"], enriched["preview_risk_amount_u"])
self.assertGreater(rows[2]["profit_u"], rows[1]["profit_u"])
if __name__ == "__main__":
unittest.main()
+85 -85
View File
@@ -1,85 +1,85 @@
"""触价开仓回调/突破关键位监控单元测试"""
from lib.key_monitor.trigger_entry_key_monitor_lib import (
BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE,
CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE,
LEGACY_TRIGGER_ENTRY_MONITOR_TYPE,
TRIGGER_ENTRY_MONITOR_TYPES,
TRIGGER_ENTRY_VALIDITY_HOURS,
breakout_trigger_entry_crossed,
check_trigger_entry_intent_limit,
is_breakout_trigger_entry_key_monitor_type,
is_trigger_entry_key_monitor_type,
trigger_entry_invalidate,
trigger_entry_reached,
trigger_should_fire,
validate_trigger_entry_geometry,
)
class _FakeConn:
def execute(self, sql, params=()):
class R:
def fetchone(self_inner):
return (2,)
return R()
def test_trigger_entry_reached_long():
assert trigger_entry_reached("long", 2049.0, 2050.0) is True
assert trigger_entry_reached("long", 2051.0, 2050.0) is False
def test_breakout_cross_long_up():
assert breakout_trigger_entry_crossed("long", 99.0, 100.5, 100.0) is True
assert breakout_trigger_entry_crossed("long", None, 101.0, 100.0) is True
assert breakout_trigger_entry_crossed("long", 100.0, 100.0, 100.0) is False
def test_breakout_cross_short_down():
assert breakout_trigger_entry_crossed("short", 101.0, 99.5, 100.0) is True
assert breakout_trigger_entry_crossed("short", None, 99.0, 100.0) is True
def test_trigger_should_fire_modes():
assert trigger_should_fire(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE, "long", 2049.0, 2050.0) is True
assert trigger_should_fire(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE, "long", 100.5, 100.0, 99.0) is True
def test_validate_geometry_callback_long():
assert validate_trigger_entry_geometry("long", 2050, 2000, 2100, 2090) is None
def test_validate_geometry_breakout_short_requires_mark_above_entry():
assert (
validate_trigger_entry_geometry(
"short", 551, 568, 540, 560, monitor_type=BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE
)
is None
)
err = validate_trigger_entry_geometry(
"short", 551, 568, 540, 550, monitor_type=BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE
)
assert err is not None
assert "高于入场价" in err
def test_invalidate_breakout_sl_side():
assert trigger_entry_invalidate(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE, "long", 96, 97, 110) == "sl"
assert trigger_entry_invalidate(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE, "long", 96, 97, 110) is None
def test_intent_limit():
ok, msg = check_trigger_entry_intent_limit(_FakeConn(), "2026-06-07", 2, 3)
assert ok is False
assert "意图" in msg
def test_type_names():
assert is_trigger_entry_key_monitor_type(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_trigger_entry_key_monitor_type(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_trigger_entry_key_monitor_type(LEGACY_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_breakout_trigger_entry_key_monitor_type(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE)
assert CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE in TRIGGER_ENTRY_MONITOR_TYPES
assert BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE in TRIGGER_ENTRY_MONITOR_TYPES
assert TRIGGER_ENTRY_VALIDITY_HOURS == 24
"""触价开仓(回调/突破)关键位监控单元测试."""
from lib.key_monitor.trigger_entry_key_monitor_lib import (
BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE,
CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE,
LEGACY_TRIGGER_ENTRY_MONITOR_TYPE,
TRIGGER_ENTRY_MONITOR_TYPES,
TRIGGER_ENTRY_VALIDITY_HOURS,
breakout_trigger_entry_crossed,
check_trigger_entry_intent_limit,
is_breakout_trigger_entry_key_monitor_type,
is_trigger_entry_key_monitor_type,
trigger_entry_invalidate,
trigger_entry_reached,
trigger_should_fire,
validate_trigger_entry_geometry,
)
class _FakeConn:
def execute(self, sql, params=()):
class R:
def fetchone(self_inner):
return (2,)
return R()
def test_trigger_entry_reached_long():
assert trigger_entry_reached("long", 2049.0, 2050.0) is True
assert trigger_entry_reached("long", 2051.0, 2050.0) is False
def test_breakout_cross_long_up():
assert breakout_trigger_entry_crossed("long", 99.0, 100.5, 100.0) is True
assert breakout_trigger_entry_crossed("long", None, 101.0, 100.0) is True
assert breakout_trigger_entry_crossed("long", 100.0, 100.0, 100.0) is False
def test_breakout_cross_short_down():
assert breakout_trigger_entry_crossed("short", 101.0, 99.5, 100.0) is True
assert breakout_trigger_entry_crossed("short", None, 99.0, 100.0) is True
def test_trigger_should_fire_modes():
assert trigger_should_fire(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE, "long", 2049.0, 2050.0) is True
assert trigger_should_fire(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE, "long", 100.5, 100.0, 99.0) is True
def test_validate_geometry_callback_long():
assert validate_trigger_entry_geometry("long", 2050, 2000, 2100, 2090) is None
def test_validate_geometry_breakout_short_requires_mark_above_entry():
assert (
validate_trigger_entry_geometry(
"short", 551, 568, 540, 560, monitor_type=BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE
)
is None
)
err = validate_trigger_entry_geometry(
"short", 551, 568, 540, 550, monitor_type=BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE
)
assert err is not None
assert "高于入场价" in err
def test_invalidate_breakout_sl_side():
assert trigger_entry_invalidate(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE, "long", 96, 97, 110) == "sl"
assert trigger_entry_invalidate(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE, "long", 96, 97, 110) is None
def test_intent_limit():
ok, msg = check_trigger_entry_intent_limit(_FakeConn(), "2026-06-07", 2, 3)
assert ok is False
assert "意图" in msg
def test_type_names():
assert is_trigger_entry_key_monitor_type(CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_trigger_entry_key_monitor_type(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_trigger_entry_key_monitor_type(LEGACY_TRIGGER_ENTRY_MONITOR_TYPE)
assert is_breakout_trigger_entry_key_monitor_type(BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE)
assert CALLBACK_TRIGGER_ENTRY_MONITOR_TYPE in TRIGGER_ENTRY_MONITOR_TYPES
assert BREAKOUT_TRIGGER_ENTRY_MONITOR_TYPE in TRIGGER_ENTRY_MONITOR_TYPES
assert TRIGGER_ENTRY_VALIDITY_HOURS == 24