-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtest_strategy.py
More file actions
195 lines (154 loc) · 6.45 KB
/
Copy pathtest_strategy.py
File metadata and controls
195 lines (154 loc) · 6.45 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
# -*- coding: utf-8 -*-
"""
均值回归策略测试脚本
"""
import sys
import os
from datetime import datetime, timedelta
# 添加当前目录到Python路径
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
from mean_reversion_csv import mean_reversion_csv_strategy
from mean_reversion_mysql import mean_reversion_mysql_strategy
def test_csv_strategy():
"""测试CSV版本的策略"""
print("=" * 60)
print("测试CSV版本的均值回归策略")
print("=" * 60)
# 测试参数
ts_code = "000001.SZ"
start_date = "2023-01-03"
end_date = "2023-02-28"
try:
# 运行策略
results = mean_reversion_csv_strategy(ts_code, start_date, end_date)
# 打印结果
print("\n策略结果:")
print("-" * 40)
if "error" in results:
print(f"错误: {results['error']}")
return False
# 基本信息
print(f"股票代码: {results.get('ts_code', 'N/A')}")
print(f"开始日期: {results.get('start_date', 'N/A')}")
print(f"结束日期: {results.get('end_date', 'N/A')}")
print(f"数据点数: {results.get('data_points', 'N/A')}")
# 收益信息
print(f"\n收益统计:")
print(f"策略总收益: {results.get('total_return', 0):.4f} ({results.get('total_return', 0)*100:.2f}%)")
print(f"买入持有收益: {results.get('buy_hold_return', 0):.4f} ({results.get('buy_hold_return', 0)*100:.2f}%)")
print(f"超额收益: {results.get('excess_return', 0):.4f} ({results.get('excess_return', 0)*100:.2f}%)")
print(f"夏普比率: {results.get('sharpe_ratio', 0):.4f}")
print(f"最大回撤: {results.get('max_drawdown', 0):.4f} ({results.get('max_drawdown', 0)*100:.2f}%)")
# 交易统计
print(f"\n交易统计:")
print(f"买入信号数: {results.get('buy_signals', 0)}")
print(f"卖出信号数: {results.get('sell_signals', 0)}")
print(f"总交易次数: {results.get('total_trades', 0)}")
return True
except Exception as e:
print(f"测试失败: {e}")
return False
def test_mysql_strategy():
"""测试MySQL版本的策略"""
print("=" * 60)
print("测试MySQL版本的均值回归策略")
print("=" * 60)
# 测试参数
ts_code = "000001.SZ"
start_date = "2023-01-03"
end_date = "2023-02-28"
try:
# 运行策略
results = mean_reversion_mysql_strategy(ts_code, start_date, end_date)
# 打印结果
print("\n策略结果:")
print("-" * 40)
if "error" in results:
print(f"错误: {results['error']}")
print("注意: 这可能是正常的,因为可能没有配置MySQL数据库")
return True # 不视为错误,因为可能没有数据库
# 基本信息
print(f"股票代码: {results.get('ts_code', 'N/A')}")
print(f"开始日期: {results.get('start_date', 'N/A')}")
print(f"结束日期: {results.get('end_date', 'N/A')}")
print(f"数据点数: {results.get('data_points', 'N/A')}")
# 收益信息
print(f"\n收益统计:")
print(f"策略总收益: {results.get('total_return', 0):.4f} ({results.get('total_return', 0)*100:.2f}%)")
print(f"买入持有收益: {results.get('buy_hold_return', 0):.4f} ({results.get('buy_hold_return', 0)*100:.2f}%)")
print(f"超额收益: {results.get('excess_return', 0):.4f} ({results.get('excess_return', 0)*100:.2f}%)")
print(f"夏普比率: {results.get('sharpe_ratio', 0):.4f}")
print(f"最大回撤: {results.get('max_drawdown', 0):.4f} ({results.get('max_drawdown', 0)*100:.2f}%)")
# 交易统计
print(f"\n交易统计:")
print(f"买入信号数: {results.get('buy_signals', 0)}")
print(f"卖出信号数: {results.get('sell_signals', 0)}")
print(f"总交易次数: {results.get('total_trades', 0)}")
return True
except Exception as e:
print(f"测试失败: {e}")
print("注意: 这可能是正常的,因为可能没有配置MySQL数据库")
return True # 不视为错误,因为可能没有数据库
def test_data_validation():
"""测试数据验证功能"""
print("=" * 60)
print("测试数据验证功能")
print("=" * 60)
# 测试无效日期
print("\n测试无效日期格式:")
results = mean_reversion_csv_strategy("000001.SZ", "2023-13-01", "2023-12-31")
if "error" in results and "未获取到数据" in results["error"]:
print("✓ 正确检测到无效日期")
else:
print("错误: 应该检测到无效日期")
return False
# 测试无效股票代码
print("\n测试无效股票代码:")
results = mean_reversion_csv_strategy("INVALID_CODE", "2023-01-01", "2023-12-31")
if "error" in results:
print("✓ 正确检测到无效股票代码")
else:
print("错误: 应该检测到无效股票代码")
return False
return True
def main():
"""主测试函数"""
print("开始测试均值回归策略...")
print(f"测试时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
# 测试计数器
passed_tests = 0
total_tests = 3
# 测试1: CSV策略
if test_csv_strategy():
print("\n✓ CSV策略测试通过")
passed_tests += 1
else:
print("\n✗ CSV策略测试失败")
# 测试2: MySQL策略
if test_mysql_strategy():
print("\n✓ MySQL策略测试通过")
passed_tests += 1
else:
print("\n✗ MySQL策略测试失败")
# 测试3: 数据验证
if test_data_validation():
print("\n✓ 数据验证测试通过")
passed_tests += 1
else:
print("\n✗ 数据验证测试失败")
# 总结
print("\n" + "=" * 60)
print("测试总结")
print("=" * 60)
print(f"总测试数: {total_tests}")
print(f"通过测试: {passed_tests}")
print(f"失败测试: {total_tests - passed_tests}")
print(f"成功率: {passed_tests/total_tests*100:.1f}%")
if passed_tests == total_tests:
print("\n🎉 所有测试通过!")
else:
print("\n⚠️ 部分测试失败,请检查代码")
return passed_tests == total_tests
if __name__ == "__main__":
success = main()
sys.exit(0 if success else 1)