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/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()