xtquant-share 1.1.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- xqshare/__init__.py +38 -0
- xqshare/auth.py +466 -0
- xqshare/client.py +682 -0
- xqshare/server.py +868 -0
- xqshare/tools/__init__.py +7 -0
- xqshare/tools/common.py +457 -0
- xqshare/tools/xtdata.py +98 -0
- xqshare/tools/xttrader.py +122 -0
- xtquant_share-1.1.2.dist-info/METADATA +756 -0
- xtquant_share-1.1.2.dist-info/RECORD +14 -0
- xtquant_share-1.1.2.dist-info/WHEEL +5 -0
- xtquant_share-1.1.2.dist-info/entry_points.txt +4 -0
- xtquant_share-1.1.2.dist-info/licenses/LICENSE +21 -0
- xtquant_share-1.1.2.dist-info/top_level.txt +1 -0
xqshare/server.py
ADDED
|
@@ -0,0 +1,868 @@
|
|
|
1
|
+
"""
|
|
2
|
+
XtQuant Share (xqshare) Server - Run on Windows to provide xtquant proxy service
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import rpyc
|
|
6
|
+
from rpyc.utils.server import ThreadedServer
|
|
7
|
+
import time
|
|
8
|
+
import os
|
|
9
|
+
import ssl
|
|
10
|
+
import logging
|
|
11
|
+
import functools
|
|
12
|
+
import json
|
|
13
|
+
from datetime import datetime
|
|
14
|
+
from typing import Any, Dict, Optional
|
|
15
|
+
|
|
16
|
+
# 导入权限模块
|
|
17
|
+
from .auth import (
|
|
18
|
+
PermissionChecker,
|
|
19
|
+
PermissionError,
|
|
20
|
+
AccountLevel,
|
|
21
|
+
Permission,
|
|
22
|
+
get_permission_checker,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
# Import xtquant (only available on Windows)
|
|
26
|
+
try:
|
|
27
|
+
import xtquant.xtdata as xtdata
|
|
28
|
+
import xtquant.xttrader as xttrader
|
|
29
|
+
import xtquant.xttype as xttype
|
|
30
|
+
import xtquant.xtconstant as xtconstant
|
|
31
|
+
from xtquant.xttrader import XtQuantTrader
|
|
32
|
+
XTQUANT_AVAILABLE = True
|
|
33
|
+
except ImportError:
|
|
34
|
+
XTQUANT_AVAILABLE = False
|
|
35
|
+
xtdata = None
|
|
36
|
+
xttrader = None
|
|
37
|
+
xttype = None
|
|
38
|
+
xtconstant = None
|
|
39
|
+
XtQuantTrader = None
|
|
40
|
+
|
|
41
|
+
# xtview 模块单独导入(某些版本可能不存在)
|
|
42
|
+
try:
|
|
43
|
+
import xtquant.xtview as xtview
|
|
44
|
+
XTVIEW_AVAILABLE = True
|
|
45
|
+
except ImportError:
|
|
46
|
+
xtview = None
|
|
47
|
+
XTVIEW_AVAILABLE = False
|
|
48
|
+
|
|
49
|
+
# QmtDataReader 导入(用于 datadir 文件解析能力)
|
|
50
|
+
try:
|
|
51
|
+
from utils.qmt_datadir import QmtDataReader
|
|
52
|
+
QMTDATAREADER_AVAILABLE = True
|
|
53
|
+
except ImportError:
|
|
54
|
+
try:
|
|
55
|
+
# 兼容:直接从 qmt_datadir 包导入
|
|
56
|
+
import sys as _sys
|
|
57
|
+
import os as _os
|
|
58
|
+
_pkg_dir = _os.path.join(_os.path.dirname(__file__), '..', 'utils')
|
|
59
|
+
if _os.path.isdir(_pkg_dir):
|
|
60
|
+
_sys.path.insert(0, _os.path.dirname(_pkg_dir))
|
|
61
|
+
from utils.qmt_datadir import QmtDataReader
|
|
62
|
+
QMTDATAREADER_AVAILABLE = True
|
|
63
|
+
except ImportError:
|
|
64
|
+
QmtDataReader = None
|
|
65
|
+
QMTDATAREADER_AVAILABLE = False
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
# ==================== 日志配置 ====================
|
|
69
|
+
|
|
70
|
+
def setup_logging(log_dir: str = None, log_level: str = "INFO"):
|
|
71
|
+
"""配置日志系统"""
|
|
72
|
+
if log_dir is None:
|
|
73
|
+
log_dir = os.environ.get("XQSHARE_LOG_DIR", "logs")
|
|
74
|
+
os.makedirs(log_dir, exist_ok=True)
|
|
75
|
+
|
|
76
|
+
formatter = logging.Formatter(
|
|
77
|
+
fmt='%(asctime)s.%(msecs)03d | %(levelname)-8s | %(name)s | %(message)s',
|
|
78
|
+
datefmt='%Y-%m-%d %H:%M:%S'
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
root_logger = logging.getLogger()
|
|
82
|
+
root_logger.setLevel(getattr(logging, log_level.upper()))
|
|
83
|
+
|
|
84
|
+
console_handler = logging.StreamHandler()
|
|
85
|
+
console_handler.setFormatter(formatter)
|
|
86
|
+
console_handler.setLevel(logging.INFO)
|
|
87
|
+
root_logger.addHandler(console_handler)
|
|
88
|
+
|
|
89
|
+
file_handler = logging.FileHandler(
|
|
90
|
+
os.path.join(log_dir, f"xtquant_service_{datetime.now().strftime('%Y%m%d')}.log"),
|
|
91
|
+
encoding='utf-8'
|
|
92
|
+
)
|
|
93
|
+
file_handler.setFormatter(formatter)
|
|
94
|
+
file_handler.setLevel(logging.DEBUG)
|
|
95
|
+
root_logger.addHandler(file_handler)
|
|
96
|
+
|
|
97
|
+
api_handler = logging.FileHandler(
|
|
98
|
+
os.path.join(log_dir, f"api_calls_{datetime.now().strftime('%Y%m%d')}.log"),
|
|
99
|
+
encoding='utf-8'
|
|
100
|
+
)
|
|
101
|
+
api_handler.setFormatter(formatter)
|
|
102
|
+
api_logger = logging.getLogger('api')
|
|
103
|
+
api_logger.addHandler(api_handler)
|
|
104
|
+
api_logger.setLevel(logging.DEBUG)
|
|
105
|
+
|
|
106
|
+
return logging.getLogger(__name__)
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
logger = None
|
|
110
|
+
api_logger = None
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def _init_logging(log_level="INFO"):
|
|
114
|
+
global logger, api_logger
|
|
115
|
+
logger = setup_logging(log_level=log_level)
|
|
116
|
+
api_logger = logging.getLogger('api')
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
# ==================== 日志装饰器 ====================
|
|
120
|
+
|
|
121
|
+
def _log_call(name: str, client_info: str, func, *args, **kwargs):
|
|
122
|
+
"""通用的 API 调用日志记录函数"""
|
|
123
|
+
try:
|
|
124
|
+
args_str = str(args)[:200] if args else ""
|
|
125
|
+
kwargs_str = str(kwargs)[:200] if kwargs else ""
|
|
126
|
+
except:
|
|
127
|
+
args_str = "<unserializable>"
|
|
128
|
+
kwargs_str = ""
|
|
129
|
+
|
|
130
|
+
api_logger.info(f"[CALL] {name} | client={client_info} | args={args_str} | kwargs={kwargs_str}")
|
|
131
|
+
|
|
132
|
+
start_time = time.perf_counter()
|
|
133
|
+
try:
|
|
134
|
+
result = func(*args, **kwargs)
|
|
135
|
+
elapsed_ms = (time.perf_counter() - start_time) * 1000
|
|
136
|
+
result_summary = _summarize_result(result)
|
|
137
|
+
api_logger.info(f"[OK] {name} | elapsed={elapsed_ms:.2f}ms | result={result_summary}")
|
|
138
|
+
return result
|
|
139
|
+
except Exception as e:
|
|
140
|
+
elapsed_ms = (time.perf_counter() - start_time) * 1000
|
|
141
|
+
api_logger.error(f"[ERROR] {name} | elapsed={elapsed_ms:.2f}ms | error={type(e).__name__}: {str(e)[:200]}")
|
|
142
|
+
raise
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def log_api_call(func_name: str = None):
|
|
146
|
+
"""记录 API 调用的装饰器"""
|
|
147
|
+
def decorator(func):
|
|
148
|
+
def wrapper(self, *args, **kwargs):
|
|
149
|
+
name = func_name or func.__name__
|
|
150
|
+
client_info = getattr(self, '_client_info', 'unknown')
|
|
151
|
+
return _log_call(name, client_info, func, self, *args, **kwargs)
|
|
152
|
+
return wrapper
|
|
153
|
+
return decorator
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def _summarize_result(result: Any, max_len: int = 200) -> str:
|
|
157
|
+
"""生成返回值摘要"""
|
|
158
|
+
try:
|
|
159
|
+
if result is None:
|
|
160
|
+
return "None"
|
|
161
|
+
elif isinstance(result, (int, float, bool, str)):
|
|
162
|
+
s = str(result)
|
|
163
|
+
return s if len(s) <= max_len else s[:max_len] + "..."
|
|
164
|
+
elif isinstance(result, (list, tuple)):
|
|
165
|
+
return f"{type(result).__name__}[len={len(result)}]"
|
|
166
|
+
elif isinstance(result, dict):
|
|
167
|
+
keys = list(result.keys())[:5]
|
|
168
|
+
return f"dict{{{', '.join(map(str, keys))}{'...' if len(result) > 5 else ''}}}"
|
|
169
|
+
elif hasattr(result, '__class__'):
|
|
170
|
+
return f"<{result.__class__.__module__}.{result.__class__.__name__}>"
|
|
171
|
+
else:
|
|
172
|
+
return str(type(result))
|
|
173
|
+
except:
|
|
174
|
+
return "<unserializable>"
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
# ==================== 异常定义 ====================
|
|
178
|
+
|
|
179
|
+
class AuthError(Exception):
|
|
180
|
+
"""认证错误"""
|
|
181
|
+
pass
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
# ==================== 序列化传输优化 ====================
|
|
185
|
+
|
|
186
|
+
# 需要序列化传输的类型标记
|
|
187
|
+
SERIALIZED_MARKER = "__xqshare_serialized__"
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def _serialize_for_transfer(result):
|
|
191
|
+
"""将结果序列化以优化 RPyC 传输性能
|
|
192
|
+
|
|
193
|
+
对于大型列表/字典/DataFrame,序列化后传输比逐元素传输快很多。
|
|
194
|
+
|
|
195
|
+
Args:
|
|
196
|
+
result: API 调用返回值
|
|
197
|
+
|
|
198
|
+
Returns:
|
|
199
|
+
序列化后的数据结构,包含类型标记和序列化数据
|
|
200
|
+
"""
|
|
201
|
+
import io
|
|
202
|
+
|
|
203
|
+
if result is None:
|
|
204
|
+
return {SERIALIZED_MARKER: "none", "data": None}
|
|
205
|
+
|
|
206
|
+
# DataFrame: 转为 CSV 字符串
|
|
207
|
+
try:
|
|
208
|
+
import pandas as pd
|
|
209
|
+
if isinstance(result, pd.DataFrame):
|
|
210
|
+
csv_str = result.to_csv(index=True)
|
|
211
|
+
return {SERIALIZED_MARKER: "dataframe_csv", "data": csv_str}
|
|
212
|
+
except ImportError:
|
|
213
|
+
pass
|
|
214
|
+
|
|
215
|
+
# 字典: 检查是否包含 DataFrame(递归检查)
|
|
216
|
+
if isinstance(result, dict):
|
|
217
|
+
try:
|
|
218
|
+
import pandas as pd
|
|
219
|
+
|
|
220
|
+
def has_dataframe_recursive(obj):
|
|
221
|
+
"""递归检查对象中是否包含 DataFrame"""
|
|
222
|
+
if isinstance(obj, pd.DataFrame):
|
|
223
|
+
return True
|
|
224
|
+
if isinstance(obj, dict):
|
|
225
|
+
return any(has_dataframe_recursive(v) for v in obj.values())
|
|
226
|
+
if isinstance(obj, (list, tuple)):
|
|
227
|
+
return any(has_dataframe_recursive(item) for item in obj)
|
|
228
|
+
return False
|
|
229
|
+
|
|
230
|
+
def serialize_dataframes(obj):
|
|
231
|
+
"""递归序列化 DataFrame"""
|
|
232
|
+
if isinstance(obj, pd.DataFrame):
|
|
233
|
+
return {"__df__": True, "csv": obj.to_csv(index=True)}
|
|
234
|
+
if isinstance(obj, dict):
|
|
235
|
+
return {k: serialize_dataframes(v) for k, v in obj.items()}
|
|
236
|
+
if isinstance(obj, (list, tuple)):
|
|
237
|
+
return [serialize_dataframes(item) for item in obj]
|
|
238
|
+
return obj
|
|
239
|
+
|
|
240
|
+
if has_dataframe_recursive(result):
|
|
241
|
+
serialized_dict = serialize_dataframes(result)
|
|
242
|
+
json_str = json.dumps(serialized_dict, ensure_ascii=False, default=str)
|
|
243
|
+
return {SERIALIZED_MARKER: "dict_with_dataframe", "data": json_str}
|
|
244
|
+
except ImportError:
|
|
245
|
+
pass
|
|
246
|
+
|
|
247
|
+
# 普通字典: JSON 序列化
|
|
248
|
+
try:
|
|
249
|
+
json_str = json.dumps(result, ensure_ascii=False, default=str)
|
|
250
|
+
return {SERIALIZED_MARKER: "json", "data": json_str}
|
|
251
|
+
except (TypeError, ValueError):
|
|
252
|
+
pass
|
|
253
|
+
|
|
254
|
+
# 列表: JSON 序列化
|
|
255
|
+
if isinstance(result, (list, tuple)):
|
|
256
|
+
try:
|
|
257
|
+
json_str = json.dumps(result, ensure_ascii=False, default=str)
|
|
258
|
+
return {SERIALIZED_MARKER: "json", "data": json_str}
|
|
259
|
+
except (TypeError, ValueError):
|
|
260
|
+
# 无法 JSON 序列化,检查是否需要包装列表元素
|
|
261
|
+
# 对于包含复杂对象的列表,不进行序列化,让 RPyC 原样传输
|
|
262
|
+
pass
|
|
263
|
+
|
|
264
|
+
# 其他类型原样返回
|
|
265
|
+
return result
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
# ==================== 模块代理(带日志和权限检查) ====================
|
|
269
|
+
|
|
270
|
+
class LoggingProxy:
|
|
271
|
+
"""通用代理:拦截模块/对象的方法调用并记录日志,支持递归包装返回对象和权限检查"""
|
|
272
|
+
|
|
273
|
+
def __init__(self, target, target_name: str, client_info_getter, permission_checker=None, account_level=None):
|
|
274
|
+
object.__setattr__(self, '_target', target)
|
|
275
|
+
object.__setattr__(self, '_target_name', target_name)
|
|
276
|
+
object.__setattr__(self, '_get_client_info', client_info_getter)
|
|
277
|
+
object.__setattr__(self, '_permission_checker', permission_checker)
|
|
278
|
+
object.__setattr__(self, '_account_level', account_level)
|
|
279
|
+
|
|
280
|
+
def __getattr__(self, name):
|
|
281
|
+
target = object.__getattribute__(self, '_target')
|
|
282
|
+
target_name = object.__getattribute__(self, '_target_name')
|
|
283
|
+
get_client_info = object.__getattribute__(self, '_get_client_info')
|
|
284
|
+
permission_checker = object.__getattribute__(self, '_permission_checker')
|
|
285
|
+
account_level = object.__getattribute__(self, '_account_level')
|
|
286
|
+
|
|
287
|
+
attr = getattr(target, name)
|
|
288
|
+
|
|
289
|
+
# 如果是可调用对象,包装成带日志和权限检查的版本
|
|
290
|
+
if callable(attr):
|
|
291
|
+
def wrapper(*args, **kwargs):
|
|
292
|
+
full_name = f"{target_name}.{name}"
|
|
293
|
+
|
|
294
|
+
# 权限检查
|
|
295
|
+
if permission_checker and account_level:
|
|
296
|
+
error = permission_checker.check_api_permission(
|
|
297
|
+
account_level, full_name, args, kwargs
|
|
298
|
+
)
|
|
299
|
+
if error:
|
|
300
|
+
api_logger.warning(f"[权限拒绝] {full_name} | client={get_client_info()} | {error}")
|
|
301
|
+
raise error
|
|
302
|
+
|
|
303
|
+
result = _log_call(full_name, get_client_info(), attr, *args, **kwargs)
|
|
304
|
+
|
|
305
|
+
# 如果返回的是复杂对象(非基本类型),递归包装
|
|
306
|
+
if result is not None and hasattr(result, '__class__'):
|
|
307
|
+
if not isinstance(result, (int, float, str, bool, list, dict, tuple, type(None), bytes)):
|
|
308
|
+
if not result.__class__.__module__.startswith('builtins'):
|
|
309
|
+
return LoggingProxy(result, full_name, get_client_info, permission_checker, account_level)
|
|
310
|
+
|
|
311
|
+
# 处理列表:检查是否包含复杂对象
|
|
312
|
+
if isinstance(result, list):
|
|
313
|
+
wrapped_list = []
|
|
314
|
+
has_complex_obj = False
|
|
315
|
+
for item in result:
|
|
316
|
+
if item is not None and hasattr(item, '__class__'):
|
|
317
|
+
if not isinstance(item, (int, float, str, bool, dict, tuple, type(None), bytes)):
|
|
318
|
+
if not item.__class__.__module__.startswith('builtins'):
|
|
319
|
+
wrapped_list.append(LoggingProxy(item, full_name, get_client_info, permission_checker, account_level))
|
|
320
|
+
has_complex_obj = True
|
|
321
|
+
continue
|
|
322
|
+
wrapped_list.append(item)
|
|
323
|
+
if has_complex_obj:
|
|
324
|
+
return wrapped_list
|
|
325
|
+
|
|
326
|
+
# 序列化传输优化:将列表/字典/DataFrame 序列化以减少远程调用
|
|
327
|
+
return _serialize_for_transfer(result)
|
|
328
|
+
wrapper.__name__ = name
|
|
329
|
+
return wrapper
|
|
330
|
+
|
|
331
|
+
return attr
|
|
332
|
+
|
|
333
|
+
def __setattr__(self, name, value):
|
|
334
|
+
return setattr(object.__getattribute__(self, '_target'), name, value)
|
|
335
|
+
|
|
336
|
+
def __dir__(self):
|
|
337
|
+
return dir(object.__getattribute__(self, '_target'))
|
|
338
|
+
|
|
339
|
+
def __repr__(self):
|
|
340
|
+
return repr(object.__getattribute__(self, '_target'))
|
|
341
|
+
|
|
342
|
+
# 兼容别名
|
|
343
|
+
LoggingModuleProxy = LoggingProxy
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
# ==================== 服务类 ====================
|
|
347
|
+
|
|
348
|
+
class XtQuantService(rpyc.Service):
|
|
349
|
+
"""完全透明代理服务"""
|
|
350
|
+
|
|
351
|
+
_xtdata = xtdata
|
|
352
|
+
_xttrader = xttrader
|
|
353
|
+
_xttype = xttype
|
|
354
|
+
_xtconstant = xtconstant
|
|
355
|
+
_xtview = xtview
|
|
356
|
+
_permission_checker = None # 类级别的权限检查器
|
|
357
|
+
_datadir_reader = None # QmtDataReader 单例(server 级)
|
|
358
|
+
_datadir_path = None # datadir 路径(用于错误提示)
|
|
359
|
+
_datadir_error = None # 初始化错误信息(路径不存在等)
|
|
360
|
+
|
|
361
|
+
def on_connect(self, conn):
|
|
362
|
+
self._conn = conn
|
|
363
|
+
self._authenticated = False
|
|
364
|
+
self._client_id = None
|
|
365
|
+
self._account_level = AccountLevel.FREE # 默认为免费等级
|
|
366
|
+
self._traders = [] # 跟踪本次连接创建的所有 trader 实例
|
|
367
|
+
# 权限检查器在服务启动时已加载
|
|
368
|
+
# 兼容不同版本 rpyc:尝试获取客户端地址
|
|
369
|
+
try:
|
|
370
|
+
if hasattr(conn, 'peer'):
|
|
371
|
+
self._client_info = f"{conn.peer}"
|
|
372
|
+
elif hasattr(conn, '_channel') and hasattr(conn._channel, 'stream'):
|
|
373
|
+
stream = conn._channel.stream
|
|
374
|
+
if hasattr(stream, 'sock'):
|
|
375
|
+
peer = stream.sock.getpeername()
|
|
376
|
+
self._client_info = f"{peer[0]}:{peer[1]}"
|
|
377
|
+
else:
|
|
378
|
+
self._client_info = "unknown"
|
|
379
|
+
else:
|
|
380
|
+
self._client_info = "unknown"
|
|
381
|
+
except Exception:
|
|
382
|
+
self._client_info = "unknown"
|
|
383
|
+
logger.info(f"[连接] 客户端接入: {self._client_info}")
|
|
384
|
+
|
|
385
|
+
def on_disconnect(self, conn):
|
|
386
|
+
client_info = getattr(self, '_client_info', 'unknown')
|
|
387
|
+
logger.info(f"[断开] 客户端离开: {client_info}")
|
|
388
|
+
# 自动清理本次连接创建的所有 trader 实例,防止 session 资源泄漏
|
|
389
|
+
traders = getattr(self, '_traders', [])
|
|
390
|
+
for trader in traders:
|
|
391
|
+
try:
|
|
392
|
+
trader.stop()
|
|
393
|
+
logger.info(f"[清理Trader] 已自动 stop trader | client={client_info}")
|
|
394
|
+
except Exception as e:
|
|
395
|
+
logger.warning(f"[清理Trader] stop trader 失败: {e} | client={client_info}")
|
|
396
|
+
|
|
397
|
+
def _delayed_disconnect(self, delay: float = 0.5):
|
|
398
|
+
"""延迟断开连接,确保异常能传输到客户端"""
|
|
399
|
+
import threading
|
|
400
|
+
def _close():
|
|
401
|
+
try:
|
|
402
|
+
self._conn.close()
|
|
403
|
+
except:
|
|
404
|
+
pass
|
|
405
|
+
threading.Timer(delay, _close).start()
|
|
406
|
+
|
|
407
|
+
def _require_auth(self):
|
|
408
|
+
"""检查认证状态,未认证则抛出异常并断开连接"""
|
|
409
|
+
if not self._authenticated:
|
|
410
|
+
logger.warning(f"[未授权] 未认证的访问尝试: {self._client_info}")
|
|
411
|
+
self._delayed_disconnect()
|
|
412
|
+
raise AuthError("未授权访问,请先认证")
|
|
413
|
+
|
|
414
|
+
# ==================== 认证接口 ====================
|
|
415
|
+
|
|
416
|
+
@log_api_call("authenticate")
|
|
417
|
+
def exposed_authenticate(self, client_id, client_secret):
|
|
418
|
+
checker = XtQuantService._permission_checker
|
|
419
|
+
|
|
420
|
+
# 检查配置文件是否变更,如果变更则重新加载
|
|
421
|
+
checker.check_and_reload_if_changed()
|
|
422
|
+
|
|
423
|
+
# 验证密钥并获取账号等级
|
|
424
|
+
valid, account_level = checker.verify_secret(client_id, client_secret)
|
|
425
|
+
|
|
426
|
+
if not valid:
|
|
427
|
+
logger.warning(f"[认证失败] client_id={client_id}")
|
|
428
|
+
self._delayed_disconnect()
|
|
429
|
+
raise AuthError("认证失败:无效的客户端凭证")
|
|
430
|
+
|
|
431
|
+
self._authenticated = True
|
|
432
|
+
self._client_id = client_id
|
|
433
|
+
self._account_level = account_level
|
|
434
|
+
self._client_info = f"{client_id}@{self._client_info}"
|
|
435
|
+
logger.info(f"[认证成功] client_id={client_id} | level={account_level.value}")
|
|
436
|
+
return {"success": True, "level": account_level.value}
|
|
437
|
+
|
|
438
|
+
@log_api_call("heartbeat")
|
|
439
|
+
def exposed_heartbeat(self):
|
|
440
|
+
return "pong"
|
|
441
|
+
|
|
442
|
+
# ==================== 模块代理接口 ====================
|
|
443
|
+
|
|
444
|
+
@log_api_call("get_xtdata")
|
|
445
|
+
def exposed_get_xtdata(self):
|
|
446
|
+
self._require_auth()
|
|
447
|
+
return LoggingModuleProxy(
|
|
448
|
+
self._xtdata, 'xtdata',
|
|
449
|
+
lambda: self._client_info,
|
|
450
|
+
XtQuantService._permission_checker,
|
|
451
|
+
self._account_level
|
|
452
|
+
)
|
|
453
|
+
|
|
454
|
+
@log_api_call("get_xttype")
|
|
455
|
+
def exposed_get_xttype(self):
|
|
456
|
+
self._require_auth()
|
|
457
|
+
return self._xttype
|
|
458
|
+
|
|
459
|
+
def exposed_get_xtconstant(self):
|
|
460
|
+
self._require_auth()
|
|
461
|
+
return self._xtconstant
|
|
462
|
+
|
|
463
|
+
@log_api_call("get_datadir")
|
|
464
|
+
def exposed_get_datadir(self):
|
|
465
|
+
"""
|
|
466
|
+
获取 QmtDataReader 远程代理(文件解析能力)。
|
|
467
|
+
|
|
468
|
+
Returns:
|
|
469
|
+
LoggingProxy 包装的 QmtDataReader 实例
|
|
470
|
+
|
|
471
|
+
Raises:
|
|
472
|
+
AuthError: 未认证时抛出
|
|
473
|
+
RuntimeError: datadir 不可用时抛出(含路径信息)
|
|
474
|
+
"""
|
|
475
|
+
self._require_auth()
|
|
476
|
+
if XtQuantService._datadir_reader is None:
|
|
477
|
+
path_info = XtQuantService._datadir_path or "<未配置>"
|
|
478
|
+
err_info = XtQuantService._datadir_error or "QmtDataReader 未初始化"
|
|
479
|
+
raise RuntimeError(
|
|
480
|
+
f"datadir 不可用:{err_info}\n"
|
|
481
|
+
f"路径:{path_info}\n"
|
|
482
|
+
f"请在 server 端 .env 中配置 QMT_DATADIR_PATH,"
|
|
483
|
+
f"或确认 xtdata.get_data_dir() 可用。"
|
|
484
|
+
)
|
|
485
|
+
return LoggingProxy(
|
|
486
|
+
XtQuantService._datadir_reader,
|
|
487
|
+
'datadir',
|
|
488
|
+
lambda: self._client_info,
|
|
489
|
+
XtQuantService._permission_checker,
|
|
490
|
+
self._account_level,
|
|
491
|
+
)
|
|
492
|
+
|
|
493
|
+
@log_api_call("get_xtview")
|
|
494
|
+
def exposed_get_xtview(self):
|
|
495
|
+
self._require_auth()
|
|
496
|
+
if self._xtview is None:
|
|
497
|
+
raise RuntimeError("xtview 模块不可用,请检查 xtquant 版本是否支持")
|
|
498
|
+
return LoggingModuleProxy(
|
|
499
|
+
self._xtview, 'xtview',
|
|
500
|
+
lambda: self._client_info,
|
|
501
|
+
XtQuantService._permission_checker,
|
|
502
|
+
self._account_level
|
|
503
|
+
)
|
|
504
|
+
|
|
505
|
+
@log_api_call("create_trader")
|
|
506
|
+
def exposed_create_trader(self, userdata_path: str = None, session_id: int = None):
|
|
507
|
+
"""
|
|
508
|
+
创建交易实例(不自动启动,由客户端控制生命周期)
|
|
509
|
+
|
|
510
|
+
Args:
|
|
511
|
+
userdata_path: QMT 客户端 userdata_mini 目录路径(可选,可通过环境变量配置)
|
|
512
|
+
session_id: 会话ID(可选,默认自动生成时间戳)
|
|
513
|
+
|
|
514
|
+
Returns:
|
|
515
|
+
XtQuantTrader 实例(需客户端调用 start() 和 connect())
|
|
516
|
+
"""
|
|
517
|
+
self._require_auth()
|
|
518
|
+
# 检查 trade 权限
|
|
519
|
+
if self._account_level:
|
|
520
|
+
error = XtQuantService._permission_checker.check_api_permission(
|
|
521
|
+
self._account_level, "create_xttrader"
|
|
522
|
+
)
|
|
523
|
+
if error:
|
|
524
|
+
logger.warning(f"[权限拒绝] create_xttrader | client={self._client_info} | {error}")
|
|
525
|
+
raise error
|
|
526
|
+
if not XTQUANT_AVAILABLE:
|
|
527
|
+
raise RuntimeError("xtquant 库未安装")
|
|
528
|
+
|
|
529
|
+
# 从环境变量获取默认值
|
|
530
|
+
if userdata_path is None:
|
|
531
|
+
userdata_path = os.environ.get("QMT_USERDATA_PATH")
|
|
532
|
+
if userdata_path is None:
|
|
533
|
+
raise ValueError("必须提供 userdata_path 参数或设置 QMT_USERDATA_PATH 环境变量")
|
|
534
|
+
|
|
535
|
+
# 自动生成 session_id(使用毫秒级时间戳,避免同一秒内多次创建时冲突)
|
|
536
|
+
if session_id is None:
|
|
537
|
+
session_id = int(time.time() * 1000) % 1000000 # 毫秒级,取后6位避免超出int范围
|
|
538
|
+
|
|
539
|
+
# 创建 trader(不自动启动,由客户端控制生命周期)
|
|
540
|
+
trader = XtQuantTrader(userdata_path, session_id)
|
|
541
|
+
|
|
542
|
+
logger.info(f"[创建Trader] userdata_path={userdata_path} | session_id={session_id}")
|
|
543
|
+
|
|
544
|
+
# 记录 trader 实例,供 on_disconnect 时自动清理
|
|
545
|
+
self._traders.append(trader)
|
|
546
|
+
|
|
547
|
+
# 用 LoggingProxy 包装 trader,支持日志记录和序列化
|
|
548
|
+
return LoggingProxy(
|
|
549
|
+
trader, 'xttrader',
|
|
550
|
+
lambda: self._client_info,
|
|
551
|
+
XtQuantService._permission_checker,
|
|
552
|
+
self._account_level
|
|
553
|
+
)
|
|
554
|
+
|
|
555
|
+
# ==================== 辅助接口 ====================
|
|
556
|
+
|
|
557
|
+
@log_api_call("get_all_stocks")
|
|
558
|
+
def exposed_get_all_stocks(self):
|
|
559
|
+
self._require_auth()
|
|
560
|
+
return self._xtdata.get_stock_list_in_sector("沪深A股")
|
|
561
|
+
|
|
562
|
+
@log_api_call("get_index_list")
|
|
563
|
+
def exposed_get_index_list(self):
|
|
564
|
+
self._require_auth()
|
|
565
|
+
return self._xtdata.get_stock_list_in_sector("沪深指数")
|
|
566
|
+
|
|
567
|
+
# ==================== 服务端封装接口 ====================
|
|
568
|
+
|
|
569
|
+
@log_api_call("download_history_data2")
|
|
570
|
+
def exposed_download_history_data2(self, stock_list: list, period: str = "1d",
|
|
571
|
+
start_time: str = "", end_time: str = "", incrementally: bool = None):
|
|
572
|
+
"""
|
|
573
|
+
下载历史数据(服务端封装,避免回调传输问题)
|
|
574
|
+
返回: {'finished': n, 'total': n, 'result': {...}}
|
|
575
|
+
"""
|
|
576
|
+
status = {'finished': 0, 'total': 0, 'done': False, 'result': {}, 'message': ''}
|
|
577
|
+
|
|
578
|
+
def on_progress(data):
|
|
579
|
+
status['finished'] = data.get('finished', 0)
|
|
580
|
+
status['total'] = data.get('total', 0)
|
|
581
|
+
status['done'] = status['finished'] >= status['total']
|
|
582
|
+
status['message'] = data.get('message', '')
|
|
583
|
+
if 'result' in data:
|
|
584
|
+
import datetime as dt
|
|
585
|
+
from xtquant import xtbson as bson
|
|
586
|
+
regino_result = bson.BSON.decode(data.get('result'))
|
|
587
|
+
for stock, info in regino_result.items():
|
|
588
|
+
info['start_time'] = str(dt.datetime.fromtimestamp(info.get('start_time') / 1000))
|
|
589
|
+
info['end_time'] = str(dt.datetime.fromtimestamp(info.get('end_time') / 1000))
|
|
590
|
+
status['result'][stock] = info
|
|
591
|
+
|
|
592
|
+
# 调用原始方法(incrementally 参数需要转换为 None 或 bool)
|
|
593
|
+
inc = incrementally
|
|
594
|
+
self._xtdata.download_history_data2(
|
|
595
|
+
stock_list, period, start_time, end_time,
|
|
596
|
+
callback=on_progress, incrementally=inc
|
|
597
|
+
)
|
|
598
|
+
|
|
599
|
+
return status
|
|
600
|
+
|
|
601
|
+
# ==================== 服务状态 ====================
|
|
602
|
+
|
|
603
|
+
@log_api_call("get_service_status")
|
|
604
|
+
def exposed_get_service_status(self):
|
|
605
|
+
self._require_auth()
|
|
606
|
+
return {
|
|
607
|
+
"uptime": time.time() - getattr(self, '_start_time', time.time()),
|
|
608
|
+
"client_id": self._client_id,
|
|
609
|
+
}
|
|
610
|
+
|
|
611
|
+
@log_api_call("ping")
|
|
612
|
+
def exposed_ping(self):
|
|
613
|
+
return "pong"
|
|
614
|
+
|
|
615
|
+
@log_api_call("test_async_callback")
|
|
616
|
+
def exposed_test_async_callback(self, callback_func, delay: float = 2.0, count: int = 5):
|
|
617
|
+
"""
|
|
618
|
+
测试 RPyC netref 异步回调机制
|
|
619
|
+
:param callback_func: 客户端传递的回调函数(netref)
|
|
620
|
+
:param delay: 每次回调间隔秒数
|
|
621
|
+
:param count: 回调次数
|
|
622
|
+
:return: 立即返回 "已启动"
|
|
623
|
+
"""
|
|
624
|
+
self._require_auth()
|
|
625
|
+
# 检查 callback 权限
|
|
626
|
+
if self._account_level:
|
|
627
|
+
error = XtQuantService._permission_checker.check_api_permission(
|
|
628
|
+
self._account_level, "test_async_callback"
|
|
629
|
+
)
|
|
630
|
+
if error:
|
|
631
|
+
logger.warning(f"[权限拒绝] test_async_callback | client={self._client_info} | {error}")
|
|
632
|
+
raise error
|
|
633
|
+
|
|
634
|
+
import threading
|
|
635
|
+
import time
|
|
636
|
+
|
|
637
|
+
def async_call():
|
|
638
|
+
for i in range(count):
|
|
639
|
+
time.sleep(delay)
|
|
640
|
+
try:
|
|
641
|
+
result = callback_func(f"异步回调 #{i+1}/{count},时间: {time.strftime('%H:%M:%S')}")
|
|
642
|
+
api_logger.info(f"[异步回调] #{i+1} 执行成功,返回: {result}")
|
|
643
|
+
except Exception as e:
|
|
644
|
+
api_logger.error(f"[异步回调] #{i+1} 执行失败: {e}")
|
|
645
|
+
|
|
646
|
+
thread = threading.Thread(target=async_call, daemon=True)
|
|
647
|
+
thread.start()
|
|
648
|
+
return f"已启动异步回调,共 {count} 次,间隔 {delay} 秒"
|
|
649
|
+
|
|
650
|
+
|
|
651
|
+
# ==================== datadir 初始化 ====================
|
|
652
|
+
|
|
653
|
+
def _init_datadir_reader():
|
|
654
|
+
"""
|
|
655
|
+
初始化 QmtDataReader 单例(server 级,仅执行一次)。
|
|
656
|
+
|
|
657
|
+
路径解析优先级:
|
|
658
|
+
1. 环境变量 QMT_DATADIR_PATH(显式配置)
|
|
659
|
+
2. xtdata.get_data_dir() 自动推断(userdata_mini\\datadir → datadir)
|
|
660
|
+
3. 均不可用时记录 WARNING,_datadir_reader 保持 None
|
|
661
|
+
"""
|
|
662
|
+
if not QMTDATAREADER_AVAILABLE:
|
|
663
|
+
logger.warning("[datadir] QmtDataReader 不可用(utils.qmt_datadir 未安装),datadir 功能已禁用")
|
|
664
|
+
return
|
|
665
|
+
|
|
666
|
+
# 步骤1:从环境变量读取
|
|
667
|
+
datadir_path = os.environ.get("QMT_DATADIR_PATH", "").strip()
|
|
668
|
+
|
|
669
|
+
# 步骤2:自动推断
|
|
670
|
+
if not datadir_path and xtdata is not None:
|
|
671
|
+
try:
|
|
672
|
+
default_dir = xtdata.get_data_dir()
|
|
673
|
+
# 复用 env.py 中的替换逻辑:userdata_mini\datadir → datadir
|
|
674
|
+
inferred = default_dir.replace(r'\userdata_mini\datadir', r'\datadir')
|
|
675
|
+
if inferred != default_dir:
|
|
676
|
+
datadir_path = inferred
|
|
677
|
+
logger.info(f"[datadir] 自动推断 datadir 路径(主数据目录):{datadir_path}")
|
|
678
|
+
else:
|
|
679
|
+
datadir_path = default_dir
|
|
680
|
+
logger.info(f"[datadir] 自动推断 datadir 路径(默认目录):{datadir_path}")
|
|
681
|
+
except Exception as _e:
|
|
682
|
+
logger.warning(f"[datadir] xtdata.get_data_dir() 调用失败:{_e}")
|
|
683
|
+
|
|
684
|
+
XtQuantService._datadir_path = datadir_path or None
|
|
685
|
+
|
|
686
|
+
if not datadir_path:
|
|
687
|
+
logger.warning(
|
|
688
|
+
"[datadir] 未配置 QMT_DATADIR_PATH 且无法自动推断,datadir 功能不可用。\n"
|
|
689
|
+
"请在 .env 中添加:QMT_DATADIR_PATH=D:\\QMT\\datadir"
|
|
690
|
+
)
|
|
691
|
+
XtQuantService._datadir_error = "未配置 QMT_DATADIR_PATH 且无法自动推断路径"
|
|
692
|
+
return
|
|
693
|
+
|
|
694
|
+
if not os.path.isdir(datadir_path):
|
|
695
|
+
logger.error(f"[datadir] 配置的路径不存在:{datadir_path},datadir 功能不可用")
|
|
696
|
+
XtQuantService._datadir_error = f"路径不存在:{datadir_path}"
|
|
697
|
+
return
|
|
698
|
+
|
|
699
|
+
try:
|
|
700
|
+
XtQuantService._datadir_reader = QmtDataReader(datadir_path)
|
|
701
|
+
logger.info(f"[datadir] QmtDataReader 初始化成功:{datadir_path}")
|
|
702
|
+
except Exception as _e:
|
|
703
|
+
logger.error(f"[datadir] QmtDataReader 初始化失败:{_e}")
|
|
704
|
+
XtQuantService._datadir_error = str(_e)
|
|
705
|
+
|
|
706
|
+
|
|
707
|
+
# ==================== 服务启动 ====================
|
|
708
|
+
|
|
709
|
+
def create_ssl_context(certfile=None, keyfile=None):
|
|
710
|
+
if not certfile or not keyfile:
|
|
711
|
+
return None
|
|
712
|
+
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
|
713
|
+
ctx.load_cert_chain(certfile, keyfile)
|
|
714
|
+
return ctx
|
|
715
|
+
|
|
716
|
+
|
|
717
|
+
def start_server(host="0.0.0.0", port=None, use_ssl=False, certfile=None, keyfile=None, log_level="INFO", env_file=None):
|
|
718
|
+
"""启动服务
|
|
719
|
+
|
|
720
|
+
Args:
|
|
721
|
+
host: 监听地址
|
|
722
|
+
port: 监听端口
|
|
723
|
+
use_ssl: 是否启用 SSL
|
|
724
|
+
certfile: SSL 证书文件
|
|
725
|
+
keyfile: SSL 密钥文件
|
|
726
|
+
log_level: 日志级别
|
|
727
|
+
env_file: 环境变量文件路径(None 时自动查找 .env)
|
|
728
|
+
"""
|
|
729
|
+
# 加载环境变量文件(None 时自动查找 .env)
|
|
730
|
+
try:
|
|
731
|
+
from dotenv import load_dotenv
|
|
732
|
+
load_dotenv(env_file)
|
|
733
|
+
except ImportError:
|
|
734
|
+
pass
|
|
735
|
+
|
|
736
|
+
if port is None:
|
|
737
|
+
port = int(os.environ.get("XQSHARE_PORT", "18812"))
|
|
738
|
+
|
|
739
|
+
if not XTQUANT_AVAILABLE:
|
|
740
|
+
print("错误: xtquant 库未安装,请先安装 xtquant")
|
|
741
|
+
return
|
|
742
|
+
|
|
743
|
+
_init_logging(log_level)
|
|
744
|
+
XtQuantService._start_time = time.time()
|
|
745
|
+
|
|
746
|
+
print("=" * 70)
|
|
747
|
+
print(" XtQuant Share (xqshare) 服务")
|
|
748
|
+
print("=" * 70)
|
|
749
|
+
print(f" 监听地址: {host}:{port}")
|
|
750
|
+
print(f" SSL 加密: {'启用' if use_ssl else '禁用'}")
|
|
751
|
+
print(f" 日志级别: {log_level}")
|
|
752
|
+
print("=" * 70)
|
|
753
|
+
|
|
754
|
+
# 预加载权限检查器(加载 clients.yaml 配置)
|
|
755
|
+
if XtQuantService._permission_checker is None:
|
|
756
|
+
XtQuantService._permission_checker = get_permission_checker()
|
|
757
|
+
|
|
758
|
+
# ── 初始化 QmtDataReader 单例 ──────────────────────────────────
|
|
759
|
+
_init_datadir_reader()
|
|
760
|
+
|
|
761
|
+
logger.info(f"服务启动 | host={host} | port={port} | ssl={use_ssl}")
|
|
762
|
+
|
|
763
|
+
config = {
|
|
764
|
+
'allow_public_attrs': True,
|
|
765
|
+
'allow_pickle': True,
|
|
766
|
+
'allow_getattr': True,
|
|
767
|
+
'allow_setattr': True,
|
|
768
|
+
'allow_delattr': True,
|
|
769
|
+
'allow_all_attrs': True,
|
|
770
|
+
'sync_request_timeout': 300,
|
|
771
|
+
}
|
|
772
|
+
|
|
773
|
+
ssl_context = None
|
|
774
|
+
if use_ssl:
|
|
775
|
+
ssl_context = create_ssl_context(certfile, keyfile)
|
|
776
|
+
if ssl_context:
|
|
777
|
+
logger.info("SSL 证书加载成功")
|
|
778
|
+
print(" ✓ SSL 证书加载成功")
|
|
779
|
+
else:
|
|
780
|
+
logger.warning("SSL 证书加载失败")
|
|
781
|
+
print(" ⚠ SSL 证书加载失败")
|
|
782
|
+
|
|
783
|
+
# 构建 ThreadedServer 参数(兼容不同 rpyc 版本)
|
|
784
|
+
server_kwargs = {
|
|
785
|
+
'hostname': host,
|
|
786
|
+
'port': port,
|
|
787
|
+
'protocol_config': config,
|
|
788
|
+
}
|
|
789
|
+
|
|
790
|
+
# 尝试使用 ssl_context(新版本 rpyc)
|
|
791
|
+
try:
|
|
792
|
+
server = ThreadedServer(XtQuantService, ssl_context=ssl_context, **server_kwargs)
|
|
793
|
+
except TypeError:
|
|
794
|
+
# 旧版本 rpyc 不支持 ssl_context,使用其他方式
|
|
795
|
+
if ssl_context:
|
|
796
|
+
# 对于旧版本,通过 protocol_config 传递 SSL
|
|
797
|
+
import socket
|
|
798
|
+
import ssl as ssl_module
|
|
799
|
+
|
|
800
|
+
# 创建 SSL 包装的 socket
|
|
801
|
+
class SSLThreadedServer(ThreadedServer):
|
|
802
|
+
def _accept_method(self, sock):
|
|
803
|
+
try:
|
|
804
|
+
return ssl_context.wrap_socket(sock, server_side=True)
|
|
805
|
+
except Exception as e:
|
|
806
|
+
logger.error(f"SSL 包装失败: {e}")
|
|
807
|
+
raise
|
|
808
|
+
|
|
809
|
+
server = SSLThreadedServer(XtQuantService, **server_kwargs)
|
|
810
|
+
logger.info("使用兼容模式启动 SSL")
|
|
811
|
+
else:
|
|
812
|
+
server = ThreadedServer(XtQuantService, **server_kwargs)
|
|
813
|
+
|
|
814
|
+
print("\n 服务已启动,等待客户端连接...")
|
|
815
|
+
print(" 按 Ctrl+C 停止服务\n")
|
|
816
|
+
|
|
817
|
+
try:
|
|
818
|
+
server.start()
|
|
819
|
+
except KeyboardInterrupt:
|
|
820
|
+
logger.info("服务停止(用户中断)")
|
|
821
|
+
print("\n 服务已停止")
|
|
822
|
+
server.close()
|
|
823
|
+
except Exception as e:
|
|
824
|
+
logger.error(f"服务异常: {e}")
|
|
825
|
+
raise
|
|
826
|
+
|
|
827
|
+
|
|
828
|
+
def main():
|
|
829
|
+
"""命令行入口函数"""
|
|
830
|
+
import argparse
|
|
831
|
+
|
|
832
|
+
parser = argparse.ArgumentParser(
|
|
833
|
+
description="XtQuant Share (xqshare) 服务",
|
|
834
|
+
formatter_class=argparse.RawDescriptionHelpFormatter,
|
|
835
|
+
epilog="""
|
|
836
|
+
示例:
|
|
837
|
+
xqshare-server # 使用默认配置启动
|
|
838
|
+
xqshare-server --port 18813 # 指定端口
|
|
839
|
+
xqshare-server --ssl --cert cert.pem --key key.pem # 启用 SSL
|
|
840
|
+
|
|
841
|
+
环境变量:
|
|
842
|
+
XQSHARE_PORT 服务端口 (默认: 18812)
|
|
843
|
+
QMT_USERDATA_PATH QMT userdata_mini 目录路径
|
|
844
|
+
"""
|
|
845
|
+
)
|
|
846
|
+
parser.add_argument("--host", default="0.0.0.0", help="监听地址 (默认: 0.0.0.0)")
|
|
847
|
+
parser.add_argument("--port", type=int, default=None, help="监听端口 (默认: 18812 或 XQSHARE_PORT)")
|
|
848
|
+
parser.add_argument("--ssl", action="store_true", help="启用 SSL 加密")
|
|
849
|
+
parser.add_argument("--cert", help="SSL 证书文件")
|
|
850
|
+
parser.add_argument("--key", help="SSL 私钥文件")
|
|
851
|
+
parser.add_argument("--log-level", default="INFO", help="日志级别 (默认: INFO)")
|
|
852
|
+
parser.add_argument("--env-file", default=".env", help="环境变量文件 (默认: .env)")
|
|
853
|
+
|
|
854
|
+
args = parser.parse_args()
|
|
855
|
+
|
|
856
|
+
start_server(
|
|
857
|
+
host=args.host,
|
|
858
|
+
port=args.port,
|
|
859
|
+
use_ssl=args.ssl,
|
|
860
|
+
certfile=args.cert,
|
|
861
|
+
keyfile=args.key,
|
|
862
|
+
log_level=args.log_level,
|
|
863
|
+
env_file=args.env_file
|
|
864
|
+
)
|
|
865
|
+
|
|
866
|
+
|
|
867
|
+
if __name__ == "__main__":
|
|
868
|
+
main()
|